Keep checkpoints aligned with PromptKit profiles

This commit is contained in:
2026-07-28 13:36:01 -05:00
parent f1a6574013
commit de046a8f13
14 changed files with 308 additions and 399 deletions

View File

@@ -20,6 +20,7 @@ import (
"gitea.maximumdirect.net/eric/notarius/internal/core/config"
"gitea.maximumdirect.net/eric/notarius/internal/framework/chunkmap"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/llm"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd"
"gitea.maximumdirect.net/eric/notarius/internal/modules/dnd/chunk/scenes"
@@ -390,6 +391,14 @@ func TestProductionLLMClientFactoriesBuildOfflineRuntime(t *testing.T) {
if len(manifests) != 0 {
t.Fatalf("eager profile manifests = %#v, want none", manifests)
}
fingerprintProvider, ok := client.(llm.CheckpointFingerprintProvider)
if !ok {
t.Fatalf("production LLM client %T does not provide one profile-source checkpoint fingerprint", client)
}
fingerprints, err := fingerprintProvider.LLMCheckpointFingerprints()
if err != nil || len(fingerprints) != 1 {
t.Fatalf("production LLM checkpoint fingerprints = %#v, error = %v, want one profile-source identity", fingerprints, err)
}
if _, ok := client.(contracts.LLMProfileManifestProvider); !ok {
t.Fatalf("production LLM client %T does not provide profile manifests", client)
}

View File

@@ -23,6 +23,7 @@ import (
"gitea.maximumdirect.net/eric/notarius/internal/framework/chunkplan"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
frameworkdebug "gitea.maximumdirect.net/eric/notarius/internal/framework/debug"
frameworkllm "gitea.maximumdirect.net/eric/notarius/internal/framework/llm"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
)
@@ -367,6 +368,13 @@ func runPipelineCommand(args []string, stdout, stderr io.Writer, opts Options) i
if err != nil {
return failPipelineCommand(stderr, commandState, terminalWriter, fmt.Errorf("create LLM client for profile %q: %w", factoryProfileID, err))
}
var llmFingerprints []checkpoint.Fingerprint
if effective.Config.Cache.Checkpoints.Enabled {
llmFingerprints, err = llmCheckpointFingerprints(llmClient)
if err != nil {
return failPipelineCommand(stderr, commandState, terminalWriter, fmt.Errorf("prepare LLM checkpoint identity: %w", err))
}
}
llmClient = pipeline.WithDebugLLMRecording(llmClient, debugRecorder)
prepared, err := pipeline.Prepare(effective.ResolvedPipeline, registries, pipeline.ModuleDependencies{LLM: llmClient})
if err != nil {
@@ -380,7 +388,7 @@ func runPipelineCommand(args []string, stdout, stderr io.Writer, opts Options) i
if err != nil {
return failPipelineCommand(stderr, commandState, terminalWriter, err)
}
checkpointRecorder, checkpointLoader, err := checkpointHandlersForRun(effective.Config.Cache.Checkpoints, opts, effective.ResolvedPipeline, prepared.CheckpointFingerprints(), rawInput, only, llmProfiles, strings.TrimSpace(*llmProfile), strings.TrimSpace(sessionID.value), *resume)
checkpointRecorder, checkpointLoader, err := checkpointHandlersForRun(effective.Config.Cache.Checkpoints, opts, effective.ResolvedPipeline, prepared.CheckpointFingerprints(), llmFingerprints, rawInput, only, llmProfiles, strings.TrimSpace(*llmProfile), strings.TrimSpace(sessionID.value), *resume)
if err != nil {
return failPipelineCommand(stderr, commandState, terminalWriter, err)
}
@@ -481,6 +489,7 @@ func checkpointHandlersForRun(
opts Options,
resolved pipeline.ResolvedPipeline,
componentFingerprints []pipeline.CheckpointFingerprint,
llmFingerprints []checkpoint.Fingerprint,
rawInput []byte,
only []string,
llmProfiles []artifacts.LLMProfileManifest,
@@ -495,13 +504,17 @@ func checkpointHandlersForRun(
return pipeline.NoopCheckpointRecorder(), pipeline.NoopCheckpointLoader(), nil
}
identity, err := checkpoint.NewIdentity(checkpoint.IdentityInput{
Pipeline: resolved,
InputKey: resolved.Input.Module,
RawInputDigest: rawInputDigest(rawInput),
SelectedLanes: only,
RuntimeOverrides: runtimeOverrideFingerprints(llmProfileOverride, sessionID),
References: pipeline.ReferenceProvenance(resolved),
ProvenanceFingerprints: append(llmProfileFingerprints(llmProfiles), checkpointIdentityFingerprints(componentFingerprints)...),
Pipeline: resolved,
InputKey: resolved.Input.Module,
RawInputDigest: rawInputDigest(rawInput),
SelectedLanes: only,
RuntimeOverrides: runtimeOverrideFingerprints(llmProfileOverride, sessionID),
References: pipeline.ReferenceProvenance(resolved),
ProvenanceFingerprints: combineCheckpointFingerprints(
llmProfileFingerprints(llmProfiles),
llmFingerprints,
checkpointIdentityFingerprints(componentFingerprints),
),
})
if err != nil {
return nil, nil, fmt.Errorf("create checkpoint identity: %w", err)
@@ -527,6 +540,30 @@ func checkpointHandlersForRun(
return recorder, loader, nil
}
func llmCheckpointFingerprints(client contracts.StructuredLLMClient) ([]checkpoint.Fingerprint, error) {
provider, ok := client.(frameworkllm.CheckpointFingerprintProvider)
if !ok {
return nil, nil
}
values, err := provider.LLMCheckpointFingerprints()
if err != nil {
return nil, err
}
out := make([]checkpoint.Fingerprint, 0, len(values))
for _, value := range values {
out = append(out, checkpoint.Fingerprint{Name: value.Name, Value: value.Value})
}
return out, nil
}
func combineCheckpointFingerprints(sources ...[]checkpoint.Fingerprint) []checkpoint.Fingerprint {
var out []checkpoint.Fingerprint
for _, source := range sources {
out = append(out, source...)
}
return out
}
func recomputePolicy(resolved pipeline.ResolvedPipeline, requestedStep string) (pipeline.CheckpointExecutionPolicy, error) {
requestedStep = strings.TrimSpace(requestedStep)
if requestedStep == "" {

View File

@@ -214,8 +214,9 @@ func TestChangedSemanticSpellCatalogFingerprintCannotResumeRecordedCheckpoint(t
t.Fatal(err)
}
fingerprints := prepared.CheckpointFingerprints()
llmFingerprints := []checkpoint.Fingerprint{{Name: "promptkit_profile_source", Value: "sha256:profile-source-one"}}
settings := config.CheckpointCacheConfig{Enabled: true, Directory: t.TempDir()}
recorder, _, err := checkpointHandlersForRun(settings, Options{}, materialized, fingerprints, []byte("same input"), nil, nil, "", "", false)
recorder, _, err := checkpointHandlersForRun(settings, Options{}, materialized, fingerprints, llmFingerprints, []byte("same input"), nil, nil, "", "", false)
if err != nil {
t.Fatal(err)
}
@@ -241,7 +242,7 @@ func TestChangedSemanticSpellCatalogFingerprintCannotResumeRecordedCheckpoint(t
t.Fatal(err)
}
_, sameLoader, err := checkpointHandlersForRun(settings, Options{}, materialized, fingerprints, []byte("same input"), nil, nil, "", "", true)
_, sameLoader, err := checkpointHandlersForRun(settings, Options{}, materialized, fingerprints, llmFingerprints, []byte("same input"), nil, nil, "", "", true)
if err != nil {
t.Fatal(err)
}
@@ -253,7 +254,7 @@ func TestChangedSemanticSpellCatalogFingerprintCannotResumeRecordedCheckpoint(t
}
changed := replaceCheckpointFingerprintValue(t, fingerprints, normalizeSpellCatalogFingerprintName(), "sha256:changed-effective-catalog")
assertOnlyCheckpointFingerprintChanged(t, fingerprints, changed, normalizeSpellCatalogFingerprintName())
_, changedLoader, err := checkpointHandlersForRun(settings, Options{}, materialized, changed, []byte("same input"), nil, nil, "", "", true)
_, changedLoader, err := checkpointHandlersForRun(settings, Options{}, materialized, changed, llmFingerprints, []byte("same input"), nil, nil, "", "", true)
if err != nil {
t.Fatal(err)
}
@@ -265,13 +266,25 @@ func TestChangedSemanticSpellCatalogFingerprintCannotResumeRecordedCheckpoint(t
}
changedMapping := replaceCheckpointFingerprintValue(t, fingerprints, extractSpellMappingFingerprintName(), "dnd.spells.extract_mapping.v3")
assertOnlyCheckpointFingerprintChanged(t, fingerprints, changedMapping, extractSpellMappingFingerprintName())
_, mappingLoader, err := checkpointHandlersForRun(settings, Options{}, materialized, changedMapping, []byte("same input"), nil, nil, "", "", true)
_, mappingLoader, err := checkpointHandlersForRun(settings, Options{}, materialized, changedMapping, llmFingerprints, []byte("same input"), nil, nil, "", "", true)
if err != nil {
t.Fatal(err)
}
if _, decision := mappingLoader.Source(materialized.Input.Module); decision.Reused {
t.Fatalf("changed mapping policy decision = %#v, want cold miss", decision)
}
changedLLMFingerprints := []checkpoint.Fingerprint{{Name: "promptkit_profile_source", Value: "sha256:profile-source-two"}}
_, profileLoader, err := checkpointHandlersForRun(settings, Options{}, materialized, fingerprints, changedLLMFingerprints, []byte("same input"), nil, nil, "", "", true)
if err != nil {
t.Fatal(err)
}
if _, decision := profileLoader.Source(materialized.Input.Module); decision.Reused {
t.Fatalf("changed PromptKit profile source decision = %#v, want cold miss", decision)
}
if _, decision := profileLoader.Normalize("spells", spellnormalize.Key, normalizeDependencies); decision.Reused {
t.Fatalf("changed PromptKit profile normalize decision = %#v, want cold miss", decision)
}
}
func normalizeSpellCatalogFingerprintName() string {

View File

@@ -0,0 +1,14 @@
package llm
// CheckpointFingerprint is a stable, non-secret semantic identity contributed
// by the LLM runtime before pipeline execution.
type CheckpointFingerprint struct {
Name string
Value string
}
// CheckpointFingerprintProvider exposes LLM-runtime identities that must
// participate in checkpoint composition.
type CheckpointFingerprintProvider interface {
LLMCheckpointFingerprints() ([]CheckpointFingerprint, error)
}

View File

@@ -29,8 +29,10 @@ type PromptKitClientConfig struct {
}
type PromptKitClient struct {
engine *promptkit.Engine
recorder *LLMProfileRecorder
engine *promptkit.Engine
recorder *LLMProfileRecorder
profileDir string
profileFile string
}
type LLMProfileRecorder struct {
@@ -70,8 +72,10 @@ func NewPromptKitClient(cfg PromptKitClientConfig) (*PromptKitClient, error) {
recorder = NewLLMProfileRecorder()
}
return &PromptKitClient{
engine: engine,
recorder: recorder,
engine: engine,
recorder: recorder,
profileDir: strings.TrimSpace(cfg.ProfileDir),
profileFile: strings.TrimSpace(cfg.ProfileFile),
}, nil
}
@@ -261,6 +265,17 @@ func (c *PromptKitClient) LLMProfileManifests() []artifacts.LLMProfileManifest {
return c.recorder.Manifests()
}
func (c *PromptKitClient) LLMCheckpointFingerprints() ([]CheckpointFingerprint, error) {
if c == nil {
return nil, nil
}
fingerprint, err := promptKitProfileFingerprint(c.profileDir, c.profileFile)
if err != nil {
return nil, err
}
return []CheckpointFingerprint{fingerprint}, nil
}
func NewLLMProfileRecorder() *LLMProfileRecorder {
return &LLMProfileRecorder{profiles: map[string]artifacts.LLMProfileManifest{}}
}

View File

@@ -6,6 +6,8 @@ import (
"errors"
"io"
"net/http"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
@@ -119,6 +121,63 @@ func TestNewPromptKitClientReportsAssetAndEngineConstructionFailures(t *testing.
})
}
func TestPromptKitClientCheckpointFingerprintTracksProfileSource(t *testing.T) {
profilePath := filepath.Join(t.TempDir(), "profiles.yml")
writeProfile := func(model string) {
t.Helper()
content := "id: checkpoint-profile\nendpoint: http://promptkit.test/v1\nmodel: " + model + "\n"
if err := os.WriteFile(profilePath, []byte(content), 0o600); err != nil {
t.Fatal(err)
}
}
fingerprintFor := func() CheckpointFingerprint {
t.Helper()
client, err := NewPromptKitClient(PromptKitClientConfig{
Assets: newTestPromptKitAssets(t),
ProfileFile: profilePath,
})
if err != nil {
t.Fatal(err)
}
values, err := client.LLMCheckpointFingerprints()
if err != nil {
t.Fatal(err)
}
if len(values) != 1 || values[0].Name != promptKitProfileFingerprintName {
t.Fatalf("checkpoint fingerprints = %#v, want one profile-source identity", values)
}
return values[0]
}
writeProfile("model-one")
first := fingerprintFor()
writeProfile("model-two")
second := fingerprintFor()
if first == second {
t.Fatalf("profile-source fingerprint = %#v for both profile models", first)
}
if strings.Contains(first.Value, profilePath) || strings.Contains(first.Value, "model-one") {
t.Fatalf("profile-source fingerprint exposes source details: %#v", first)
}
client, err := NewPromptKitClient(PromptKitClientConfig{Assets: newTestPromptKitAssets(t)})
if err != nil {
t.Fatal(err)
}
copy, err := client.LLMCheckpointFingerprints()
if err != nil {
t.Fatal(err)
}
copy[0].Value = "mutated"
fresh, err := client.LLMCheckpointFingerprints()
if err != nil {
t.Fatal(err)
}
if fresh[0].Value == "mutated" {
t.Fatal("LLMCheckpointFingerprints exposed mutable backing storage")
}
}
func TestPromptKitClientUsesPromptDefaultProfileWhenRequestProfileEmpty(t *testing.T) {
fake := &fakePromptKitLLM{content: `{"ok":true}`}
client := newTestPromptKitClient(t, fake)

View File

@@ -0,0 +1,82 @@
package llm
import (
"crypto/sha256"
"encoding/hex"
"fmt"
"io/fs"
"os"
"path/filepath"
"sort"
"strings"
)
const (
promptKitProfileFingerprintName = "promptkit_profile_source"
// The built-in profile catalog is compiled into this pinned PromptKit
// release. Update this identity when the dependency is upgraded.
promptKitBuiltinProfileCatalogID = "promptkit:v0.1.0:builtin-profiles"
)
func promptKitProfileFingerprint(profileDir, profileFile string) (CheckpointFingerprint, error) {
hasher := sha256.New()
writeFingerprintPart(hasher, []byte(promptKitBuiltinProfileCatalogID))
switch {
case strings.TrimSpace(profileFile) != "":
data, err := os.ReadFile(strings.TrimSpace(profileFile))
if err != nil {
return CheckpointFingerprint{}, fmt.Errorf("read PromptKit profile file for checkpoint identity: %w", err)
}
writeFingerprintPart(hasher, data)
case strings.TrimSpace(profileDir) != "":
digests, err := promptKitProfileFileDigests(strings.TrimSpace(profileDir))
if err != nil {
return CheckpointFingerprint{}, err
}
for _, digest := range digests {
writeFingerprintPart(hasher, digest)
}
}
return CheckpointFingerprint{
Name: promptKitProfileFingerprintName,
Value: "sha256:" + hex.EncodeToString(hasher.Sum(nil)),
}, nil
}
func promptKitProfileFileDigests(root string) ([][]byte, error) {
var digests [][]byte
err := filepath.WalkDir(root, func(name string, entry fs.DirEntry, walkErr error) error {
if walkErr != nil {
return walkErr
}
if entry.IsDir() {
return nil
}
extension := filepath.Ext(entry.Name())
if extension != ".yaml" && extension != ".yml" {
return nil
}
data, err := os.ReadFile(name)
if err != nil {
return err
}
sum := sha256.Sum256(data)
digests = append(digests, append([]byte(nil), sum[:]...))
return nil
})
if err != nil {
return nil, fmt.Errorf("read PromptKit profile directory for checkpoint identity: %w", err)
}
sort.Slice(digests, func(i, j int) bool {
return string(digests[i]) < string(digests[j])
})
return digests, nil
}
func writeFingerprintPart(hasher interface{ Write([]byte) (int, error) }, value []byte) {
length := []byte(fmt.Sprintf("%d:", len(value)))
_, _ = hasher.Write(length)
_, _ = hasher.Write(value)
}

View File

@@ -53,3 +53,14 @@ func (c *scheduledClient) LLMProfileManifests() []artifacts.LLMProfileManifest {
}
return provider.LLMProfileManifests()
}
func (c *scheduledClient) LLMCheckpointFingerprints() ([]CheckpointFingerprint, error) {
if c == nil || c.client == nil {
return nil, nil
}
provider, ok := c.client.(CheckpointFingerprintProvider)
if !ok {
return nil, nil
}
return provider.LLMCheckpointFingerprints()
}

View File

@@ -75,6 +75,28 @@ func TestScheduledClientPropagatesSchedulerError(t *testing.T) {
}
}
func TestScheduledClientPreservesCheckpointFingerprints(t *testing.T) {
scheduler, err := NewScheduler(1)
if err != nil {
t.Fatal(err)
}
inner := &fingerprintedStructuredClient{
fingerprints: []CheckpointFingerprint{{Name: "profile_source", Value: "sha256:one"}},
}
client := NewScheduledClient(inner, scheduler)
provider, ok := client.(CheckpointFingerprintProvider)
if !ok {
t.Fatalf("scheduled client %T does not preserve checkpoint fingerprints", client)
}
got, err := provider.LLMCheckpointFingerprints()
if err != nil {
t.Fatal(err)
}
if len(got) != 1 || got[0] != inner.fingerprints[0] {
t.Fatalf("checkpoint fingerprints = %#v, want %#v", got, inner.fingerprints)
}
}
type blockingStructuredClient struct {
release chan struct{}
inFlight int32
@@ -110,6 +132,15 @@ type errorStructuredClient struct {
err error
}
type fingerprintedStructuredClient struct {
errorStructuredClient
fingerprints []CheckpointFingerprint
}
func (c *fingerprintedStructuredClient) LLMCheckpointFingerprints() ([]CheckpointFingerprint, error) {
return append([]CheckpointFingerprint(nil), c.fingerprints...), nil
}
func (c *errorStructuredClient) CompleteStructured(ctx context.Context, req contracts.StructuredCompletionRequest, out any) (contracts.StructuredCompletionResponse, error) {
return contracts.StructuredCompletionResponse{}, c.err
}

View File

@@ -11,15 +11,15 @@ import (
"time"
)
// ModulePromptFile maps a module-owned embedded prompt file into the
// PromptKit-visible module prompt directory.
// ModulePromptFile maps a module-owned embedded prompt file into the registered
// module prompt directory.
type ModulePromptFile struct {
Name string
Path string
}
// SharedPromptFile maps a caller-owned shared prompt file into a module's
// PromptKit-visible sharedassets prompt subdirectory.
// registered sharedassets prompt subdirectory.
type SharedPromptFile struct {
Name string
FS fs.FS