From 984836ba379d92065e61a12149525010dd5c4778 Mon Sep 17 00:00:00 2001 From: sususu Date: Mon, 10 Aug 2026 17:27:16 +0800 Subject: [PATCH] refactor(antigravity): route replay merge through the request index The indexing change left two parallel implementations of the same replay matching logic: the filter path used the request index while the merge and insert paths still rescanned the payload. Filter and merge must agree on which part a ledger item targets, so the duplication was a latent source of signatures being replayed onto the wrong part. Move the write path onto the index. Every lookup in the merge path already ran before the first mutation, so one index describes the whole call; insert rebuilds it after each mutation to keep sequential semantics. Collapse the thought signature locator into thoughtSignaturePartIndex, shared by the eligibility check and the write path. Delete the superseded payload-scanning implementations along with functions that had no production callers left: antigravityNeedsSignatureReplayForExistingFunctionCall, antigravityRequestHasMatchingFunctionResponse, antigravityPayloadHasFunctionCallID, filterAntigravityReasoningReplayItemsForRequest, insertAntigravityReasoningReplayItems, mergeAntigravityFunctionCallPartReplay and antigravitySetReplayItemContextHash. Freeze the pre-index implementations into a dedicated oracle test file. The differential tests previously called the production functions they were meant to check, so consolidating the logic would have silently turned them into self-comparisons. Malformed non-array parts now fail closed consistently. gjson's Result.Array() yields a one-element slice for a non-array non-null value, so the old fallback scans could match a functionCall inside a parts object while the primary ID lookup could not. The end state was already identical because such a payload cannot be written to; the behavior is pinned by a test. Also cover the positional fallback for pre-targetHash cache entries, which had no test at all, and add a write-path benchmark. --- .../executor/antigravity_reasoning_replay.go | 449 +++------------- ...antigravity_reasoning_replay_index_test.go | 239 ++++++++- ...ity_reasoning_replay_legacy_oracle_test.go | 485 ++++++++++++++++++ .../antigravity_reasoning_replay_test.go | 15 +- 4 files changed, 787 insertions(+), 401 deletions(-) create mode 100644 internal/runtime/executor/antigravity_reasoning_replay_legacy_oracle_test.go diff --git a/internal/runtime/executor/antigravity_reasoning_replay.go b/internal/runtime/executor/antigravity_reasoning_replay.go index f6fb7f50..c2240f2b 100644 --- a/internal/runtime/executor/antigravity_reasoning_replay.go +++ b/internal/runtime/executor/antigravity_reasoning_replay.go @@ -321,7 +321,7 @@ func applyAntigravityReasoningReplayItems(payload []byte, items [][]byte, toolSc if len(eligible) != 1 { continue } - next, applied := insertAntigravityReasoningReplayItemsWithSchemas(updated, eligible, toolSchemas) + next, applied := insertAntigravityReasoningReplayItemsWithSchemas(index, updated, eligible, toolSchemas) if !applied { continue } @@ -334,10 +334,6 @@ func applyAntigravityReasoningReplayItems(payload []byte, items [][]byte, toolSc return updated, changed } -func filterAntigravityReasoningReplayItemsForRequest(payload []byte, items [][]byte) [][]byte { - return filterAntigravityReasoningReplayItemsForRequestWithSchemas(payload, items, nil) -} - func filterAntigravityReasoningReplayItemsForRequestWithSchemas(payload []byte, items [][]byte, toolSchemas map[string]any) [][]byte { index := newAntigravityReplayRequestIndex(payload) return filterAntigravityReasoningReplayItemsForRequestWithIndex(index, items, toolSchemas) @@ -463,67 +459,6 @@ func antigravityAnyKeyExists(existing map[string]bool, keys []string) bool { return false } -func antigravityNeedsSignatureReplayForExistingFunctionCall(payload []byte, itemResult gjson.Result) bool { - if strings.TrimSpace(itemResult.Get("thoughtSignature").String()) == "" { - return false - } - ci, pi, ok := antigravityFunctionCallPartLocationForReplay(payload, itemResult) - if !ok { - return false - } - pathSig := fmt.Sprintf("request.contents.%d.parts.%d.thoughtSignature", ci, pi) - return !antigravityHasNativeThoughtSignature(gjson.GetBytes(payload, pathSig).String()) -} - -func antigravityRequestHasMatchingFunctionResponse(payload []byte, itemResult gjson.Result) bool { - callID := strings.TrimSpace(itemResult.Get("call_id").String()) - if callID == "" { - return true - } - _, _, ok := antigravityFunctionResponseContentIndexForReplay(payload, itemResult) - return ok -} - -func antigravityFunctionResponseContentIndexForReplay(payload []byte, itemResult gjson.Result) (int, string, bool) { - callID := strings.TrimSpace(itemResult.Get("call_id").String()) - name := strings.TrimSpace(itemResult.Get("name").String()) - args := itemResult.Get("args") - candidateIDs := []string{callID} - if stableID := util.GeminiClaudeToolUseID(callID, name, args.Raw); stableID != "" && stableID != callID { - candidateIDs = append(candidateIDs, stableID) - } - for _, candidateID := range candidateIDs { - if contentIndex, ok := antigravityFunctionResponseContentIndex(payload, candidateID); ok { - return contentIndex, candidateID, true - } - } - return -1, "", false -} - -func antigravityFunctionResponseContentIndex(payload []byte, callID string) (int, bool) { - callID = strings.TrimSpace(callID) - if callID == "" { - return -1, false - } - contents := util.GetGJSONBytesNoCopy(payload, "request.contents") - if !contents.IsArray() { - return -1, false - } - for i, content := range contents.Array() { - parts := content.Get("parts") - if !parts.IsArray() { - continue - } - for _, part := range parts.Array() { - fr := part.Get("functionResponse") - if fr.Exists() && strings.TrimSpace(fr.Get("id").String()) == callID { - return i, true - } - } - } - return -1, false -} - func restoreAntigravityFunctionResponseReplayIdentity(payload []byte, currentID, nativeID, nativeName string) []byte { currentID = strings.TrimSpace(currentID) nativeID = strings.TrimSpace(nativeID) @@ -549,154 +484,6 @@ func restoreAntigravityFunctionResponseReplayIdentity(payload []byte, currentID, return out } -func antigravityPayloadHasFunctionCallID(payload []byte, callID string) bool { - _, _, ok := antigravityFunctionCallPartLocation(payload, callID) - return ok -} - -func antigravityFunctionCallPartLocation(payload []byte, callID string) (contentIndex int, partIndex int, ok bool) { - callID = strings.TrimSpace(callID) - if callID == "" { - return -1, -1, false - } - contents := util.GetGJSONBytesNoCopy(payload, "request.contents") - if !contents.IsArray() { - return -1, -1, false - } - for ci, content := range contents.Array() { - parts := content.Get("parts") - if !parts.IsArray() { - continue - } - for pi, part := range parts.Array() { - fc := part.Get("functionCall") - if fc.Exists() && strings.TrimSpace(fc.Get("id").String()) == callID { - return ci, pi, true - } - } - } - return -1, -1, false -} - -func antigravityFunctionCallPartLocationForReplay(payload []byte, itemResult gjson.Result) (contentIndex int, partIndex int, ok bool) { - return antigravityFunctionCallPartLocationForReplayWithSchemas(payload, itemResult, nil) -} - -func antigravityFunctionCallPartLocationForReplayWithSchemas(payload []byte, itemResult gjson.Result, toolSchemas map[string]any) (contentIndex int, partIndex int, ok bool) { - name := strings.TrimSpace(itemResult.Get("name").String()) - args := itemResult.Get("args") - if name == "" || !args.Exists() { - return -1, -1, false - } - callID := strings.TrimSpace(itemResult.Get("call_id").String()) - if callID == "" { - callID = strings.TrimSpace(itemResult.Get("id").String()) - } - candidateIDs := []string{callID} - if stableID := util.GeminiClaudeToolUseID(callID, name, args.Raw); stableID != "" && stableID != callID { - candidateIDs = append(candidateIDs, stableID) - } - for _, candidateID := range candidateIDs { - if candidateID == "" { - continue - } - ci, pi, found := antigravityFunctionCallPartLocation(payload, candidateID) - if !found { - continue - } - if antigravityReplayItemContextMatches(payload, itemResult, ci) { - fc := gjson.GetBytes(payload, fmt.Sprintf("request.contents.%d.parts.%d.functionCall", ci, pi)) - if antigravityFunctionCallMatchesReplayItem(fc, itemResult, toolSchemas) { - return ci, pi, true - } - log.Debugf("antigravity replay: located call %q at contents[%d].parts[%d] but name/args did not match ledger item (opaque_id=%t)", - name, ci, pi, util.IsGeminiClaudeToolUseID(candidateID)) - return -1, -1, false - } - // The candidate ID matched exactly, so callID+name+args are already proven - // identical. Only the surrounding context drifted, which invalidates the - // cached signature but not the tool identity. - log.Debugf("antigravity replay: exact tool ID match for %q at contents[%d].parts[%d] rejected by context hash (opaque_id=%t)", - name, ci, pi, util.IsGeminiClaudeToolUseID(candidateID)) - return -1, -1, false - } - contents := util.GetGJSONBytesNoCopy(payload, "request.contents") - if !contents.IsArray() { - return -1, -1, false - } - contentArr := contents.Array() - cachedCI := int(itemResult.Get("contentIndex").Int()) - if targetOccurrence := itemResult.Get("targetOccurrence"); targetOccurrence.Exists() { - if cachedCI < 0 || cachedCI >= len(contentArr) || !antigravityReplayItemContextMatches(payload, itemResult, cachedCI) { - return -1, -1, false - } - wantedOccurrence := int(targetOccurrence.Int()) - occurrence := 0 - for pi, part := range contentArr[cachedCI].Get("parts").Array() { - fc := part.Get("functionCall") - if !fc.Exists() || (util.IsGeminiClaudeToolUseID(fc.Get("id").String()) && fc.Get("id").String() != util.GeminiClaudeToolUseID(callID, name, args.Raw)) || !antigravityFunctionCallMatchesReplayItem(fc, itemResult, toolSchemas) { - continue - } - if occurrence == wantedOccurrence { - return cachedCI, pi, true - } - occurrence++ - } - return -1, -1, false - } - - matches := make([][2]int, 0, 1) - for ci, content := range contentArr { - if !antigravityReplayItemContextMatches(payload, itemResult, ci) { - continue - } - for pi, part := range content.Get("parts").Array() { - fc := part.Get("functionCall") - if !fc.Exists() || (util.IsGeminiClaudeToolUseID(fc.Get("id").String()) && fc.Get("id").String() != util.GeminiClaudeToolUseID(callID, name, args.Raw)) { - continue - } - if antigravityFunctionCallMatchesReplayItem(fc, itemResult, toolSchemas) { - matches = append(matches, [2]int{ci, pi}) - } - } - } - if len(matches) == 1 { - return matches[0][0], matches[0][1], true - } - return -1, -1, false -} - -// antigravityFunctionCallProvenanceLocation locates the function call whose -// Claude-facing opaque ID was derived from this exact ledger item. -// -// The opaque ID is sha256(call_id, name, args), so an exact match already proves -// that the call ID, tool name and arguments are identical to the provider-native -// call. The surrounding context hash adds nothing to that proof; it only decides -// whether the cached thoughtSignature is still valid. Callers therefore use this -// to recover tool identity after the context has drifted, without replaying any -// signature. -func antigravityFunctionCallProvenanceLocation(payload []byte, itemResult gjson.Result, toolSchemas map[string]any) (contentIndex int, partIndex int, ok bool) { - name := strings.TrimSpace(itemResult.Get("name").String()) - args := itemResult.Get("args") - callID := strings.TrimSpace(itemResult.Get("call_id").String()) - if name == "" || !args.Exists() || callID == "" { - return -1, -1, false - } - stableID := util.GeminiClaudeToolUseID(callID, name, args.Raw) - if stableID == "" || stableID == callID { - return -1, -1, false - } - ci, pi, found := antigravityFunctionCallPartLocation(payload, stableID) - if !found { - return -1, -1, false - } - fc := gjson.GetBytes(payload, fmt.Sprintf("request.contents.%d.parts.%d.functionCall", ci, pi)) - if !antigravityFunctionCallMatchesReplayItem(fc, itemResult, toolSchemas) { - return -1, -1, false - } - return ci, pi, true -} - func (i *antigravityReplayRequestIndex) functionResponseContentIndexForReplay(itemResult gjson.Result) (int, string, bool) { callID := strings.TrimSpace(itemResult.Get("call_id").String()) name := strings.TrimSpace(itemResult.Get("name").String()) @@ -832,19 +619,29 @@ func (i *antigravityReplayRequestIndex) functionCallProvenanceLocation( return location, true } -func (i *antigravityReplayRequestIndex) hasThoughtSignatureAt(itemResult gjson.Result) bool { - contentIndex := int(itemResult.Get("contentIndex").Int()) +// thoughtSignaturePartIndex resolves the part a thought_signature item belongs +// to. It is the single locator shared by the eligibility check and the write +// path, so the two can never disagree about the target part. +// +// A target hash pins the signature to a part whose own bytes are unchanged, +// which is all Gemini validates: the signature's own integrity, never its +// binding to the surrounding history. Drift elsewhere in the conversation +// therefore costs this signature nothing, so it is deliberately not gated on +// the context fingerprint. The positional fallback below has no such proof and +// stays gated. +func (i *antigravityReplayRequestIndex) thoughtSignaturePartIndex(itemResult gjson.Result) (contentIndex int, partIndex int, ok bool) { + contentIndex = int(itemResult.Get("contentIndex").Int()) if i == nil || contentIndex < 0 || contentIndex >= len(i.contents) { - return false + return -1, -1, false } content := i.contents[contentIndex] if !strings.EqualFold(strings.TrimSpace(content.content.Get("role").String()), "model") { - return false + return -1, -1, false } parts := content.parts targetKind := strings.TrimSpace(itemResult.Get("targetKind").String()) targetHash := strings.TrimSpace(itemResult.Get("targetHash").String()) - partIndex := -1 + partIndex = -1 if targetHash != "" { if targetOccurrence := itemResult.Get("targetOccurrence"); targetOccurrence.Exists() { wantedOccurrence := int(targetOccurrence.Int()) @@ -879,8 +676,10 @@ func (i *antigravityReplayRequestIndex) hasThoughtSignatureAt(itemResult gjson.R } } } else { + // No target hash: nothing proves which part this signature belongs to, so + // only a matching context fingerprint makes the positional guess safe. if !i.contextMatches(itemResult, contentIndex) { - return false + return -1, -1, false } candidateIndex := int(itemResult.Get("partIndex").Int()) if candidateIndex >= 0 && candidateIndex < len(parts) && parts[candidateIndex].Type != gjson.Null { @@ -888,6 +687,9 @@ func (i *antigravityReplayRequestIndex) hasThoughtSignatureAt(itemResult gjson.R partIndex = candidateIndex } } + // Legacy cache entries may point at a streamed signature-only part after + // multiple text chunks. Attach them to the last semantic part in the same + // model content, never to a different turn. if partIndex < 0 { for candidateIndex := len(parts) - 1; candidateIndex >= 0; candidateIndex-- { if kind, _ := antigravityReplayPartFingerprint(parts[candidateIndex]); kind != "" { @@ -898,9 +700,26 @@ func (i *antigravityReplayRequestIndex) hasThoughtSignatureAt(itemResult gjson.R } } if partIndex < 0 { + return -1, -1, false + } + return contentIndex, partIndex, true +} + +func (i *antigravityReplayRequestIndex) hasThoughtSignatureAt(itemResult gjson.Result) bool { + contentIndex, partIndex, ok := i.thoughtSignaturePartIndex(itemResult) + if !ok { return false } - return antigravityHasNativeThoughtSignature(parts[partIndex].Get("thoughtSignature").String()) + part := i.contents[contentIndex].parts[partIndex] + return antigravityHasNativeThoughtSignature(part.Get("thoughtSignature").String()) +} + +func (i *antigravityReplayRequestIndex) thoughtSignatureReplayPartPath(itemResult gjson.Result) (string, bool) { + contentIndex, partIndex, ok := i.thoughtSignaturePartIndex(itemResult) + if !ok { + return "", false + } + return fmt.Sprintf("request.contents.%d.parts.%d", contentIndex, partIndex), true } func insertAntigravityModelFunctionCallBeforeContent(payload []byte, beforeIndex int, name, callID, thoughtSig string, args gjson.Result) ([]byte, bool) { @@ -996,14 +815,6 @@ func antigravityRemoveThoughtSignatureFromOtherParts(payload []byte, contentInde return out } -func antigravityRequestHasThoughtSignatureAt(payload []byte, itemResult gjson.Result) bool { - partPath, ok := antigravityThoughtSignatureReplayPartPath(payload, itemResult) - if !ok { - return false - } - return antigravityHasNativeThoughtSignature(gjson.GetBytes(payload, partPath+".thoughtSignature").String()) -} - func antigravityHasNativeThoughtSignature(signature string) bool { signature = strings.TrimSpace(signature) return signature != "" && signature != "skip_thought_signature_validator" @@ -1241,49 +1052,6 @@ func (f *antigravityReplayContextFingerprints) at(beforeContentIndex int) string return f.sums[beforeContentIndex] } -func antigravityReplayContextFingerprint(payload []byte, beforeContentIndex int) string { - contents := util.GetGJSONBytesNoCopy(payload, "request.contents") - if !contents.IsArray() || beforeContentIndex < 0 { - return "" - } - contentArr := contents.Array() - if beforeContentIndex > len(contentArr) { - return "" - } - var context strings.Builder - for _, path := range []string{"request.systemInstruction", "request.tools", "request.toolConfig"} { - if value := gjson.GetBytes(payload, path); value.Exists() { - context.WriteString(path) - context.WriteByte('\x00') - context.Write(antigravityCanonicalReplayJSON([]byte(value.Raw))) - context.WriteByte('\x00') - } - } - for ci := 0; ci < beforeContentIndex; ci++ { - content := contentArr[ci] - context.WriteString(strings.ToLower(strings.TrimSpace(content.Get("role").String()))) - context.WriteByte('\x00') - parts := content.Get("parts") - if !parts.IsArray() { - continue - } - parts.ForEach(func(_, part gjson.Result) bool { - normalized := []byte(part.Raw) - for _, signaturePath := range []string{"thoughtSignature", "thought_signature", "extra_content.google.thought_signature"} { - normalized, _ = sjson.DeleteBytes(normalized, signaturePath) - } - context.Write(antigravityCanonicalReplayJSON(normalized)) - context.WriteByte('\x00') - return true - }) - } - if context.Len() == 0 { - return "" - } - sum := sha256.Sum256([]byte(context.String())) - return fmt.Sprintf("%x", sum[:]) -} - func antigravityReplayToolSchemasFromRequests(rawRequests ...[]byte) map[string]any { toolSchemas := make(map[string]any) for _, raw := range rawRequests { @@ -1524,15 +1292,6 @@ func antigravityCanonicalReplayJSON(raw []byte) []byte { return canonical } -func antigravityReplayItemContextMatches(payload []byte, itemResult gjson.Result, contentIndex int) bool { - expected := strings.TrimSpace(itemResult.Get("contextHash").String()) - return expected == "" || expected == antigravityReplayContextFingerprint(payload, contentIndex) -} - -func antigravitySetReplayItemContextHash(item []byte, payload []byte, contentIndex int) []byte { - return antigravitySetReplayItemContextHashValue(item, antigravityReplayContextFingerprint(payload, contentIndex)) -} - func antigravitySetReplayItemContextHashValue(item []byte, contextHash string) []byte { if contextHash != "" { item, _ = sjson.SetBytes(item, "contextHash", contextHash) @@ -1540,83 +1299,6 @@ func antigravitySetReplayItemContextHashValue(item []byte, contextHash string) [ return item } -func antigravityThoughtSignatureReplayPartPath(payload []byte, itemResult gjson.Result) (string, bool) { - ci := int(itemResult.Get("contentIndex").Int()) - contents := util.GetGJSONBytesNoCopy(payload, "request.contents") - if !contents.IsArray() { - return "", false - } - contentArr := contents.Array() - if ci < 0 || ci >= len(contentArr) || !strings.EqualFold(strings.TrimSpace(contentArr[ci].Get("role").String()), "model") { - return "", false - } - parts := contentArr[ci].Get("parts") - if !parts.IsArray() { - return "", false - } - partArr := parts.Array() - targetKind := strings.TrimSpace(itemResult.Get("targetKind").String()) - targetHash := strings.TrimSpace(itemResult.Get("targetHash").String()) - // A target hash pins the signature to a part whose own bytes are unchanged, - // which is all Gemini validates: the signature's own integrity, never its - // binding to the surrounding history. Drift elsewhere in the conversation - // therefore costs this signature nothing, so it is deliberately not gated on - // the context fingerprint. The fallback below has no such proof and stays - // gated. - if targetHash != "" { - if targetOccurrence := itemResult.Get("targetOccurrence"); targetOccurrence.Exists() { - wanted := int(targetOccurrence.Int()) - occurrence := 0 - for pi, part := range partArr { - kind, fingerprint := antigravityReplayPartFingerprint(part) - if fingerprint != targetHash || (targetKind != "" && kind != targetKind) { - continue - } - if occurrence == wanted { - return fmt.Sprintf("request.contents.%d.parts.%d", ci, pi), true - } - occurrence++ - } - return "", false - } - pi := int(itemResult.Get("partIndex").Int()) - if pi >= 0 && pi < len(partArr) { - kind, fingerprint := antigravityReplayPartFingerprint(partArr[pi]) - if fingerprint == targetHash && (targetKind == "" || kind == targetKind) { - return fmt.Sprintf("request.contents.%d.parts.%d", ci, pi), true - } - } - for pi, part := range partArr { - kind, fingerprint := antigravityReplayPartFingerprint(part) - if fingerprint == targetHash && (targetKind == "" || kind == targetKind) { - return fmt.Sprintf("request.contents.%d.parts.%d", ci, pi), true - } - } - return "", false - } - - // No target hash: nothing proves which part this signature belongs to, so - // only a matching context fingerprint makes the positional guess safe. - if !antigravityReplayItemContextMatches(payload, itemResult, ci) { - return "", false - } - pi := int(itemResult.Get("partIndex").Int()) - if pi >= 0 && pi < len(partArr) && partArr[pi].Type != gjson.Null { - if kind, _ := antigravityReplayPartFingerprint(partArr[pi]); kind != "" { - return fmt.Sprintf("request.contents.%d.parts.%d", ci, pi), true - } - } - // Legacy cache entries may point at a streamed signature-only part after - // multiple text chunks. Attach them to the last semantic part in the same - // model content, never to a different turn. - for candidate := len(partArr) - 1; candidate >= 0; candidate-- { - if kind, _ := antigravityReplayPartFingerprint(partArr[candidate]); kind != "" { - return fmt.Sprintf("request.contents.%d.parts.%d", ci, candidate), true - } - } - return "", false -} - func antigravityExistingReplayPartPath(payload []byte, contentIndex int, partIndex int) (string, bool) { if contentIndex < 0 || partIndex < 0 { return "", false @@ -1644,11 +1326,10 @@ func antigravityReplayPartWritePath(payload []byte, contentIndex int, partIndex return partsPath + ".0" } -func insertAntigravityReasoningReplayItems(payload []byte, items [][]byte) ([]byte, bool) { - return insertAntigravityReasoningReplayItemsWithSchemas(payload, items, nil) -} - -func insertAntigravityReasoningReplayItemsWithSchemas(payload []byte, items [][]byte, toolSchemas map[string]any) ([]byte, bool) { +// insertAntigravityReasoningReplayItemsWithSchemas applies items sequentially. +// index must describe payload on entry and is rebuilt after any mutation so each +// item observes exactly the payload the previous item produced. +func insertAntigravityReasoningReplayItemsWithSchemas(index *antigravityReplayRequestIndex, payload []byte, items [][]byte, toolSchemas map[string]any) ([]byte, bool) { out := payload changed := false for _, item := range items { @@ -1659,7 +1340,7 @@ func insertAntigravityReasoningReplayItemsWithSchemas(payload []byte, items [][] if sig == "" { continue } - partPath, exists := antigravityThoughtSignatureReplayPartPath(out, itemResult) + partPath, exists := index.thoughtSignatureReplayPartPath(itemResult) if !exists { continue } @@ -1671,15 +1352,20 @@ func insertAntigravityReasoningReplayItemsWithSchemas(payload []byte, items [][] out = antigravityRemoveThoughtSignatureFromOtherParts(out, ci, sig, partPath) updated, err := sjson.SetBytes(out, path, sig) if err != nil { + // antigravityRemoveThoughtSignatureFromOtherParts may already have + // rewritten out, so the index has to be refreshed regardless. + index = newAntigravityReplayRequestIndex(out) continue } out = updated changed = true + index = newAntigravityReplayRequestIndex(out) case "function_call_part": - updated, ok := mergeAntigravityFunctionCallPartReplayWithSchemas(out, itemResult, toolSchemas) + updated, ok := mergeAntigravityFunctionCallPartReplayWithSchemas(index, out, itemResult, toolSchemas) if ok { out = updated changed = true + index = newAntigravityReplayRequestIndex(out) } } } @@ -1798,11 +1484,11 @@ func restoreAntigravityNativeFunctionCallReplay(payload []byte, contentIndex, pa return out, !bytes.Equal(out, payload) } -func mergeAntigravityFunctionCallPartReplay(payload []byte, itemResult gjson.Result) ([]byte, bool) { - return mergeAntigravityFunctionCallPartReplayWithSchemas(payload, itemResult, nil) -} - -func mergeAntigravityFunctionCallPartReplayWithSchemas(payload []byte, itemResult gjson.Result, toolSchemas map[string]any) ([]byte, bool) { +// mergeAntigravityFunctionCallPartReplayWithSchemas locates the target call via +// index, which must describe exactly the payload passed alongside it. Every +// lookup happens before the first mutation, so one index is valid for the whole +// call. +func mergeAntigravityFunctionCallPartReplayWithSchemas(index *antigravityReplayRequestIndex, payload []byte, itemResult gjson.Result, toolSchemas map[string]any) ([]byte, bool) { name := strings.TrimSpace(itemResult.Get("name").String()) args := itemResult.Get("args") callID := strings.TrimSpace(itemResult.Get("call_id").String()) @@ -1810,34 +1496,39 @@ func mergeAntigravityFunctionCallPartReplayWithSchemas(payload []byte, itemResul if name == "" || !args.Exists() { return payload, false } - if ci, pi, exists := antigravityFunctionCallPartLocationForReplayWithSchemas(payload, itemResult, toolSchemas); exists { + if location, exists := index.functionCallPartLocationForReplayWithSchemas(itemResult, toolSchemas); exists { _, allowLegacyIDRestore := toolSchemas[name] - return restoreAntigravityNativeFunctionCallReplay(payload, ci, pi, itemResult, allowLegacyIDRestore, true) + return restoreAntigravityNativeFunctionCallReplay(payload, location.contentIndex, location.partIndex, itemResult, allowLegacyIDRestore, true) } // The context drifted, but an exact opaque ID match still proves this call's // identity. Gemini validates a thought signature's own integrity and nothing // about the history around it, so the drift costs the signature nothing: restore // the native call and its signature rather than making the model re-reason. - if ci, pi, exists := antigravityFunctionCallProvenanceLocation(payload, itemResult, toolSchemas); exists { - return restoreAntigravityNativeFunctionCallReplay(payload, ci, pi, itemResult, false, true) + if location, exists := index.functionCallProvenanceLocation(itemResult, toolSchemas); exists { + return restoreAntigravityNativeFunctionCallReplay(payload, location.contentIndex, location.partIndex, itemResult, false, true) } if callID != "" { stableID := util.GeminiClaudeToolUseID(callID, name, args.Raw) - if antigravityPayloadHasFunctionCallID(payload, callID) || (stableID != "" && antigravityPayloadHasFunctionCallID(payload, stableID)) { + _, hasNativeID := index.functionCallPartLocation(callID) + hasStableID := false + if stableID != "" { + _, hasStableID = index.functionCallPartLocation(stableID) + } + if hasNativeID || hasStableID { // The call is already in the history under its native or Claude-facing // ID, and neither lookup above accepted it, so the client changed it. // Never replay an opaque signature onto that changed call, and never // insert a second copy of it further down. return payload, false } - if frIndex, currentResponseID, ok := antigravityFunctionResponseContentIndexForReplay(payload, itemResult); ok { + if frIndex, currentResponseID, ok := index.functionResponseContentIndexForReplay(itemResult); ok { parallelModelIndex := frIndex - 1 - if parallelModelIndex >= 0 && strings.EqualFold(strings.TrimSpace(gjson.GetBytes(payload, fmt.Sprintf("request.contents.%d.role", parallelModelIndex)).String()), "model") && antigravityReplayItemContextMatches(payload, itemResult, parallelModelIndex) { + if parallelModelIndex >= 0 && strings.EqualFold(strings.TrimSpace(index.contents[parallelModelIndex].content.Get("role").String()), "model") && index.contextMatches(itemResult, parallelModelIndex) { if updated, appended := appendAntigravityFunctionCallToModelContent(payload, parallelModelIndex, name, callID, sig, args); appended { return restoreAntigravityFunctionResponseReplayIdentity(updated, currentResponseID, callID, name), true } } - if antigravityReplayItemContextMatches(payload, itemResult, frIndex) { + if index.contextMatches(itemResult, frIndex) { if updated, inserted := insertAntigravityModelFunctionCallBeforeContent(payload, frIndex, name, callID, sig, args); inserted { return restoreAntigravityFunctionResponseReplayIdentity(updated, currentResponseID, callID, name), true } @@ -1850,7 +1541,7 @@ func mergeAntigravityFunctionCallPartReplayWithSchemas(payload []byte, itemResul } ci := antigravityReasoningReplayResolveContentIndex(payload, int(itemResult.Get("contentIndex").Int())) - if ci < 0 || !antigravityReplayItemContextMatches(payload, itemResult, ci) { + if ci < 0 || !index.contextMatches(itemResult, ci) { return payload, false } pi := int(itemResult.Get("partIndex").Int()) diff --git a/internal/runtime/executor/antigravity_reasoning_replay_index_test.go b/internal/runtime/executor/antigravity_reasoning_replay_index_test.go index 0cc5aaa5..9230e29a 100644 --- a/internal/runtime/executor/antigravity_reasoning_replay_index_test.go +++ b/internal/runtime/executor/antigravity_reasoning_replay_index_test.go @@ -11,6 +11,13 @@ import ( "github.com/tidwall/sjson" ) +// antigravityReplayItemContextHashForTest stamps an item with the production +// context fingerprint for contentIndex, going through the request index exactly +// as production does. +func antigravityReplayItemContextHashForTest(item, payload []byte, contentIndex int) []byte { + return antigravitySetReplayItemContextHashValue(item, newAntigravityReplayRequestIndex(payload).contextFingerprint(contentIndex)) +} + func legacyAntigravityReasoningReplayItemsFromRequest(payload []byte) [][]byte { contents := gjson.GetBytes(payload, "request.contents") if !contents.IsArray() { @@ -40,7 +47,7 @@ func legacyAntigravityReasoningReplayItemsFromRequest(payload []byte) [][]byte { functionCallOccurrences[key] = occurrence + 1 } if item := buildAntigravityFunctionCallPartItem(contentIndex, partIndex, occurrence, functionCall, signature); len(item) > 0 { - items = append(items, antigravitySetReplayItemContextHash(item, payload, contentIndex)) + items = append(items, legacyAntigravitySetReplayItemContextHash(item, payload, contentIndex)) } continue } @@ -60,7 +67,7 @@ func legacyAntigravityReasoningReplayItemsFromRequest(payload []byte) [][]byte { } item := buildAntigravityThoughtSignatureItem(contentIndex, targetPartIndex, signature, kind, fingerprint) item, _ = sjson.SetBytes(item, "targetOccurrence", antigravityReplayPartOccurrence(partArray, targetPartIndex, kind, fingerprint)) - items = append(items, antigravitySetReplayItemContextHash(item, payload, contentIndex)) + items = append(items, legacyAntigravitySetReplayItemContextHash(item, payload, contentIndex)) } return true }) @@ -74,7 +81,7 @@ func legacyFilterAntigravityReasoningReplayItemsForRequestWithSchemas(payload [] switch strings.TrimSpace(itemResult.Get("type").String()) { case "function_call_part": signature := strings.TrimSpace(itemResult.Get("thoughtSignature").String()) - if contentIndex, partIndex, foundCall := antigravityFunctionCallPartLocationForReplayWithSchemas(payload, itemResult, toolSchemas); foundCall { + if contentIndex, partIndex, foundCall := legacyAntigravityFunctionCallPartLocationForReplayWithSchemas(payload, itemResult, toolSchemas); foundCall { part := gjson.GetBytes(payload, fmt.Sprintf("request.contents.%d.parts.%d", contentIndex, partIndex)) currentID := strings.TrimSpace(part.Get("functionCall.id").String()) nativeID := strings.TrimSpace(itemResult.Get("call_id").String()) @@ -87,27 +94,27 @@ func legacyFilterAntigravityReasoningReplayItemsForRequestWithSchemas(payload [] } break } - if _, _, foundProvenance := antigravityFunctionCallProvenanceLocation(payload, itemResult, toolSchemas); foundProvenance { + if _, _, foundProvenance := legacyAntigravityFunctionCallProvenanceLocation(payload, itemResult, toolSchemas); foundProvenance { break } callID := strings.TrimSpace(itemResult.Get("call_id").String()) if callID == "" { continue } - responseIndex, _, foundResponse := antigravityFunctionResponseContentIndexForReplay(payload, itemResult) + responseIndex, _, foundResponse := legacyAntigravityFunctionResponseContentIndexForReplay(payload, itemResult) if !foundResponse { continue } - contextMatches := antigravityReplayItemContextMatches(payload, itemResult, responseIndex) + contextMatches := legacyAntigravityReplayItemContextMatches(payload, itemResult, responseIndex) if !contextMatches && responseIndex > 0 { previousRole := gjson.GetBytes(payload, fmt.Sprintf("request.contents.%d.role", responseIndex-1)).String() - contextMatches = strings.EqualFold(strings.TrimSpace(previousRole), "model") && antigravityReplayItemContextMatches(payload, itemResult, responseIndex-1) + contextMatches = strings.EqualFold(strings.TrimSpace(previousRole), "model") && legacyAntigravityReplayItemContextMatches(payload, itemResult, responseIndex-1) } if !contextMatches { continue } case "thought_signature": - if antigravityRequestHasThoughtSignatureAt(payload, itemResult) { + if legacyAntigravityRequestHasThoughtSignatureAt(payload, itemResult) { continue } default: @@ -126,7 +133,7 @@ func legacyApplyAntigravityReasoningReplayItems(payload []byte, items [][]byte, if len(eligible) != 1 { continue } - next, applied := insertAntigravityReasoningReplayItemsWithSchemas(updated, eligible, toolSchemas) + next, applied := legacyInsertAntigravityReasoningReplayItemsWithSchemas(updated, eligible, toolSchemas) if !applied { continue } @@ -167,7 +174,7 @@ func TestAntigravityReplayContextFingerprintsMatchLegacy(t *testing.T) { t.Run(test.name, func(t *testing.T) { index := newAntigravityReplayRequestIndex(test.payload) for beforeContentIndex := -1; beforeContentIndex <= len(index.contents)+1; beforeContentIndex++ { - want := antigravityReplayContextFingerprint(test.payload, beforeContentIndex) + want := legacyAntigravityReplayContextFingerprint(test.payload, beforeContentIndex) if got := index.contextFingerprint(beforeContentIndex); got != want { t.Fatalf("contextFingerprint(%d) = %q, want %q", beforeContentIndex, got, want) } @@ -342,7 +349,7 @@ func TestAntigravityReasoningReplayAccumulatorUsesIndexedContextHash(t *testing. if accumulator == nil { t.Fatal("accumulator is nil") } - wantContextHash := antigravityReplayContextFingerprint(payload, 1) + wantContextHash := legacyAntigravityReplayContextFingerprint(payload, 1) if accumulator.responseContextHash != wantContextHash { t.Fatalf("response context hash = %q, want %q", accumulator.responseContextHash, wantContextHash) } @@ -447,3 +454,213 @@ func syntheticAntigravityReplayBenchmarkPayload(inlineBytes, turns int) []byte { payload.WriteString(`]}}`) return []byte(payload.String()) } + +// syntheticAntigravityReplayMixedPayload builds a history whose model turns mix +// thought parts, text parts and function calls, so the extracted ledger contains +// both thought_signature and function_call_part items. +func syntheticAntigravityReplayMixedPayload(turns int) []byte { + var payload strings.Builder + payload.WriteString(`{"request":{"systemInstruction":{"parts":[{"text":"sys"}]},"contents":[`) + payload.WriteString(`{"role":"user","parts":[{"text":"start"}]}`) + for turn := range turns { + fmt.Fprintf(&payload, + `,{"role":"model","parts":[`+ + `{"thought":true,"text":"reason-%d","thoughtSignature":"tsig-%d"},`+ + `{"text":"say-%d","thoughtSignature":"xsig-%d"},`+ + `{"functionCall":{"id":"call-%d","name":"lookup","args":{"turn":%d}},"thoughtSignature":"csig-%d"}`+ + `]}`, + turn, turn, turn, turn, turn, turn, turn) + fmt.Fprintf(&payload, + `,{"role":"user","parts":[{"functionResponse":{"id":"call-%d","name":"lookup","response":{"result":"ok-%d"}}}]}`, + turn, turn) + } + payload.WriteString(`]}}`) + return []byte(payload.String()) +} + +func TestAntigravityReplayMergeRandomizedDifferential(t *testing.T) { + const randomSeed = 20260811 + const turns = 6 + randomSource := rand.New(rand.NewSource(randomSeed)) + basePayload := syntheticAntigravityReplayMixedPayload(turns) + + items := legacyAntigravityReasoningReplayItemsFromRequest(basePayload) + thoughtItems, callItems := 0, 0 + for _, item := range items { + switch gjson.GetBytes(item, "type").String() { + case "thought_signature": + thoughtItems++ + case "function_call_part": + callItems++ + } + } + if thoughtItems == 0 || callItems == 0 { + t.Fatalf("ledger must mix item kinds: thought=%d call=%d", thoughtItems, callItems) + } + + applied := 0 + for caseIndex := range 300 { + payload := bytes.Clone(basePayload) + for range 1 + randomSource.Intn(5) { + turn := randomSource.Intn(turns) + modelIndex := 1 + turn*2 + responseIndex := modelIndex + 1 + part := randomSource.Intn(3) + var errSet error + switch randomSource.Intn(9) { + case 0: + payload, errSet = sjson.DeleteBytes(payload, fmt.Sprintf("request.contents.%d.parts.%d.thoughtSignature", modelIndex, part)) + case 1: + payload, errSet = sjson.SetBytes(payload, fmt.Sprintf("request.contents.%d.parts.2.functionCall.args.turn", modelIndex), turn+50) + case 2: + payload, errSet = sjson.DeleteBytes(payload, fmt.Sprintf("request.contents.%d.parts.2.functionCall.id", modelIndex)) + case 3: + payload, errSet = sjson.SetBytes(payload, fmt.Sprintf("request.contents.%d.parts.1.text", modelIndex), "drifted") + case 4: + payload, errSet = sjson.SetBytes(payload, fmt.Sprintf("request.contents.%d.role", responseIndex), "model") + case 5: + payload, errSet = sjson.DeleteBytes(payload, fmt.Sprintf("request.contents.%d.parts.%d", modelIndex, part)) + case 6: + payload, errSet = sjson.SetBytes(payload, fmt.Sprintf("request.contents.%d.parts.0.thought", modelIndex), false) + case 7: + payload, errSet = sjson.SetBytes(payload, fmt.Sprintf("request.contents.%d.parts.2.functionCall.id", modelIndex), "call-0") + case 8: + payload, errSet = sjson.SetBytes(payload, "request.toolConfig.functionCallingConfig.mode", "ANY") + } + if errSet != nil { + t.Fatalf("case %d mutation failed: %v", caseIndex, errSet) + } + } + + wantPayload, wantChanged := legacyApplyAntigravityReasoningReplayItems(payload, items, nil) + gotPayload, gotChanged := applyAntigravityReasoningReplayItems(payload, items, nil) + if gotChanged != wantChanged { + t.Fatalf("seed=%d case=%d changed=%t want=%t", randomSeed, caseIndex, gotChanged, wantChanged) + } + if !bytes.Equal(gotPayload, wantPayload) { + t.Fatalf("seed=%d case=%d payload differs\n got: %s\nwant: %s", randomSeed, caseIndex, gotPayload, wantPayload) + } + if wantChanged { + applied++ + } + } + if applied == 0 { + t.Fatal("no case applied a replay item; the differential proved nothing") + } + t.Logf("cases=300 casesThatApplied=%d ledgerItems=%d (thought=%d call=%d)", applied, len(items), thoughtItems, callItems) +} + +// TestAntigravityReplayNonArrayPartsFailsClosed pins an INTENTIONAL behavior +// change made when the merge path moved onto the request index. +// +// gjson's Result.Array() returns a one-element slice for a value that exists but +// is neither null nor an array, so the pre-index fallback scans could "locate" a +// functionCall inside a parts OBJECT. The index only walks parts when IsArray(), +// which is what the primary ID lookup always did, so malformed parts now fail +// closed consistently instead of depending on which branch ran. +func TestAntigravityReplayNonArrayPartsFailsClosed(t *testing.T) { + payload := []byte(`{"request":{"contents":[` + + `{"role":"user","parts":[{"text":"hi"}]},` + + `{"role":"model","parts":{"functionCall":{"name":"lookup","args":{"value":1}}}}` + + `]}}`) + items := [][]byte{ + []byte(`{"type":"function_call_part","contentIndex":1,"partIndex":0,"name":"lookup","call_id":"call-1","args":{"value":1},"thoughtSignature":"sig-x"}`), + } + + if kept := filterAntigravityReasoningReplayItemsForRequestWithSchemas(payload, items, nil); len(kept) != 0 { + t.Fatalf("malformed parts must not yield an eligible item, kept=%d", len(kept)) + } + got, changed := applyAntigravityReasoningReplayItems(payload, items, nil) + if changed || !bytes.Equal(got, payload) { + t.Fatalf("malformed parts must not be mutated: changed=%t body=%s", changed, got) + } + // The legacy oracle accepted the item at the filter layer but could not write + // it either, so the observable end state was already identical. + if _, legacyChanged := legacyApplyAntigravityReasoningReplayItems(payload, items, nil); legacyChanged { + t.Fatal("legacy oracle unexpectedly mutated malformed parts") + } +} + +// TestAntigravityReplayLegacyItemWithoutTargetHash covers the positional +// fallback used by pre-targetHash cache entries. Such an item carries no proof +// of which part owns the signature, so it must attach to the LAST semantic part +// of the model content (the streamed chunks collapse into one part on replay), +// and only when the context fingerprint still matches. +func TestAntigravityReplayLegacyItemWithoutTargetHash(t *testing.T) { + payload := []byte(`{"request":{"contents":[` + + `{"role":"user","parts":[{"text":"ask"}]},` + + `{"role":"model","parts":[{"text":"chunk-a"},{"text":"chunk-b"},{"text":"chunk-c"}]}` + + `]}}`) + const signature = "legacy-positional-signature-12345" + // partIndex 7 is out of range on purpose: legacy entries pointed at a + // streamed signature-only part that no longer exists. + item := buildAntigravityThoughtSignatureItem(1, 7, signature, "", "") + item = antigravityReplayItemContextHashForTest(item, payload, 1) + if gjson.GetBytes(item, "targetHash").Exists() { + t.Fatal("this test must exercise the no-targetHash path") + } + + got, changed := applyAntigravityReasoningReplayItems(payload, [][]byte{item}, nil) + if !changed { + t.Fatalf("legacy positional item was not applied: %s", got) + } + if sig := gjson.GetBytes(got, "request.contents.1.parts.2.thoughtSignature").String(); sig != signature { + t.Fatalf("signature must attach to the LAST semantic part, got parts.2=%q body=%s", sig, got) + } + for _, path := range []string{"request.contents.1.parts.0.thoughtSignature", "request.contents.1.parts.1.thoughtSignature"} { + if gjson.GetBytes(got, path).Exists() { + t.Fatalf("signature leaked to %s: %s", path, got) + } + } + + want, wantChanged := legacyApplyAntigravityReasoningReplayItems(payload, [][]byte{item}, nil) + if wantChanged != changed || !bytes.Equal(want, got) { + t.Fatalf("legacy oracle disagrees\n got: %s\nwant: %s", got, want) + } + + // Context drift must reject the positional guess entirely. + drifted, errSet := sjson.SetBytes(payload, "request.contents.0.parts.0.text", "different question") + if errSet != nil { + t.Fatal(errSet) + } + driftedOut, driftedChanged := applyAntigravityReasoningReplayItems(drifted, [][]byte{item}, nil) + if driftedChanged || !bytes.Equal(driftedOut, drifted) { + t.Fatalf("context drift must reject a positional legacy item: changed=%t body=%s", driftedChanged, driftedOut) + } +} + +// BenchmarkApplyAntigravityReasoningReplayItems measures the WRITE path, where +// every ledger item actually mutates the payload. This is the worst case for the +// request index because it is rebuilt after each mutation. +func BenchmarkApplyAntigravityReasoningReplayItems(b *testing.B) { + const turns = 32 + base := syntheticAntigravityReplayBenchmarkPayload(1<<20, turns) + items := antigravityReasoningReplayItemsFromRequest(base) + if len(items) != turns { + b.Fatalf("items = %d, want %d", len(items), turns) + } + payload := base + for turn := range turns { + var err error + payload, err = sjson.DeleteBytes(payload, fmt.Sprintf("request.contents.%d.parts.0.thoughtSignature", 1+turn*2)) + if err != nil { + b.Fatal(err) + } + } + if _, changed := applyAntigravityReasoningReplayItems(payload, items, nil); !changed { + b.Fatal("benchmark payload applies nothing") + } + + b.Run("legacy", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + _, _ = legacyApplyAntigravityReasoningReplayItems(payload, items, nil) + } + }) + b.Run("indexed", func(b *testing.B) { + b.ReportAllocs() + for b.Loop() { + _, _ = applyAntigravityReasoningReplayItems(payload, items, nil) + } + }) +} diff --git a/internal/runtime/executor/antigravity_reasoning_replay_legacy_oracle_test.go b/internal/runtime/executor/antigravity_reasoning_replay_legacy_oracle_test.go new file mode 100644 index 00000000..e111fd95 --- /dev/null +++ b/internal/runtime/executor/antigravity_reasoning_replay_legacy_oracle_test.go @@ -0,0 +1,485 @@ +package executor + +// This file is a frozen, pre-index copy of the Antigravity reasoning replay +// location, context-fingerprint and merge logic. It exists so the differential +// tests compare the indexed implementation against an INDEPENDENT oracle rather +// than against itself. +// +// Do not refactor these functions, do not make them delegate to the production +// implementation, and do not "fix" them. If a production behavior change is +// intentional, assert the new behavior explicitly in a test instead of editing +// this oracle. + +import ( + "crypto/sha256" + "encoding/json" + "fmt" + "strings" + + "github.com/router-for-me/CLIProxyAPI/v7/internal/util" + log "github.com/sirupsen/logrus" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" +) + +func legacyAntigravityFunctionResponseContentIndexForReplay(payload []byte, itemResult gjson.Result) (int, string, bool) { + callID := strings.TrimSpace(itemResult.Get("call_id").String()) + name := strings.TrimSpace(itemResult.Get("name").String()) + args := itemResult.Get("args") + candidateIDs := []string{callID} + if stableID := util.GeminiClaudeToolUseID(callID, name, args.Raw); stableID != "" && stableID != callID { + candidateIDs = append(candidateIDs, stableID) + } + for _, candidateID := range candidateIDs { + if contentIndex, ok := legacyAntigravityFunctionResponseContentIndex(payload, candidateID); ok { + return contentIndex, candidateID, true + } + } + return -1, "", false +} + +func legacyAntigravityFunctionResponseContentIndex(payload []byte, callID string) (int, bool) { + callID = strings.TrimSpace(callID) + if callID == "" { + return -1, false + } + contents := util.GetGJSONBytesNoCopy(payload, "request.contents") + if !contents.IsArray() { + return -1, false + } + for i, content := range contents.Array() { + parts := content.Get("parts") + if !parts.IsArray() { + continue + } + for _, part := range parts.Array() { + fr := part.Get("functionResponse") + if fr.Exists() && strings.TrimSpace(fr.Get("id").String()) == callID { + return i, true + } + } + } + return -1, false +} + +func legacyAntigravityPayloadHasFunctionCallID(payload []byte, callID string) bool { + _, _, ok := legacyAntigravityFunctionCallPartLocation(payload, callID) + return ok +} + +func legacyAntigravityFunctionCallPartLocation(payload []byte, callID string) (contentIndex int, partIndex int, ok bool) { + callID = strings.TrimSpace(callID) + if callID == "" { + return -1, -1, false + } + contents := util.GetGJSONBytesNoCopy(payload, "request.contents") + if !contents.IsArray() { + return -1, -1, false + } + for ci, content := range contents.Array() { + parts := content.Get("parts") + if !parts.IsArray() { + continue + } + for pi, part := range parts.Array() { + fc := part.Get("functionCall") + if fc.Exists() && strings.TrimSpace(fc.Get("id").String()) == callID { + return ci, pi, true + } + } + } + return -1, -1, false +} + +func legacyAntigravityFunctionCallPartLocationForReplayWithSchemas(payload []byte, itemResult gjson.Result, toolSchemas map[string]any) (contentIndex int, partIndex int, ok bool) { + name := strings.TrimSpace(itemResult.Get("name").String()) + args := itemResult.Get("args") + if name == "" || !args.Exists() { + return -1, -1, false + } + callID := strings.TrimSpace(itemResult.Get("call_id").String()) + if callID == "" { + callID = strings.TrimSpace(itemResult.Get("id").String()) + } + candidateIDs := []string{callID} + if stableID := util.GeminiClaudeToolUseID(callID, name, args.Raw); stableID != "" && stableID != callID { + candidateIDs = append(candidateIDs, stableID) + } + for _, candidateID := range candidateIDs { + if candidateID == "" { + continue + } + ci, pi, found := legacyAntigravityFunctionCallPartLocation(payload, candidateID) + if !found { + continue + } + if legacyAntigravityReplayItemContextMatches(payload, itemResult, ci) { + fc := gjson.GetBytes(payload, fmt.Sprintf("request.contents.%d.parts.%d.functionCall", ci, pi)) + if antigravityFunctionCallMatchesReplayItem(fc, itemResult, toolSchemas) { + return ci, pi, true + } + log.Debugf("antigravity replay: located call %q at contents[%d].parts[%d] but name/args did not match ledger item (opaque_id=%t)", + name, ci, pi, util.IsGeminiClaudeToolUseID(candidateID)) + return -1, -1, false + } + // The candidate ID matched exactly, so callID+name+args are already proven + // identical. Only the surrounding context drifted, which invalidates the + // cached signature but not the tool identity. + log.Debugf("antigravity replay: exact tool ID match for %q at contents[%d].parts[%d] rejected by context hash (opaque_id=%t)", + name, ci, pi, util.IsGeminiClaudeToolUseID(candidateID)) + return -1, -1, false + } + contents := util.GetGJSONBytesNoCopy(payload, "request.contents") + if !contents.IsArray() { + return -1, -1, false + } + contentArr := contents.Array() + cachedCI := int(itemResult.Get("contentIndex").Int()) + if targetOccurrence := itemResult.Get("targetOccurrence"); targetOccurrence.Exists() { + if cachedCI < 0 || cachedCI >= len(contentArr) || !legacyAntigravityReplayItemContextMatches(payload, itemResult, cachedCI) { + return -1, -1, false + } + wantedOccurrence := int(targetOccurrence.Int()) + occurrence := 0 + for pi, part := range contentArr[cachedCI].Get("parts").Array() { + fc := part.Get("functionCall") + if !fc.Exists() || (util.IsGeminiClaudeToolUseID(fc.Get("id").String()) && fc.Get("id").String() != util.GeminiClaudeToolUseID(callID, name, args.Raw)) || !antigravityFunctionCallMatchesReplayItem(fc, itemResult, toolSchemas) { + continue + } + if occurrence == wantedOccurrence { + return cachedCI, pi, true + } + occurrence++ + } + return -1, -1, false + } + + matches := make([][2]int, 0, 1) + for ci, content := range contentArr { + if !legacyAntigravityReplayItemContextMatches(payload, itemResult, ci) { + continue + } + for pi, part := range content.Get("parts").Array() { + fc := part.Get("functionCall") + if !fc.Exists() || (util.IsGeminiClaudeToolUseID(fc.Get("id").String()) && fc.Get("id").String() != util.GeminiClaudeToolUseID(callID, name, args.Raw)) { + continue + } + if antigravityFunctionCallMatchesReplayItem(fc, itemResult, toolSchemas) { + matches = append(matches, [2]int{ci, pi}) + } + } + } + if len(matches) == 1 { + return matches[0][0], matches[0][1], true + } + return -1, -1, false +} + +func legacyAntigravityFunctionCallProvenanceLocation(payload []byte, itemResult gjson.Result, toolSchemas map[string]any) (contentIndex int, partIndex int, ok bool) { + name := strings.TrimSpace(itemResult.Get("name").String()) + args := itemResult.Get("args") + callID := strings.TrimSpace(itemResult.Get("call_id").String()) + if name == "" || !args.Exists() || callID == "" { + return -1, -1, false + } + stableID := util.GeminiClaudeToolUseID(callID, name, args.Raw) + if stableID == "" || stableID == callID { + return -1, -1, false + } + ci, pi, found := legacyAntigravityFunctionCallPartLocation(payload, stableID) + if !found { + return -1, -1, false + } + fc := gjson.GetBytes(payload, fmt.Sprintf("request.contents.%d.parts.%d.functionCall", ci, pi)) + if !antigravityFunctionCallMatchesReplayItem(fc, itemResult, toolSchemas) { + return -1, -1, false + } + return ci, pi, true +} + +func legacyAntigravityRequestHasThoughtSignatureAt(payload []byte, itemResult gjson.Result) bool { + partPath, ok := legacyAntigravityThoughtSignatureReplayPartPath(payload, itemResult) + if !ok { + return false + } + return antigravityHasNativeThoughtSignature(gjson.GetBytes(payload, partPath+".thoughtSignature").String()) +} + +func legacyAntigravityThoughtSignatureReplayPartPath(payload []byte, itemResult gjson.Result) (string, bool) { + ci := int(itemResult.Get("contentIndex").Int()) + contents := util.GetGJSONBytesNoCopy(payload, "request.contents") + if !contents.IsArray() { + return "", false + } + contentArr := contents.Array() + if ci < 0 || ci >= len(contentArr) || !strings.EqualFold(strings.TrimSpace(contentArr[ci].Get("role").String()), "model") { + return "", false + } + parts := contentArr[ci].Get("parts") + if !parts.IsArray() { + return "", false + } + partArr := parts.Array() + targetKind := strings.TrimSpace(itemResult.Get("targetKind").String()) + targetHash := strings.TrimSpace(itemResult.Get("targetHash").String()) + // A target hash pins the signature to a part whose own bytes are unchanged, + // which is all Gemini validates: the signature's own integrity, never its + // binding to the surrounding history. Drift elsewhere in the conversation + // therefore costs this signature nothing, so it is deliberately not gated on + // the context fingerprint. The fallback below has no such proof and stays + // gated. + if targetHash != "" { + if targetOccurrence := itemResult.Get("targetOccurrence"); targetOccurrence.Exists() { + wanted := int(targetOccurrence.Int()) + occurrence := 0 + for pi, part := range partArr { + kind, fingerprint := antigravityReplayPartFingerprint(part) + if fingerprint != targetHash || (targetKind != "" && kind != targetKind) { + continue + } + if occurrence == wanted { + return fmt.Sprintf("request.contents.%d.parts.%d", ci, pi), true + } + occurrence++ + } + return "", false + } + pi := int(itemResult.Get("partIndex").Int()) + if pi >= 0 && pi < len(partArr) { + kind, fingerprint := antigravityReplayPartFingerprint(partArr[pi]) + if fingerprint == targetHash && (targetKind == "" || kind == targetKind) { + return fmt.Sprintf("request.contents.%d.parts.%d", ci, pi), true + } + } + for pi, part := range partArr { + kind, fingerprint := antigravityReplayPartFingerprint(part) + if fingerprint == targetHash && (targetKind == "" || kind == targetKind) { + return fmt.Sprintf("request.contents.%d.parts.%d", ci, pi), true + } + } + return "", false + } + + // No target hash: nothing proves which part this signature belongs to, so + // only a matching context fingerprint makes the positional guess safe. + if !legacyAntigravityReplayItemContextMatches(payload, itemResult, ci) { + return "", false + } + pi := int(itemResult.Get("partIndex").Int()) + if pi >= 0 && pi < len(partArr) && partArr[pi].Type != gjson.Null { + if kind, _ := antigravityReplayPartFingerprint(partArr[pi]); kind != "" { + return fmt.Sprintf("request.contents.%d.parts.%d", ci, pi), true + } + } + // Legacy cache entries may point at a streamed signature-only part after + // multiple text chunks. Attach them to the last semantic part in the same + // model content, never to a different turn. + for candidate := len(partArr) - 1; candidate >= 0; candidate-- { + if kind, _ := antigravityReplayPartFingerprint(partArr[candidate]); kind != "" { + return fmt.Sprintf("request.contents.%d.parts.%d", ci, candidate), true + } + } + return "", false +} + +func legacyAntigravityReplayContextFingerprint(payload []byte, beforeContentIndex int) string { + contents := util.GetGJSONBytesNoCopy(payload, "request.contents") + if !contents.IsArray() || beforeContentIndex < 0 { + return "" + } + contentArr := contents.Array() + if beforeContentIndex > len(contentArr) { + return "" + } + var context strings.Builder + for _, path := range []string{"request.systemInstruction", "request.tools", "request.toolConfig"} { + if value := gjson.GetBytes(payload, path); value.Exists() { + context.WriteString(path) + context.WriteByte('\x00') + context.Write(antigravityCanonicalReplayJSON([]byte(value.Raw))) + context.WriteByte('\x00') + } + } + for ci := 0; ci < beforeContentIndex; ci++ { + content := contentArr[ci] + context.WriteString(strings.ToLower(strings.TrimSpace(content.Get("role").String()))) + context.WriteByte('\x00') + parts := content.Get("parts") + if !parts.IsArray() { + continue + } + parts.ForEach(func(_, part gjson.Result) bool { + normalized := []byte(part.Raw) + for _, signaturePath := range []string{"thoughtSignature", "thought_signature", "extra_content.google.thought_signature"} { + normalized, _ = sjson.DeleteBytes(normalized, signaturePath) + } + context.Write(antigravityCanonicalReplayJSON(normalized)) + context.WriteByte('\x00') + return true + }) + } + if context.Len() == 0 { + return "" + } + sum := sha256.Sum256([]byte(context.String())) + return fmt.Sprintf("%x", sum[:]) +} + +func legacyAntigravityReplayItemContextMatches(payload []byte, itemResult gjson.Result, contentIndex int) bool { + expected := strings.TrimSpace(itemResult.Get("contextHash").String()) + return expected == "" || expected == legacyAntigravityReplayContextFingerprint(payload, contentIndex) +} + +func legacyAntigravitySetReplayItemContextHash(item []byte, payload []byte, contentIndex int) []byte { + if contextHash := legacyAntigravityReplayContextFingerprint(payload, contentIndex); contextHash != "" { + item, _ = sjson.SetBytes(item, "contextHash", contextHash) + } + return item +} + +func legacyInsertAntigravityReasoningReplayItemsWithSchemas(payload []byte, items [][]byte, toolSchemas map[string]any) ([]byte, bool) { + out := payload + changed := false + for _, item := range items { + itemResult := gjson.ParseBytes(item) + switch strings.TrimSpace(itemResult.Get("type").String()) { + case "thought_signature": + sig := strings.TrimSpace(itemResult.Get("thoughtSignature").String()) + if sig == "" { + continue + } + partPath, exists := legacyAntigravityThoughtSignatureReplayPartPath(out, itemResult) + if !exists { + continue + } + path := partPath + ".thoughtSignature" + if antigravityHasNativeThoughtSignature(gjson.GetBytes(out, path).String()) { + continue + } + ci := int(itemResult.Get("contentIndex").Int()) + out = antigravityRemoveThoughtSignatureFromOtherParts(out, ci, sig, partPath) + updated, err := sjson.SetBytes(out, path, sig) + if err != nil { + continue + } + out = updated + changed = true + case "function_call_part": + updated, ok := legacyMergeAntigravityFunctionCallPartReplayWithSchemas(out, itemResult, toolSchemas) + if ok { + out = updated + changed = true + } + } + } + return out, changed +} + +func legacyMergeAntigravityFunctionCallPartReplayWithSchemas(payload []byte, itemResult gjson.Result, toolSchemas map[string]any) ([]byte, bool) { + name := strings.TrimSpace(itemResult.Get("name").String()) + args := itemResult.Get("args") + callID := strings.TrimSpace(itemResult.Get("call_id").String()) + sig := strings.TrimSpace(itemResult.Get("thoughtSignature").String()) + if name == "" || !args.Exists() { + return payload, false + } + if ci, pi, exists := legacyAntigravityFunctionCallPartLocationForReplayWithSchemas(payload, itemResult, toolSchemas); exists { + _, allowLegacyIDRestore := toolSchemas[name] + return restoreAntigravityNativeFunctionCallReplay(payload, ci, pi, itemResult, allowLegacyIDRestore, true) + } + // The context drifted, but an exact opaque ID match still proves this call's + // identity. Gemini validates a thought signature's own integrity and nothing + // about the history around it, so the drift costs the signature nothing: restore + // the native call and its signature rather than making the model re-reason. + if ci, pi, exists := legacyAntigravityFunctionCallProvenanceLocation(payload, itemResult, toolSchemas); exists { + return restoreAntigravityNativeFunctionCallReplay(payload, ci, pi, itemResult, false, true) + } + if callID != "" { + stableID := util.GeminiClaudeToolUseID(callID, name, args.Raw) + if legacyAntigravityPayloadHasFunctionCallID(payload, callID) || (stableID != "" && legacyAntigravityPayloadHasFunctionCallID(payload, stableID)) { + // The call is already in the history under its native or Claude-facing + // ID, and neither lookup above accepted it, so the client changed it. + // Never replay an opaque signature onto that changed call, and never + // insert a second copy of it further down. + return payload, false + } + if frIndex, currentResponseID, ok := legacyAntigravityFunctionResponseContentIndexForReplay(payload, itemResult); ok { + parallelModelIndex := frIndex - 1 + if parallelModelIndex >= 0 && strings.EqualFold(strings.TrimSpace(gjson.GetBytes(payload, fmt.Sprintf("request.contents.%d.role", parallelModelIndex)).String()), "model") && legacyAntigravityReplayItemContextMatches(payload, itemResult, parallelModelIndex) { + if updated, appended := appendAntigravityFunctionCallToModelContent(payload, parallelModelIndex, name, callID, sig, args); appended { + return restoreAntigravityFunctionResponseReplayIdentity(updated, currentResponseID, callID, name), true + } + } + if legacyAntigravityReplayItemContextMatches(payload, itemResult, frIndex) { + if updated, inserted := insertAntigravityModelFunctionCallBeforeContent(payload, frIndex, name, callID, sig, args); inserted { + return restoreAntigravityFunctionResponseReplayIdentity(updated, currentResponseID, callID, name), true + } + } + } + } else { + // Without a native call ID, only an exact semantic match is safe. Never + // put an opaque signature on a different call at the old numeric slot. + return payload, false + } + + ci := antigravityReasoningReplayResolveContentIndex(payload, int(itemResult.Get("contentIndex").Int())) + if ci < 0 || !legacyAntigravityReplayItemContextMatches(payload, itemResult, ci) { + return payload, false + } + pi := int(itemResult.Get("partIndex").Int()) + out := payload + changed := false + + partPath, exists := antigravityExistingReplayPartPath(out, ci, pi) + if !exists { + fc := map[string]any{"name": name} + if callID != "" { + fc["id"] = callID + } + if args.Type == gjson.String { + fc["args"] = args.String() + } else { + var parsed any + if json.Unmarshal([]byte(args.Raw), &parsed) == nil { + fc["args"] = parsed + } + } + part := map[string]any{"functionCall": fc} + if sig != "" { + part["thoughtSignature"] = sig + } + if updated, err := sjson.SetBytes(out, antigravityReplayPartWritePath(out, ci, pi), part); err == nil { + return updated, true + } + return payload, false + } + + pathSig := partPath + ".thoughtSignature" + if sig != "" && !antigravityHasNativeThoughtSignature(gjson.GetBytes(out, pathSig).String()) { + out = antigravityRemoveThoughtSignatureFromOtherParts(out, ci, sig, partPath) + if updated, err := sjson.SetBytes(out, pathSig, sig); err == nil { + out = updated + changed = true + } + } + pathFC := partPath + ".functionCall" + if !gjson.GetBytes(out, pathFC).Exists() { + fc := map[string]any{"name": name} + if callID != "" { + fc["id"] = callID + } + if args.Type == gjson.String { + fc["args"] = args.String() + } else { + var parsed any + if json.Unmarshal([]byte(args.Raw), &parsed) == nil { + fc["args"] = parsed + } + } + if updated, err := sjson.SetBytes(out, pathFC, fc); err == nil { + out = updated + changed = true + } + } + return out, changed +} diff --git a/internal/runtime/executor/antigravity_reasoning_replay_test.go b/internal/runtime/executor/antigravity_reasoning_replay_test.go index 7ab99fce..29e2f5de 100644 --- a/internal/runtime/executor/antigravity_reasoning_replay_test.go +++ b/internal/runtime/executor/antigravity_reasoning_replay_test.go @@ -110,7 +110,7 @@ func TestPrepareAntigravityGeminiReasoningReplayPayloadKeepsCacheForClientMalfor payload := []byte(`{"sessionId":"client-malformed-history","request":{"contents":[{"role":"model","parts":[{"text":"answer"}]},{"role":"model","parts":[{"functionResponse":{"id":"orphan","name":"run","response":{"result":"bad"}}}]}]}}`) kind, fingerprint := antigravityReplayPartFingerprint(gjson.Parse(`{"text":"answer"}`)) item := buildAntigravityThoughtSignatureItem(0, 0, "valid-cache-signature-123456789", kind, fingerprint) - item = antigravitySetReplayItemContextHash(item, payload, 0) + item = antigravityReplayItemContextHashForTest(item, payload, 0) if !internalcache.CacheAntigravityReasoningReplayItems(model, sessionKey, [][]byte{item}) { t.Fatal("cache write failed") } @@ -832,7 +832,7 @@ func TestPrepareAntigravityGeminiReasoningReplayReplacesIDLessFunctionCallBypass func TestAntigravityReasoningReplayContextFingerprintCanonicalizesJSON(t *testing.T) { payload1 := []byte(`{"request":{"tools":[{"functionDeclarations":[{"name":"run","parameters":{"type":"object","properties":{"a":{"type":"string"},"b":{"type":"number"}}}}]}],"contents":[{"role":"user","parts":[{"text":"turn"}]},{"role":"model","parts":[{"functionCall":{"name":"run","args":{"a":"x","b":2}}}]}]}}`) payload2 := []byte(`{"request":{"tools":[{"functionDeclarations":[{"parameters":{"properties":{"b":{"type":"number"},"a":{"type":"string"}},"type":"object"},"name":"run"}]}],"contents":[{"parts":[{"text":"turn"}],"role":"user"},{"parts":[{"functionCall":{"args":{"b":2,"a":"x"},"name":"run"}}],"role":"model"}]}}`) - if got1, got2 := antigravityReplayContextFingerprint(payload1, 2), antigravityReplayContextFingerprint(payload2, 2); got1 == "" || got1 != got2 { + if got1, got2 := newAntigravityReplayRequestIndex(payload1).contextFingerprint(2), newAntigravityReplayRequestIndex(payload2).contextFingerprint(2); got1 == "" || got1 != got2 { t.Fatalf("canonical context hashes differ: %q vs %q", got1, got2) } key1 := antigravityFunctionCallKey("run", `{"a":"x","b":2}`, "") @@ -980,7 +980,7 @@ func TestAntigravityReasoningReplayPreservesRepeatedIDLessCallsAcrossSplitSSEPar func TestAntigravityReasoningReplayLegacyAmbiguousIDLessCallFailsClosed(t *testing.T) { item := []byte(`{"type":"function_call_part","contentIndex":1,"partIndex":1,"name":"run_command","args":{"command":"same"},"thoughtSignature":"legacy-ambiguous-signature-123456"}`) payload := []byte(`{"request":{"contents":[{"role":"user","parts":[{"text":"run"}]},{"role":"model","parts":[{"functionCall":{"name":"run_command","args":{"command":"same"}}},{"functionCall":{"name":"run_command","args":{"command":"same"}}}]}]}}`) - out, changed := insertAntigravityReasoningReplayItems(payload, [][]byte{item}) + out, changed := insertAntigravityReasoningReplayItemsWithSchemas(newAntigravityReplayRequestIndex(payload), payload, [][]byte{item}, nil) if changed || strings.Contains(string(out), "legacy-ambiguous-signature") { t.Fatalf("legacy ambiguous ID-less replay must fail closed: changed=%v body=%s", changed, out) } @@ -1047,7 +1047,7 @@ func TestPrepareAntigravityGeminiReasoningReplayKeepsTextSignatureOnContextDrift kind, fingerprint := antigravityReplayPartFingerprint(gjson.Parse(`{"text":"same answer"}`)) item := buildAntigravityThoughtSignatureItem(1, 0, "fingerprinted-signature-123456", kind, fingerprint) originalPayload := []byte(`{"sessionId":"rebuilt","request":{"contents":[{"role":"user","parts":[{"text":"old context"}]},{"role":"model","parts":[{"text":"same answer"}]},{"role":"user","parts":[{"text":"old next"}]}]}}`) - item = antigravitySetReplayItemContextHash(item, originalPayload, 1) + item = antigravityReplayItemContextHashForTest(item, originalPayload, 1) internalcache.CacheAntigravityReasoningReplayItems("gemini-3.6-flash-high", "session:rebuilt", [][]byte{item}) payload := []byte(`{"sessionId":"rebuilt","request":{"contents":[{"role":"user","parts":[{"text":"new context"}]},{"role":"model","parts":[{"text":"same answer"}]},{"role":"user","parts":[{"text":"next"}]}]}}`) @@ -1268,13 +1268,6 @@ func TestAntigravityReplayToolCallKeysUsesNativeFunctionCallID(t *testing.T) { } } -func TestAntigravityRequestHasMatchingFunctionResponseWhitespaceCallID(t *testing.T) { - item := gjson.Parse(`{"call_id":" "}`) - if !antigravityRequestHasMatchingFunctionResponse(nil, item) { - t.Fatal("whitespace-only call_id should be treated as empty => true") - } -} - func TestPrepareAntigravityGeminiReasoningReplayRepairsSequentialCompactedUnknownResponseName(t *testing.T) { internalcache.ClearAntigravityReasoningReplayCache() t.Cleanup(internalcache.ClearAntigravityReasoningReplayCache) -- 2.51.2