Include local backend target in checkpoint identity

This commit is contained in:
2026-07-30 05:17:39 +00:00
parent 715fff7b72
commit 241e9d2a89
4 changed files with 106 additions and 4 deletions

View File

@@ -41,6 +41,7 @@ type PromptKitClient struct {
recorder *LLMProfileRecorder
profileDir string
profileFile string
localEndpoint string
reasoningEffort *string
}
@@ -70,9 +71,11 @@ func NewPromptKitClient(cfg PromptKitClientConfig) (*PromptKitClient, error) {
if profileFile := strings.TrimSpace(cfg.ProfileFile); profileFile != "" {
options = append(options, promptkit.WithProfileFile(profileFile))
}
var localEndpoint string
if cfg.LocalBackend != nil {
localBackend := *cfg.LocalBackend
localBackend.Endpoint = strings.TrimSpace(localBackend.Endpoint)
localEndpoint = localBackend.Endpoint
options = append(options, PromptKitLocalBackendOption(localBackend))
}
options = append(options, cfg.EngineOptions...)
@@ -99,6 +102,7 @@ func NewPromptKitClient(cfg PromptKitClientConfig) (*PromptKitClient, error) {
recorder: recorder,
profileDir: strings.TrimSpace(cfg.ProfileDir),
profileFile: strings.TrimSpace(cfg.ProfileFile),
localEndpoint: localEndpoint,
reasoningEffort: reasoningEffort,
}, nil
}
@@ -317,7 +321,11 @@ func (c *PromptKitClient) LLMCheckpointFingerprints() ([]CheckpointFingerprint,
if err != nil {
return nil, err
}
return []CheckpointFingerprint{fingerprint}, nil
fingerprints := []CheckpointFingerprint{fingerprint}
if c.localEndpoint != "" {
fingerprints = append(fingerprints, promptKitLocalBackendFingerprint(c.localEndpoint))
}
return fingerprints, nil
}
func NewLLMProfileRecorder() *LLMProfileRecorder {

View File

@@ -312,6 +312,83 @@ func TestPromptKitClientCheckpointFingerprintTracksProfileSource(t *testing.T) {
}
}
func TestPromptKitClientCheckpointFingerprintTracksLocalBackendTarget(t *testing.T) {
const (
firstEndpoint = "http://localhost:8000/v1"
secondEndpoint = "https://inference.example.test/v1"
)
fingerprintsFor := func(endpoint string, concurrencyLimit int) []CheckpointFingerprint {
t.Helper()
client, err := NewPromptKitClient(PromptKitClientConfig{
Assets: newTestPromptKitAssets(t),
LocalBackend: &PromptKitLocalBackendConfig{
Endpoint: endpoint,
ConcurrencyLimit: concurrencyLimit,
},
})
if err != nil {
t.Fatal(err)
}
values, err := client.LLMCheckpointFingerprints()
if err != nil {
t.Fatal(err)
}
return values
}
baseline := fingerprintsFor(firstEndpoint, 0)
if len(baseline) != 2 ||
baseline[0].Name != promptKitProfileFingerprintName ||
baseline[1].Name != promptKitLocalBackendFingerprintName {
t.Fatalf("checkpoint fingerprints = %#v, want profile source then local backend target", baseline)
}
endpointChanged := fingerprintsFor(secondEndpoint, 0)
if baseline[0] != endpointChanged[0] || baseline[1] == endpointChanged[1] {
t.Fatalf("endpoint fingerprints = %#v and %#v, want only local target to change", baseline, endpointChanged)
}
concurrencyChanged := fingerprintsFor(firstEndpoint, 4)
if !reflect.DeepEqual(baseline, concurrencyChanged) {
t.Fatalf("concurrency fingerprints = %#v, want %#v", concurrencyChanged, baseline)
}
for _, values := range [][]CheckpointFingerprint{baseline, endpointChanged} {
for _, fingerprint := range values {
if strings.Contains(fingerprint.Value, firstEndpoint) ||
strings.Contains(fingerprint.Value, secondEndpoint) {
t.Fatalf("checkpoint fingerprint exposes endpoint: %#v", fingerprint)
}
}
}
localBackend := &PromptKitLocalBackendConfig{
Endpoint: " " + firstEndpoint + " ",
ConcurrencyLimit: 0,
}
client, err := NewPromptKitClient(PromptKitClientConfig{
Assets: newTestPromptKitAssets(t),
LocalBackend: localBackend,
})
if err != nil {
t.Fatal(err)
}
localBackend.Endpoint = secondEndpoint
copy, err := client.LLMCheckpointFingerprints()
if err != nil {
t.Fatal(err)
}
if !reflect.DeepEqual(copy, baseline) {
t.Fatalf("fingerprints after input mutation = %#v, want retained target %#v", copy, baseline)
}
copy[0].Value = "mutated-profile"
copy[1].Value = "mutated-target"
fresh, err := client.LLMCheckpointFingerprints()
if err != nil {
t.Fatal(err)
}
if !reflect.DeepEqual(fresh, baseline) {
t.Fatalf("fingerprints after returned-slice mutation = %#v, want %#v", fresh, baseline)
}
}
func TestPromptKitClientUsesPromptDefaultProfileWhenRequestProfileEmpty(t *testing.T) {
fake := &fakePromptKitLLM{content: `{"ok":true}`}
client := newTestPromptKitClient(t, fake)

View File

@@ -12,7 +12,9 @@ import (
)
const (
promptKitProfileFingerprintName = "promptkit_profile_source"
promptKitProfileFingerprintName = "promptkit_profile_source"
promptKitLocalBackendFingerprintName = "promptkit_local_backend_target"
promptKitLocalBackendMarker = "notarius:promptkit-local-backend:v1"
// The built-in profile catalog is compiled into this pinned PromptKit
// release. Update this identity when the dependency is upgraded.
promptKitBuiltinProfileCatalogID = "promptkit:v0.3.0:builtin-profiles"
@@ -45,6 +47,16 @@ func promptKitProfileFingerprint(profileDir, profileFile string) (CheckpointFing
}, nil
}
func promptKitLocalBackendFingerprint(endpoint string) CheckpointFingerprint {
hasher := sha256.New()
writeFingerprintPart(hasher, []byte(promptKitLocalBackendMarker))
writeFingerprintPart(hasher, []byte(strings.TrimSpace(endpoint)))
return CheckpointFingerprint{
Name: promptKitLocalBackendFingerprintName,
Value: "sha256:" + hex.EncodeToString(hasher.Sum(nil)),
}
}
func promptKitProfileFileDigests(root string) ([][]byte, error) {
var digests [][]byte
err := filepath.WalkDir(root, func(name string, entry fs.DirEntry, walkErr error) error {

View File

@@ -81,7 +81,10 @@ func TestScheduledClientPreservesCheckpointFingerprints(t *testing.T) {
t.Fatal(err)
}
inner := &fingerprintedStructuredClient{
fingerprints: []CheckpointFingerprint{{Name: "profile_source", Value: "sha256:one"}},
fingerprints: []CheckpointFingerprint{
{Name: "profile_source", Value: "sha256:one"},
{Name: "backend_target", Value: "sha256:two"},
},
}
client := NewScheduledClient(inner, scheduler)
provider, ok := client.(CheckpointFingerprintProvider)
@@ -92,7 +95,9 @@ func TestScheduledClientPreservesCheckpointFingerprints(t *testing.T) {
if err != nil {
t.Fatal(err)
}
if len(got) != 1 || got[0] != inner.fingerprints[0] {
if len(got) != len(inner.fingerprints) ||
got[0] != inner.fingerprints[0] ||
got[1] != inner.fingerprints[1] {
t.Fatalf("checkpoint fingerprints = %#v, want %#v", got, inner.fingerprints)
}
}