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

@@ -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 == "" {