diff --git a/internal/backend/forwarder/artifacts.go b/internal/backend/forwarder/artifacts.go index c194cfd..715a50c 100644 --- a/internal/backend/forwarder/artifacts.go +++ b/internal/backend/forwarder/artifacts.go @@ -23,10 +23,10 @@ type artifactRecorder struct { } type artifactSession struct { - conversationID string - turnSeq int64 - requestPayload map[string]any - summaryPayload map[string]any + conversationID string + turnSeq int64 + requestPrefix requestArtifactPrefix + hasRequestPrefix bool } type requestArtifactPrefix struct { @@ -57,14 +57,22 @@ func (recorder *artifactRecorder) RecordLLMRequest(requestID string, _ string, m if err != nil { return "", err } - session.requestPayload = cloneStringAnyMap(payload) + prefix, hasPrefix, prefixErr := decodeRequestPrefixPayload(payload) + session.requestPrefix = requestArtifactPrefix{} + session.hasRequestPrefix = false + if prefixErr == nil && hasPrefix && prefix != nil { + session.requestPrefix = *prefix + session.hasRequestPrefix = true + } recorder.mu.Lock() recorder.sessions[artifactSessionKey(requestID, modelCallID)] = session recorder.mu.Unlock() - if prefix, ok, err := decodeRequestPrefixPayload(session.requestPayload); err == nil && ok && prefix != nil { + if prefixErr == nil && hasPrefix && prefix != nil { recorder.persistLatestRequestPrefix(session.conversationID, requestID, modelCallID, prefix) } - recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_request", session.requestPayload) + // The payload is serialized synchronously for debug output and is deliberately + // not retained in sessions: it can contain the entire replayed conversation. + recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_request", payload) return "", nil } @@ -86,17 +94,16 @@ func (recorder *artifactRecorder) RecordLLMSummary(requestID string, _ string, m if err != nil { return "", err } - session.summaryPayload = cloneStringAnyMap(payload) - recorder.mu.Lock() - recorder.sessions[artifactSessionKey(requestID, modelCallID)] = session - recorder.mu.Unlock() - if prefix, ok, err := decodeRequestPrefixPayload(session.requestPayload); err == nil && ok && prefix != nil { - if tokens := readInt64Value(session.summaryPayload["prompt_tokens_total"]); tokens > 0 { - prefix.PromptTokensTotal = tokens + if session.hasRequestPrefix { + if tokens := readInt64Value(payload["prompt_tokens_total"]); tokens > 0 { + session.requestPrefix.PromptTokensTotal = tokens } - recorder.persistLatestRequestPrefix(session.conversationID, requestID, modelCallID, prefix) + recorder.mu.Lock() + recorder.sessions[artifactSessionKey(requestID, modelCallID)] = session + recorder.mu.Unlock() + recorder.persistLatestRequestPrefix(session.conversationID, requestID, modelCallID, &session.requestPrefix) } - recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_summary", session.summaryPayload) + recorder.debug.LogProviderArtifact(context.Background(), requestID, session.conversationID, modelCallID, "llm_summary", payload) return "", nil } diff --git a/internal/backend/forwarder/json_helpers.go b/internal/backend/forwarder/json_helpers.go index 95cec74..ad06599 100644 --- a/internal/backend/forwarder/json_helpers.go +++ b/internal/backend/forwarder/json_helpers.go @@ -10,21 +10,6 @@ import ( modeladapter "cursor/internal/backend/agent/model" ) -func cloneStringAnyMap(value map[string]any) map[string]any { - if len(value) == 0 { - return nil - } - payload, err := json.Marshal(value) - if err != nil { - return nil - } - var decoded map[string]any - if err := json.Unmarshal(payload, &decoded); err != nil { - return nil - } - return decoded -} - func parseRFC3339Time(value string) time.Time { trimmed := strings.TrimSpace(value) if trimmed == "" { diff --git a/internal/backend/forwarder/provider.go b/internal/backend/forwarder/provider.go index 23ec356..b37e317 100644 --- a/internal/backend/forwarder/provider.go +++ b/internal/backend/forwarder/provider.go @@ -25,6 +25,7 @@ func (gateway *DefaultProviderGateway) StartStream(ctx context.Context, req Prov if ctx == nil { ctx = context.Background() } + defer releaseArtifactSession(req.Observer, req.RequestID, req.ModelCallID) requestKnobs := make(map[string]any, len(req.RequestKnobs)+2) for key, value := range req.RequestKnobs { requestKnobs[key] = value @@ -60,3 +61,15 @@ func (gateway *DefaultProviderGateway) StartStream(ctx context.Context, req Prov } return nil } + +type artifactSessionCleaner interface { + ClearActiveArtifacts(requestID string, modelCallID string) +} + +func releaseArtifactSession(observer modeladapter.LLMArtifactObserver, requestID string, modelCallID string) { + cleaner, ok := observer.(artifactSessionCleaner) + if !ok { + return + } + cleaner.ClearActiveArtifacts(requestID, modelCallID) +}