fix(session): derive LCP session ID from rolling prefix keys and harden matcher

- root derived session ID in rolling prefix key to disambiguate conversations with different system prompts
- sanitize matcher config to guarantee MaxPrefixes >= MaxTurns
- tie-break tool parts sorting with digest for determinism
This commit is contained in:
sususu
2026-09-01 18:00:13 +08:00
parent 5679bbf330
commit dc0f7a594b
2 changed files with 92 additions and 5 deletions

View File

@@ -586,7 +586,10 @@ func normalizeCanonicalTurn(turn CanonicalTurn) CanonicalTurn {
}
if len(toolParts) > 1 {
sort.SliceStable(toolParts, func(i, j int) bool {
return toolParts[i].Value < toolParts[j].Value
if toolParts[i].Value != toolParts[j].Value {
return toolParts[i].Value < toolParts[j].Value
}
return toolParts[i].Digest < toolParts[j].Digest
})
for index, partIndex := range toolIndexes {
parts[partIndex] = toolParts[index]
@@ -693,6 +696,9 @@ func NewMerklePrefixMatcherWithConfig(cfg MerklePrefixMatcherConfig) *MerklePref
if cfg.MaxPrefixes <= 0 {
cfg.MaxPrefixes = defaultMatcherMaxPrefixes
}
if cfg.MaxPrefixes < cfg.MaxTurns {
cfg.MaxPrefixes = cfg.MaxTurns
}
return &MerklePrefixMatcher{
ttl: cfg.TTL,
maxTurns: cfg.MaxTurns,
@@ -955,10 +961,15 @@ func (m *MerklePrefixMatcher) bindLocked(namespace string, fingerprints []string
if match, ok := m.matchLocked(namespace, fingerprints, minPrefixLength, now); ok {
sessionID = match.SessionID
}
prefixKeys := rollingPrefixKeys(fingerprints)
if sessionID == "" {
firstKey := fingerprints[0]
if minPrefixLength > 0 && minPrefixLength <= len(fingerprints) {
firstKey = fingerprints[minPrefixLength-1]
firstKey := ""
if len(prefixKeys) > 0 {
targetIndex := 0
if minPrefixLength > 0 && minPrefixLength <= len(prefixKeys) {
targetIndex = minPrefixLength - 1
}
firstKey = prefixKeys[targetIndex]
}
sessionID = newLCPSessionID(namespace, firstKey)
}
@@ -969,6 +980,7 @@ func (m *MerklePrefixMatcher) bindLocked(namespace string, fingerprints []string
sessionID: sessionID,
minPrefixLength: minPrefixLength,
fingerprints: append([]string(nil), fingerprints...),
prefixKeys: prefixKeys,
expiresAt: now.Add(m.ttl),
})
return sessionID
@@ -976,7 +988,9 @@ func (m *MerklePrefixMatcher) bindLocked(namespace string, fingerprints []string
func (m *MerklePrefixMatcher) addGroupLocked(group *lcpGroup) {
ns := m.namespaceLocked(group.namespace)
group.prefixKeys = rollingPrefixKeys(group.fingerprints)
if len(group.prefixKeys) == 0 {
group.prefixKeys = rollingPrefixKeys(group.fingerprints)
}
ns.groups[group.key] = group
for _, prefix := range group.prefixKeys {
bucket := ns.prefixes[prefix]

View File

@@ -416,6 +416,79 @@ func TestMerklePrefixMatcherConcurrentAccess(t *testing.T) {
}
}
func TestMerklePrefixMatcherRollingPrefixSessionIDDifferentSystemPrompts(t *testing.T) {
t.Parallel()
matcher := NewMerklePrefixMatcher(time.Minute)
defer matcher.Clear()
namespace := "lcp:v1:test:model:caller"
// Two conversations with different system instructions but identical first user prompt
conv1 := []CanonicalTurn{
{Role: "system", Parts: []CanonicalPart{{Kind: "text", Value: "system instruction AAA"}}},
{Role: "user", Parts: []CanonicalPart{{Kind: "text", Value: "common question"}}},
}
conv2 := []CanonicalTurn{
{Role: "system", Parts: []CanonicalPart{{Kind: "text", Value: "system instruction BBB"}}},
{Role: "user", Parts: []CanonicalPart{{Kind: "text", Value: "common question"}}},
}
sessionID1 := matcher.Bind(namespace, conv1, "auth-1")
sessionID2 := matcher.Bind(namespace, conv2, "auth-2")
if sessionID1 == "" || sessionID2 == "" {
t.Fatalf("expected non-empty session IDs, got sessionID1=%q, sessionID2=%q", sessionID1, sessionID2)
}
// Rolling prefix ensures that because turn 0 differs, the derived session ID for turn 2 differs
if sessionID1 == sessionID2 {
t.Fatalf("expected distinct session IDs for conversations with different system prompts, got same: %q", sessionID1)
}
}
func TestMerklePrefixMatcherConfigBoundsSanitization(t *testing.T) {
t.Parallel()
matcher := NewMerklePrefixMatcherWithConfig(MerklePrefixMatcherConfig{
MaxTurns: 10,
MaxPrefixes: 2, // Less than MaxTurns
})
defer matcher.Clear()
if matcher.maxPrefixes < matcher.maxTurns {
t.Fatalf("matcher.maxPrefixes = %d, want >= maxTurns (%d)", matcher.maxPrefixes, matcher.maxTurns)
}
}
func TestNormalizeCanonicalTurnToolPartsDigestTieBreak(t *testing.T) {
t.Parallel()
turn1 := CanonicalTurn{
Role: "assistant",
Parts: []CanonicalPart{
{Kind: "tool:call", Value: "same_value", Digest: "digest_b"},
{Kind: "tool:call", Value: "same_value", Digest: "digest_a"},
},
}
turn2 := CanonicalTurn{
Role: "assistant",
Parts: []CanonicalPart{
{Kind: "tool:call", Value: "same_value", Digest: "digest_a"},
{Kind: "tool:call", Value: "same_value", Digest: "digest_b"},
},
}
norm1 := normalizeCanonicalTurn(turn1)
norm2 := normalizeCanonicalTurn(turn2)
if norm1.Parts[0].Digest != "digest_a" || norm2.Parts[0].Digest != "digest_a" {
t.Fatalf("tool parts were not sorted deterministically by Digest: %#v vs %#v", norm1, norm2)
}
if FastTurnFingerprint(norm1) != FastTurnFingerprint(norm2) {
t.Fatal("fingerprints differed for tool parts in different input order")
}
}
func BenchmarkFastTurnFingerprint(b *testing.B) {
turn := CanonicalTurn{
Role: "user",