diff --git a/internal/backend/forwarder/actor.go b/internal/backend/forwarder/actor.go index 2614cc1..a712976 100644 --- a/internal/backend/forwarder/actor.go +++ b/internal/backend/forwarder/actor.go @@ -23,6 +23,7 @@ const ( TurnPhaseWaitingExternal TurnPhase = "waiting_external" TurnPhaseAwaitingUser TurnPhase = "awaiting_user" TurnPhaseCompacting TurnPhase = "compacting" + TurnPhaseCheckpointing TurnPhase = "checkpointing" TurnPhaseCompleted TurnPhase = "completed" TurnPhaseFailed TurnPhase = "failed" TurnPhaseCanceled TurnPhase = "canceled" @@ -66,6 +67,7 @@ const ( streamTimerNonStreamingRecovery streamTimerKind = "non_streaming_recovery" streamTimerShellForeground streamTimerKind = "shell_foreground" streamTimerShellTransportClose streamTimerKind = "shell_transport_close" + streamTimerCheckpointBlobs streamTimerKind = "checkpoint_blobs" streamTimerOrphanCancel streamTimerKind = "orphan_cancel" ) @@ -318,6 +320,9 @@ func (service *Service) handleStreamCommand(stream *ActiveStream, command stream case streamCommandCancel: return service.handleCancelIntent(command.Intent) case streamCommandMetadata: + if strings.TrimSpace(command.Intent.Kind) == "kv_result" { + return service.handleCheckpointBlobResult(stream, command.Intent.KVClientMessage) + } return service.handleMetadataIntent(command.Intent) case streamCommandExecResult: return service.handleExecResult(command.Intent) @@ -1003,6 +1008,8 @@ func (service *Service) handleTimerEvent(stream *ActiveStream, payload *streamTi return nil } return service.recoverShellWithoutTerminal(stream, current, shellRecoveryReasonTransportClosed) + case streamTimerCheckpointBlobs: + return service.handleCheckpointBlobTimeout(stream) case streamTimerOrphanCancel: stream.mu.Lock() subscriberCount := len(stream.Subscribers) diff --git a/internal/backend/forwarder/blob_sync.go b/internal/backend/forwarder/blob_sync.go new file mode 100644 index 0000000..eb79b33 --- /dev/null +++ b/internal/backend/forwarder/blob_sync.go @@ -0,0 +1,383 @@ +package forwarder + +import ( + "encoding/hex" + "fmt" + "log" + "strings" + "time" + + "google.golang.org/protobuf/proto" + + "cursor/gen/agentv1" +) + +const ( + checkpointBlobWriteTimeout = 10 * time.Second + checkpointBlobCacheIdleTTL = 6 * time.Hour + checkpointBlobCacheMaxConversations = 256 +) + +type checkpointBlobCacheEntry struct { + Confirmed map[string]struct{} + LastAccess time.Time +} + +func checkpointBlobKey(id []byte) string { + return string(id) +} + +func checkpointBlobHex(key string) string { + return hex.EncodeToString([]byte(key)) +} + +func (service *Service) confirmedCheckpointBlob(conversationID string, key string) bool { + if service == nil || key == "" { + return false + } + service.checkpointBlobMu.Lock() + defer service.checkpointBlobMu.Unlock() + conversationID = strings.TrimSpace(conversationID) + entry := service.checkpointBlobs[conversationID] + if entry == nil { + return false + } + entry.LastAccess = time.Now().UTC() + _, ok := entry.Confirmed[key] + return ok +} + +func (service *Service) confirmCheckpointBlob(conversationID string, key string) { + if service == nil || key == "" { + return + } + service.checkpointBlobMu.Lock() + defer service.checkpointBlobMu.Unlock() + if service.checkpointBlobs == nil { + service.checkpointBlobs = make(map[string]*checkpointBlobCacheEntry) + } + conversationID = strings.TrimSpace(conversationID) + now := time.Now().UTC() + entry := service.checkpointBlobs[conversationID] + if entry == nil { + entry = &checkpointBlobCacheEntry{Confirmed: make(map[string]struct{})} + service.checkpointBlobs[conversationID] = entry + } + entry.Confirmed[key] = struct{}{} + entry.LastAccess = now + service.pruneCheckpointBlobCacheLocked(now) +} + +func (service *Service) pruneCheckpointBlobCacheLocked(now time.Time) { + if service == nil || len(service.checkpointBlobs) == 0 { + return + } + cutoff := now.Add(-checkpointBlobCacheIdleTTL) + for conversationID, entry := range service.checkpointBlobs { + if entry == nil || entry.LastAccess.Before(cutoff) { + delete(service.checkpointBlobs, conversationID) + } + } + for len(service.checkpointBlobs) > checkpointBlobCacheMaxConversations { + oldestConversationID := "" + oldestAccess := now + for conversationID, entry := range service.checkpointBlobs { + if entry == nil || oldestConversationID == "" || entry.LastAccess.Before(oldestAccess) { + oldestConversationID = conversationID + if entry != nil { + oldestAccess = entry.LastAccess + } + } + } + if oldestConversationID == "" { + return + } + delete(service.checkpointBlobs, oldestConversationID) + } +} + +func checkpointCompletionAction(completion *pendingTurnCompletion) checkpointTerminalAction { + if completion == nil { + return checkpointTerminalAction{kind: checkpointTerminalActionNone} + } + return checkpointTerminalAction{ + kind: checkpointTerminalActionComplete, + completion: *clonePendingTurnCompletion(completion), + } +} + +func checkpointCancellationAction(message string) checkpointTerminalAction { + return checkpointTerminalAction{ + kind: checkpointTerminalActionCancel, + cancelMessage: firstNonEmpty(strings.TrimSpace(message), "[canceled] User aborted request"), + } +} + +func mergeCheckpointTerminalAction(current checkpointTerminalAction, incoming checkpointTerminalAction) checkpointTerminalAction { + switch { + case incoming.kind == checkpointTerminalActionCancel: + return incoming + case current.kind == checkpointTerminalActionCancel: + return current + case incoming.kind == checkpointTerminalActionComplete: + return incoming + default: + return current + } +} + +func (action checkpointTerminalAction) completionValue() *pendingTurnCompletion { + if action.kind != checkpointTerminalActionComplete { + return nil + } + return clonePendingTurnCompletion(&action.completion) +} + +func (service *Service) queueCheckpointProjection(stream *ActiveStream, projection *CheckpointProjection, terminalAction checkpointTerminalAction) error { + if service == nil || stream == nil || projection == nil || projection.State == nil { + return nil + } + state, ok := proto.Clone(projection.State).(*agentv1.ConversationStateStructure) + if !ok || state == nil { + return fmt.Errorf("clone checkpoint state") + } + + stream.mu.Lock() + if stream.PendingCheckpointBlobWrites == nil { + stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite) + } + if stream.PendingCheckpointBlobRequests == nil { + stream.PendingCheckpointBlobRequests = make(map[string]uint32) + } + stream.NextCheckpointRevision++ + if stream.NextCheckpointRevision == 0 { + stream.NextCheckpointRevision++ + } + revision := stream.NextCheckpointRevision + if stream.PendingCheckpoint != nil { + terminalAction = mergeCheckpointTerminalAction(stream.PendingCheckpoint.TerminalAction, terminalAction) + } + required := make(map[string]struct{}, len(projection.Blobs)) + toWrite := make([]struct { + requestID uint32 + blob CheckpointBlob + }, 0, len(projection.Blobs)) + for _, blob := range projection.Blobs { + key := checkpointBlobKey(blob.ID) + if key == "" { + continue + } + required[key] = struct{}{} + if service.confirmedCheckpointBlob(stream.ConversationID, key) { + continue + } + if _, pending := stream.PendingCheckpointBlobRequests[key]; pending { + continue + } + stream.NextCheckpointBlobRequestID++ + if stream.NextCheckpointBlobRequestID == 0 { + stream.NextCheckpointBlobRequestID++ + } + requestID := stream.NextCheckpointBlobRequestID + stream.PendingCheckpointBlobWrites[requestID] = pendingCheckpointBlobWrite{ + Key: key, + Revision: revision, + } + stream.PendingCheckpointBlobRequests[key] = requestID + toWrite = append(toWrite, struct { + requestID uint32 + blob CheckpointBlob + }{requestID: requestID, blob: blob}) + } + stream.PendingCheckpoint = &pendingCheckpointPublish{ + Revision: revision, + State: state, + Required: required, + TerminalAction: terminalAction, + } + if terminalAction.kind != checkpointTerminalActionNone { + stream.Phase = TurnPhaseCheckpointing + } + for requestID, write := range stream.PendingCheckpointBlobWrites { + if _, stillRequired := required[write.Key]; stillRequired { + continue + } + delete(stream.PendingCheckpointBlobWrites, requestID) + delete(stream.PendingCheckpointBlobRequests, write.Key) + } + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + + for _, item := range toWrite { + if err := service.broker.Publish(stream.RequestID, StreamEvent{ + Message: buildSetCheckpointBlobMessage(item.requestID, item.blob), + }); err != nil { + service.discardPendingCheckpoint(stream, fmt.Errorf("publish checkpoint blob write: %w", err)) + return err + } + } + if len(toWrite) > 0 { + service.scheduleStreamTimer( + stream, + providerTimerKey(streamTimerCheckpointBlobs, ""), + checkpointBlobWriteTimeout, + streamTimerCheckpointBlobs, + "", + 0, + "checkpoint blob write timeout", + ) + } + return service.publishReadyCheckpoint(stream) +} + +func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion { + if completion == nil { + return nil + } + cloned := *completion + return &cloned +} + +func (service *Service) handleCheckpointBlobResult(stream *ActiveStream, message *agentv1.KvClientMessage) error { + if service == nil || stream == nil || message == nil { + return nil + } + result := message.GetSetBlobResult() + if result == nil { + return nil + } + stream.mu.Lock() + write, ok := stream.PendingCheckpointBlobWrites[message.GetId()] + pendingRequiresBlob := false + if ok { + delete(stream.PendingCheckpointBlobWrites, message.GetId()) + delete(stream.PendingCheckpointBlobRequests, write.Key) + if stream.PendingCheckpoint != nil { + _, pendingRequiresBlob = stream.PendingCheckpoint.Required[write.Key] + } + } + conversationID := stream.ConversationID + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + if !ok { + return nil + } + if result.GetError() != nil { + if !pendingRequiresBlob { + return service.publishReadyCheckpoint(stream) + } + return service.abandonPendingCheckpoint(stream, fmt.Errorf( + "write checkpoint blob %s: %s", + checkpointBlobHex(write.Key), + firstNonEmpty(result.GetError().GetMessage(), "client blob store rejected write"), + )) + } + service.confirmCheckpointBlob(conversationID, write.Key) + return service.publishReadyCheckpoint(stream) +} + +func (service *Service) publishReadyCheckpoint(stream *ActiveStream) error { + if service == nil || stream == nil { + return nil + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + if pending == nil { + stream.mu.Unlock() + return nil + } + for key := range pending.Required { + if !service.confirmedCheckpointBlob(stream.ConversationID, key) { + stream.mu.Unlock() + return nil + } + } + stream.PendingCheckpoint = nil + state := pending.State + terminalAction := pending.TerminalAction + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil { + return err + } + switch terminalAction.kind { + case checkpointTerminalActionComplete: + if completion := terminalAction.completionValue(); completion != nil { + return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion) + } + case checkpointTerminalActionCancel: + return service.finishCanceledTurnAfterCheckpoint(stream, terminalAction.cancelMessage) + } + return nil +} + +func (service *Service) discardPendingCheckpoint(stream *ActiveStream, cause error) { + if service == nil || stream == nil { + return + } + stream.mu.Lock() + stream.PendingCheckpoint = nil + stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite) + stream.PendingCheckpointBlobRequests = make(map[string]uint32) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if cause != nil { + log.Printf("forwarder pending checkpoint discarded request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause) + } +} + +func (service *Service) abandonPendingCheckpoint(stream *ActiveStream, cause error) error { + if service == nil || stream == nil { + return nil + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + stream.PendingCheckpoint = nil + stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite) + stream.PendingCheckpointBlobRequests = make(map[string]uint32) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if cause != nil { + log.Printf("forwarder checkpoint blob sync abandoned request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause) + } + if pending != nil { + return service.failTerminalCheckpointSync(stream, cause) + } + return nil +} + +func (service *Service) finishCanceledTurnAfterCheckpoint(stream *ActiveStream, message string) error { + if stream == nil { + return nil + } + service.setTurnPhase(stream, TurnPhaseCanceled) + return service.broker.Cancel(stream.RequestID, firstNonEmpty(strings.TrimSpace(message), "[canceled] User aborted request")) +} + +func (service *Service) failTerminalCheckpointSync(stream *ActiveStream, cause error) error { + if stream == nil { + return nil + } + message := "checkpoint synchronization failed" + if cause != nil && strings.TrimSpace(cause.Error()) != "" { + message = strings.TrimSpace(cause.Error()) + } + service.setTurnPhase(stream, TurnPhaseFailed) + return service.broker.Fail(stream.RequestID, "checkpoint_sync_error", message) +} + +func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error { + if stream == nil { + return nil + } + stream.mu.Lock() + pendingCount := len(stream.PendingCheckpointBlobWrites) + stream.mu.Unlock() + if pendingCount == 0 { + return service.publishReadyCheckpoint(stream) + } + return service.abandonPendingCheckpoint(stream, fmt.Errorf("%d checkpoint blob writes timed out", pendingCount)) +} diff --git a/internal/backend/forwarder/blob_sync_test.go b/internal/backend/forwarder/blob_sync_test.go new file mode 100644 index 0000000..d101a92 --- /dev/null +++ b/internal/backend/forwarder/blob_sync_test.go @@ -0,0 +1,546 @@ +package forwarder + +import ( + "os" + "path/filepath" + "testing" + + "cursor/gen/agentv1" +) + +func TestCheckpointBlobSyncPublishesCheckpointAfterAllWrites(t *testing.T) { + service, stream := testCheckpointBlobService(t) + projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + newAssistantTextEntry(1, "request-1", "hi", "", ""), + })) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + + if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("ReadFromCursor() error = %v", err) + } + if len(events) != len(projection.Blobs) { + t.Fatalf("events before ACK = %d, want %d blob writes", len(events), len(projection.Blobs)) + } + for _, event := range events { + if event.Message.GetKvServerMessage().GetSetBlobArgs() == nil { + t.Fatalf("event before ACK = %#v, want set_blob_args", event.Message) + } + } + + for index, event := range events { + requestID := event.Message.GetKvServerMessage().GetId() + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + t.Fatalf("handleCheckpointBlobResult(%d) error = %v", index, err) + } + } + events, err = service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("ReadFromCursor() after ACK error = %v", err) + } + if len(events) != len(projection.Blobs)+1 { + t.Fatalf("events after ACK = %d, want %d", len(events), len(projection.Blobs)+1) + } + checkpoint := events[len(events)-1].Message.GetConversationCheckpointUpdate() + if checkpoint == nil || len(checkpoint.GetTurns()) != 1 { + t.Fatalf("last event checkpoint = %#v, want one Blob-backed turn", checkpoint) + } +} + +func TestCheckpointBlobSyncRejectDoesNotPublishDanglingCheckpoint(t *testing.T) { + service, stream := testCheckpointBlobService(t) + projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + })) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil || len(events) == 0 { + t.Fatalf("blob write events = %d, err = %v", len(events), err) + } + requestID := events[0].Message.GetKvServerMessage().GetId() + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{Error: &agentv1.Error{Message: "disk full"}}, + }, + }); err != nil { + t.Fatalf("handleCheckpointBlobResult() error = %v", err) + } + events, err = service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("ReadFromCursor() after rejection error = %v", err) + } + for _, event := range events { + if event.Message.GetConversationCheckpointUpdate() != nil { + t.Fatal("rejected Blob write published a dangling checkpoint") + } + } +} + +func TestTerminalCheckpointRejectionFailsInsteadOfCompleting(t *testing.T) { + service, stream := testCheckpointBlobService(t) + projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, stream.RequestID, "hello"), + })) + if err != nil { + t.Fatalf("projection: %v", err) + } + completion := &pendingTurnCompletion{RequestID: stream.RequestID} + if err := service.queueCheckpointProjection(stream, projection, checkpointCompletionAction(completion)); err != nil { + t.Fatalf("queue checkpoint: %v", err) + } + stream.mu.Lock() + var requestID uint32 + for pendingID := range stream.PendingCheckpointBlobWrites { + requestID = pendingID + break + } + stream.mu.Unlock() + if requestID == 0 { + t.Fatal("test did not queue a Blob write") + } + if err := rejectCheckpointBlob(service, stream, requestID, "disk full"); err != nil { + t.Fatalf("reject terminal checkpoint: %v", err) + } + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("ReadFromCursor() error = %v", err) + } + var failed, completed bool + for _, event := range events { + if event.End && event.TerminalErrorCode == "checkpoint_sync_error" { + failed = true + } + if event.Message.GetInteractionUpdate().GetTurnEnded() != nil { + completed = true + } + } + if !failed || completed { + t.Fatalf("terminal checkpoint events failed=%v completed=%v", failed, completed) + } +} + +func TestCheckpointBlobSyncMergesRevisionsWithoutObsoleteFailure(t *testing.T) { + service, stream := testCheckpointBlobService(t) + first, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + })) + if err != nil { + t.Fatalf("first ProjectCheckpointProjection() error = %v", err) + } + latest, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + newAssistantTextEntry(1, "request-1", "latest answer", "", ""), + })) + if err != nil { + t.Fatalf("latest ProjectCheckpointProjection() error = %v", err) + } + if err := service.queueCheckpointProjection(stream, first, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queue first projection: %v", err) + } + firstEvents, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("read first events: %v", err) + } + if err := service.queueCheckpointProjection(stream, latest, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queue latest projection: %v", err) + } + + latestRequired := make(map[string]struct{}, len(latest.Blobs)) + for _, blob := range latest.Blobs { + latestRequired[string(blob.ID)] = struct{}{} + } + var obsoleteRequestID uint32 + for _, event := range firstEvents { + message := event.Message.GetKvServerMessage() + if message == nil || message.GetSetBlobArgs() == nil { + continue + } + if _, required := latestRequired[string(message.GetSetBlobArgs().GetBlobId())]; !required { + obsoleteRequestID = message.GetId() + break + } + } + if obsoleteRequestID == 0 { + t.Fatal("test did not find an obsolete first-revision Blob write") + } + if err := rejectCheckpointBlob(service, stream, obsoleteRequestID, "obsolete write rejected"); err != nil { + t.Fatalf("reject obsolete Blob: %v", err) + } + if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil { + t.Fatalf("acknowledge latest Blob writes: %v", err) + } + + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("read merged events: %v", err) + } + checkpoints := 0 + for _, event := range events { + if checkpoint := event.Message.GetConversationCheckpointUpdate(); checkpoint != nil { + checkpoints++ + if len(checkpoint.GetTurns()) != len(latest.State.GetTurns()) || string(checkpoint.GetTurns()[0]) != string(latest.State.GetTurns()[0]) { + t.Fatalf("published checkpoint is not the latest revision: %#v", checkpoint) + } + } + } + if checkpoints != 1 { + t.Fatalf("published checkpoints = %d, want exactly latest revision", checkpoints) + } +} + +func TestCheckpointBlobSyncCarriesCompletionIntoLatestRevision(t *testing.T) { + service, stream := testCheckpointBlobService(t) + first, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + })) + if err != nil { + t.Fatalf("first projection: %v", err) + } + latest, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + newAssistantTextEntry(1, "request-1", "done", "", ""), + })) + if err != nil { + t.Fatalf("latest projection: %v", err) + } + completion := &pendingTurnCompletion{ + RequestID: stream.RequestID, + Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7}, + } + if err := service.queueCheckpointProjection(stream, first, checkpointCompletionAction(completion)); err != nil { + t.Fatalf("queue completion projection: %v", err) + } + if err := service.queueCheckpointProjection(stream, latest, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queue latest projection: %v", err) + } + if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil { + t.Fatalf("acknowledge latest Blob writes: %v", err) + } + + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("read completion events: %v", err) + } + checkpointIndex, turnEndedIndex, endIndex := -1, -1, -1 + for index, event := range events { + switch { + case event.Message.GetConversationCheckpointUpdate() != nil: + checkpointIndex = index + case event.Message.GetInteractionUpdate().GetTurnEnded() != nil: + turnEndedIndex = index + case event.End: + endIndex = index + } + } + if checkpointIndex < 0 || turnEndedIndex <= checkpointIndex || endIndex <= turnEndedIndex { + t.Fatalf("terminal order checkpoint=%d turn_ended=%d end=%d", checkpointIndex, turnEndedIndex, endIndex) + } +} + +func TestCheckpointBlobSyncTimeoutFailsStreamWithoutPublishingDanglingCheckpoint(t *testing.T) { + service, stream := testCheckpointBlobService(t) + projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + })) + if err != nil { + t.Fatalf("projection: %v", err) + } + if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queue checkpoint: %v", err) + } + if err := service.handleCheckpointBlobTimeout(stream); err != nil { + t.Fatalf("timeout checkpoint: %v", err) + } + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("read timeout events: %v", err) + } + var failed bool + for _, event := range events { + if event.Message.GetConversationCheckpointUpdate() != nil { + t.Fatal("timed-out Blob dependency published a dangling checkpoint") + } + if event.End && event.TerminalErrorCode == "checkpoint_sync_error" { + failed = true + } + } + if !failed { + t.Fatal("timed-out checkpoint did not fail the stream explicitly") + } +} + +func TestCheckpointBlobSyncReusesConversationCacheAcrossRequests(t *testing.T) { + service, firstStream := testCheckpointBlobService(t) + projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + })) + if err != nil { + t.Fatalf("projection: %v", err) + } + if err := service.queueCheckpointProjection(firstStream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queue first request: %v", err) + } + if err := acknowledgePendingCheckpointBlobs(service, firstStream); err != nil { + t.Fatalf("acknowledge first request: %v", err) + } + + secondStream, err := service.broker.OpenStream( + "request-2", firstStream.ConversationID, 2, "default", "default", + agentv1.AgentMode_AGENT_MODE_AGENT, "continue", + ) + if err != nil { + t.Fatalf("OpenStream() second request error = %v", err) + } + if err := service.queueCheckpointProjection(secondStream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queue second request: %v", err) + } + events, err := service.broker.ReadFromCursor(secondStream.RequestID, 0) + if err != nil { + t.Fatalf("read second request events: %v", err) + } + if len(events) != 1 || events[0].Message.GetConversationCheckpointUpdate() == nil { + t.Fatalf("second request events = %#v, want cached immediate checkpoint", events) + } +} + +func TestCancellationReplacesUnconfirmedCheckpointBeforeEnding(t *testing.T) { + service, stream := testCheckpointBlobService(t) + conversation := testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, stream.RequestID, "hello"), + }) + if err := service.replaceCheckpointConversation(stream, conversation); err != nil { + t.Fatalf("replaceCheckpointConversation() error = %v", err) + } + projection, err := service.projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + stream.mu.Lock() + var staleRequestID uint32 + for requestID := range stream.PendingCheckpointBlobWrites { + staleRequestID = requestID + break + } + stream.mu.Unlock() + if staleRequestID == 0 { + t.Fatal("test did not queue an unconfirmed Blob write") + } + + if err := service.handleCancelIntent(InboundIntent{ + Kind: "cancel", + RequestID: stream.RequestID, + CancelReason: "user stopped", + }); err != nil { + t.Fatalf("handleCancelIntent() error = %v", err) + } + stream.mu.Lock() + phase := stream.Phase + status := stream.Status + pendingCheckpoint := stream.PendingCheckpoint + pendingWrites := len(stream.PendingCheckpointBlobWrites) + stream.mu.Unlock() + if phase != TurnPhaseCheckpointing || status != StreamStatusCreated { + t.Fatalf("before checkpoint ACK phase=%s status=%s, want checkpointing/created", phase, status) + } + if pendingCheckpoint == nil || pendingWrites == 0 { + t.Fatalf("before checkpoint ACK pending_checkpoint=%v pending_writes=%d", pendingCheckpoint != nil, pendingWrites) + } + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: staleRequestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + t.Fatalf("stale Blob ACK error = %v", err) + } + if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil { + t.Fatalf("acknowledge cancellation checkpoint: %v", err) + } + stream.mu.Lock() + phase = stream.Phase + status = stream.Status + stream.mu.Unlock() + if phase != TurnPhaseCanceled || status != StreamStatusCanceled { + t.Fatalf("after checkpoint ACK phase=%s status=%s, want canceled", phase, status) + } + assertCanceledEndEvent(t, service, stream) +} + +func TestCancellationMetadataFailureStillEndsStream(t *testing.T) { + service, stream := testCheckpointBlobService(t) + blockingPath := filepath.Join(t.TempDir(), "not-a-directory") + if err := os.WriteFile(blockingPath, []byte("block child creation"), 0o600); err != nil { + t.Fatalf("WriteFile() error = %v", err) + } + service.store = NewConversationFileStore(blockingPath) + if err := service.replaceCheckpointConversation(stream, testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, stream.RequestID, "hello"), + })); err != nil { + t.Fatalf("replaceCheckpointConversation() error = %v", err) + } + + if err := service.handleCancelIntent(InboundIntent{Kind: "cancel", RequestID: stream.RequestID}); err != nil { + t.Fatalf("handleCancelIntent() error = %v", err) + } + if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil { + t.Fatalf("acknowledge cancellation checkpoint: %v", err) + } + assertCanceledEndEvent(t, service, stream) +} + +func TestCheckpointTerminalActionMergePriority(t *testing.T) { + complete := checkpointCompletionAction(&pendingTurnCompletion{RequestID: "complete"}) + cancel := checkpointCancellationAction("user canceled") + none := checkpointTerminalAction{kind: checkpointTerminalActionNone} + + tests := []struct { + name string + current checkpointTerminalAction + incoming checkpointTerminalAction + wantKind checkpointTerminalActionKind + wantID string + }{ + {name: "none then complete", current: none, incoming: complete, wantKind: checkpointTerminalActionComplete, wantID: "complete"}, + {name: "complete then none", current: complete, incoming: none, wantKind: checkpointTerminalActionComplete, wantID: "complete"}, + {name: "complete then cancel", current: complete, incoming: cancel, wantKind: checkpointTerminalActionCancel}, + {name: "cancel then complete", current: cancel, incoming: complete, wantKind: checkpointTerminalActionCancel}, + {name: "cancel then none", current: cancel, incoming: none, wantKind: checkpointTerminalActionCancel}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + merged := mergeCheckpointTerminalAction(test.current, test.incoming) + if merged.kind != test.wantKind { + t.Fatalf("merged kind = %d, want %d", merged.kind, test.wantKind) + } + if test.wantID != "" { + completion := merged.completionValue() + if completion == nil || completion.RequestID != test.wantID { + t.Fatalf("merged completion = %#v, want request_id=%s", completion, test.wantID) + } + } + }) + } +} + +func TestCheckpointTerminalActionIsMutuallyExclusive(t *testing.T) { + completion := &pendingTurnCompletion{RequestID: "request-1"} + action := checkpointCompletionAction(completion) + if action.kind != checkpointTerminalActionComplete || action.completionValue() == nil { + t.Fatalf("completion action = %#v", action) + } + empty := checkpointCompletionAction(nil) + if empty.kind != checkpointTerminalActionNone || empty.completionValue() != nil { + t.Fatalf("empty action = %#v", empty) + } +} + +func assertCanceledEndEvent(t *testing.T, service *Service, stream *ActiveStream) { + t.Helper() + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("ReadFromCursor() error = %v", err) + } + for _, event := range events { + if event.End && event.TerminalErrorCode == "canceled" { + return + } + } + t.Fatal("cancellation did not publish canceled end event") +} + +func acknowledgePendingCheckpointBlobs(service *Service, stream *ActiveStream) error { + for { + stream.mu.Lock() + requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites)) + for requestID := range stream.PendingCheckpointBlobWrites { + requestIDs = append(requestIDs, requestID) + } + stream.mu.Unlock() + if len(requestIDs) == 0 { + return nil + } + for _, requestID := range requestIDs { + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + return err + } + } + } +} + +func rejectCheckpointBlob(service *Service, stream *ActiveStream, requestID uint32, message string) error { + return service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{Error: &agentv1.Error{Message: message}}, + }, + }) +} + +func TestImportedTurnIDsRemainCheckpointPrefix(t *testing.T) { + importedID := make([]byte, 32) + for index := range importedID { + importedID[index] = byte(index + 1) + } + conversation := testConversation([]HistoryEntry{ + testUserMessageEntry(t, 2, "request-2", "continued question"), + }) + conversation.ImportedTurnIDs = [][]byte{importedID} + projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + if len(projection.State.GetTurns()) != 2 { + t.Fatalf("turns = %d, want imported prefix plus projected turn", len(projection.State.GetTurns())) + } + if string(projection.State.GetTurns()[0]) != string(importedID) { + t.Fatal("imported turn ID was not preserved as the checkpoint prefix") + } +} + +func testCheckpointBlobService(t *testing.T) (*Service, *ActiveStream) { + t.Helper() + broker := NewStreamBroker() + service := &Service{ + projector: NewHistoryProjector(), + broker: broker, + checkpointBlobs: make(map[string]*checkpointBlobCacheEntry), + } + stream, err := broker.OpenStream( + "request-1", + "conversation-1", + 1, + "default", + "default", + agentv1.AgentMode_AGENT_MODE_AGENT, + "hello", + ) + if err != nil { + t.Fatalf("OpenStream() error = %v", err) + } + return service, stream +} diff --git a/internal/backend/forwarder/broker.go b/internal/backend/forwarder/broker.go index b88442c..3d704c3 100644 --- a/internal/backend/forwarder/broker.go +++ b/internal/backend/forwarder/broker.go @@ -82,28 +82,30 @@ func (broker *StreamBroker) OpenStream(requestID string, conversationID string, } now := time.Now().UTC() stream := &ActiveStream{ - RequestID: normalizedRequestID, - ConversationID: strings.TrimSpace(conversationID), - TurnSeq: turnSeq, - ModelID: strings.TrimSpace(modelID), - ModelName: strings.TrimSpace(modelName), - Mode: normalizedMode, - LatestUserText: strings.TrimSpace(latestUserText), - Status: StreamStatusCreated, - Backlog: make([]StreamEvent, 0, 64), - Subscribers: make(map[string]*StreamSubscriber), - PendingExecs: make(map[string]runtimecore.PendingExec), - PendingInteractions: make(map[string]runtimecore.PendingInteraction), - PartialToolCallIDs: make(map[string]struct{}), - PatchEditQueues: make(map[string][]queuedPatchEditOperation), - MCPToolServers: make(map[string]string), - RecentCompletedExecs: make(map[uint32]time.Time), - BackgroundShells: make(map[string]*BackgroundShellState), - BackgroundShellsByMessageID: make(map[uint32]string), - BackgroundShellsByExecID: make(map[string]string), - BackgroundShellActions: make(map[string]time.Time), - CreatedAt: now, - UpdatedAt: now, + RequestID: normalizedRequestID, + ConversationID: strings.TrimSpace(conversationID), + TurnSeq: turnSeq, + ModelID: strings.TrimSpace(modelID), + ModelName: strings.TrimSpace(modelName), + Mode: normalizedMode, + LatestUserText: strings.TrimSpace(latestUserText), + Status: StreamStatusCreated, + Backlog: make([]StreamEvent, 0, 64), + Subscribers: make(map[string]*StreamSubscriber), + PendingExecs: make(map[string]runtimecore.PendingExec), + PendingInteractions: make(map[string]runtimecore.PendingInteraction), + PartialToolCallIDs: make(map[string]struct{}), + PatchEditQueues: make(map[string][]queuedPatchEditOperation), + MCPToolServers: make(map[string]string), + RecentCompletedExecs: make(map[uint32]time.Time), + BackgroundShells: make(map[string]*BackgroundShellState), + BackgroundShellsByMessageID: make(map[uint32]string), + BackgroundShellsByExecID: make(map[string]string), + BackgroundShellActions: make(map[string]time.Time), + PendingCheckpointBlobWrites: make(map[uint32]pendingCheckpointBlobWrite), + PendingCheckpointBlobRequests: make(map[string]uint32), + CreatedAt: now, + UpdatedAt: now, } broker.streams[normalizedRequestID] = stream return stream, nil diff --git a/internal/backend/forwarder/events.go b/internal/backend/forwarder/events.go index 7b02404..ef69776 100644 --- a/internal/backend/forwarder/events.go +++ b/internal/backend/forwarder/events.go @@ -245,6 +245,22 @@ func buildCheckpointMessage(state *agentv1.ConversationStateStructure) *agentv1. } } +func buildSetCheckpointBlobMessage(id uint32, blob CheckpointBlob) *agentv1.AgentServerMessage { + return &agentv1.AgentServerMessage{ + Message: &agentv1.AgentServerMessage_KvServerMessage{ + KvServerMessage: &agentv1.KvServerMessage{ + Id: id, + Message: &agentv1.KvServerMessage_SetBlobArgs{ + SetBlobArgs: &agentv1.SetBlobArgs{ + BlobId: append([]byte(nil), blob.ID...), + BlobData: append([]byte(nil), blob.Data...), + }, + }, + }, + }, + } +} + // buildExecAbortMessage 构造对客户端执行桥的 abort 控制消息。 func buildExecAbortMessage(pending runtimecore.PendingExec) *agentv1.AgentServerMessage { return &agentv1.AgentServerMessage{ diff --git a/internal/backend/forwarder/file_store.go b/internal/backend/forwarder/file_store.go index 0732614..fb96d1b 100644 --- a/internal/backend/forwarder/file_store.go +++ b/internal/backend/forwarder/file_store.go @@ -750,6 +750,7 @@ func mergeConversationMetadata(target *ConversationFile, source *ConversationFil target.CurrentPlanText = source.CurrentPlanText target.CurrentPlans = clonePlanRegistryEntries(source.CurrentPlans) target.CurrentTodos = cloneTodoItems(source.CurrentTodos) + target.ImportedTurnIDs = cloneByteSlices(source.ImportedTurnIDs) target.LatestRequestPrefix = cloneConversationRequestPrefix(source.LatestRequestPrefix) target.LastProviderCall = cloneConversationProviderCall(source.LastProviderCall) if !source.CreatedAt.IsZero() && (target.CreatedAt.IsZero() || source.CreatedAt.Before(target.CreatedAt)) { @@ -882,6 +883,7 @@ func cloneConversationFile(conversation *ConversationFile) *ConversationFile { cloned := *conversation cloned.CurrentPlans = clonePlanRegistryEntries(conversation.CurrentPlans) cloned.CurrentTodos = cloneTodoItems(conversation.CurrentTodos) + cloned.ImportedTurnIDs = cloneByteSlices(conversation.ImportedTurnIDs) cloned.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix) cloned.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall) cloned.Entries = append([]HistoryEntry(nil), conversation.Entries...) diff --git a/internal/backend/forwarder/imported_blobs.go b/internal/backend/forwarder/imported_blobs.go new file mode 100644 index 0000000..fbd17ed --- /dev/null +++ b/internal/backend/forwarder/imported_blobs.go @@ -0,0 +1,167 @@ +package forwarder + +import ( + "crypto/sha256" + "fmt" + + "google.golang.org/protobuf/proto" + + "cursor/gen/agentv1" + modeladapter "cursor/internal/backend/agent/model" + promptengine "cursor/internal/backend/agent/prompt" +) + +type importedBlobStore map[string][]byte + +func newImportedBlobStore(items []*agentv1.PreFetchedBlob) (importedBlobStore, error) { + if len(items) == 0 { + return nil, nil + } + store := make(importedBlobStore, len(items)) + for _, item := range items { + if item == nil || len(item.GetId()) == 0 { + continue + } + if len(item.GetId()) != sha256.Size { + return nil, fmt.Errorf("prefetched blob id length %d, want %d", len(item.GetId()), sha256.Size) + } + digest := sha256.Sum256(item.GetValue()) + if string(digest[:]) != string(item.GetId()) { + return nil, fmt.Errorf("prefetched blob %x failed SHA-256 validation", item.GetId()) + } + store[string(item.GetId())] = append([]byte(nil), item.GetValue()...) + } + return store, nil +} + +func (store importedBlobStore) resolve(id []byte) ([]byte, bool) { + if len(id) == 0 || len(store) == 0 { + return nil, false + } + value, ok := store[string(id)] + return append([]byte(nil), value...), ok +} + +func decodeImportedTurn(raw []byte, blobs importedBlobStore) (*agentv1.ConversationTurnStructure, []byte, error) { + if data, ok := blobs.resolve(raw); ok { + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(data, turn); err != nil || turn.GetTurn() == nil { + return nil, nil, fmt.Errorf("decode imported turn blob %x: %w", raw, firstNonNilError(err, fmt.Errorf("turn payload is empty"))) + } + return turn, append([]byte(nil), raw...), nil + } + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(raw, turn); err == nil && turn.GetTurn() != nil { + return turn, nil, nil + } + if len(raw) == sha256.Size { + return nil, append([]byte(nil), raw...), nil + } + return nil, nil, fmt.Errorf("decode imported inline turn") +} + +func decodeImportedUserMessage(raw []byte, blobs importedBlobStore) (*agentv1.UserMessage, error) { + data := raw + if resolved, ok := blobs.resolve(raw); ok { + data = resolved + } else if len(raw) == sha256.Size { + candidate := &agentv1.UserMessage{} + if err := proto.Unmarshal(raw, candidate); err != nil || !hasKnownUserMessageContent(candidate) { + return nil, fmt.Errorf("missing prefetched user message blob %x", raw) + } + return candidate, nil + } + message := &agentv1.UserMessage{} + if err := proto.Unmarshal(data, message); err != nil { + return nil, fmt.Errorf("decode imported turn user_message: %w", err) + } + return message, nil +} + +func decodeImportedStep(raw []byte, blobs importedBlobStore) (*agentv1.ConversationStep, error) { + data := raw + if resolved, ok := blobs.resolve(raw); ok { + data = resolved + } else if len(raw) == sha256.Size { + candidate := &agentv1.ConversationStep{} + if err := proto.Unmarshal(raw, candidate); err != nil || candidate.GetMessage() == nil { + return nil, fmt.Errorf("missing prefetched conversation step blob %x", raw) + } + return candidate, nil + } + step := &agentv1.ConversationStep{} + if err := proto.Unmarshal(data, step); err != nil { + return nil, fmt.Errorf("decode imported turn step: %w", err) + } + if step.GetMessage() == nil { + return nil, fmt.Errorf("decode imported turn step: payload is empty") + } + return step, nil +} + +func importedBlobTurnMessages(turn *agentv1.ConversationTurnStructure, blobs importedBlobStore) ([]modeladapter.Message, error) { + if turn == nil || turn.GetAgentConversationTurn() == nil { + return nil, nil + } + agentTurn := turn.GetAgentConversationTurn() + messages := make([]modeladapter.Message, 0, 1+len(agentTurn.GetSteps())) + if len(agentTurn.GetUserMessage()) > 0 { + userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs) + if err != nil { + return nil, err + } + if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok { + messages = append(messages, toModelMessage(replay)) + } + } + for _, rawStep := range agentTurn.GetSteps() { + if len(rawStep) == 0 { + continue + } + step, err := decodeImportedStep(rawStep, blobs) + if err != nil { + return nil, err + } + for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) { + messages = append(messages, toModelMessage(replay)) + } + } + return messages, nil +} + +func importedTurnIDs(turns [][]byte, blobs importedBlobStore) ([][]byte, error) { + ids := make([][]byte, 0, len(turns)) + for _, raw := range turns { + if len(raw) == 0 { + continue + } + _, id, err := decodeImportedTurn(raw, blobs) + if err != nil { + return nil, err + } + if len(id) > 0 { + ids = append(ids, id) + } + } + return ids, nil +} + +func hasKnownUserMessageContent(message *agentv1.UserMessage) bool { + if message == nil { + return false + } + return message.GetText() != "" || + message.GetMessageId() != "" || + message.GetSelectedContext() != nil || + message.GetRichText() != "" || + len(message.GetConversationStateBlobId()) > 0 || + len(message.GetTextBlobId()) > 0 || + len(message.GetRichTextBlobId()) > 0 +} + +func firstNonNilError(err error, fallback error) error { + if err != nil { + return err + } + return fallback +} diff --git a/internal/backend/forwarder/projector.go b/internal/backend/forwarder/projector.go index 5726f8f..3f612d2 100644 --- a/internal/backend/forwarder/projector.go +++ b/internal/backend/forwarder/projector.go @@ -2,6 +2,7 @@ package forwarder import ( + "crypto/sha256" "encoding/json" "fmt" "strings" @@ -19,6 +20,51 @@ const projectedConversationMaxTokens = 130000 type HistoryProjector struct { } +type CheckpointBlob struct { + ID []byte + Data []byte +} + +type CheckpointProjection struct { + State *agentv1.ConversationStateStructure + Blobs []CheckpointBlob +} + +type checkpointBlobGraph struct { + blobs map[[sha256.Size]byte][]byte + order [][sha256.Size]byte +} + +func newCheckpointBlobGraph() *checkpointBlobGraph { + return &checkpointBlobGraph{blobs: make(map[[sha256.Size]byte][]byte)} +} + +func (graph *checkpointBlobGraph) add(data []byte) []byte { + if graph == nil || len(data) == 0 { + return nil + } + id := sha256.Sum256(data) + if _, exists := graph.blobs[id]; !exists { + graph.blobs[id] = append([]byte(nil), data...) + graph.order = append(graph.order, id) + } + return append([]byte(nil), id[:]...) +} + +func (graph *checkpointBlobGraph) list() []CheckpointBlob { + if graph == nil || len(graph.order) == 0 { + return nil + } + blobs := make([]CheckpointBlob, 0, len(graph.order)) + for _, id := range graph.order { + blobs = append(blobs, CheckpointBlob{ + ID: append([]byte(nil), id[:]...), + Data: append([]byte(nil), graph.blobs[id]...), + }) + } + return blobs +} + // NewHistoryProjector 创建 history 投影器。 func NewHistoryProjector() *HistoryProjector { return &HistoryProjector{} @@ -475,6 +521,16 @@ func isHistoricalReplayToolResult(conversation *ConversationFile, entry HistoryE // ProjectLegacyCheckpoint 按需从 JSON history 投影出兼容旧客户端的 checkpoint 结构。 func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *ConversationFile) (*agentv1.ConversationStateStructure, error) { + projection, err := projector.ProjectCheckpointProjection(conversation) + if err != nil || projection == nil { + return nil, err + } + return projection.State, nil +} + +// ProjectCheckpointProjection 同时返回 checkpoint 状态及其引用的内容寻址 Blob。 +func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *ConversationFile) (*CheckpointProjection, error) { + blobs := newCheckpointBlobGraph() state := &agentv1.ConversationStateStructure{ TokenDetails: &agentv1.ConversationTokenDetails{ UsedTokens: conversationTokenDetailsUsedTokens(conversation), @@ -488,7 +544,7 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers if conversation == nil { mode := agentv1.AgentMode_AGENT_MODE_AGENT state.Mode = &mode - return state, nil + return &CheckpointProjection{State: state}, nil } mode, err := parseModeAlias(conversation.Mode) if err != nil { @@ -506,158 +562,11 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers if structuredState.HasTodos { state.Todos = encodeConversationTodoBytes(structuredState.Todos) } - grouped := make(map[int64][]HistoryEntry) - order := make([]int64, 0, conversation.NextTurnSeq) - for _, entry := range checkpointProjectionEntries(conversation.Entries) { - if entry.TurnSeq <= 0 { - continue - } - if _, ok := grouped[entry.TurnSeq]; !ok { - order = append(order, entry.TurnSeq) - } - grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry) - } - - for _, turnSeq := range order { - entries := grouped[turnSeq] - var rawUserMessage []byte - var turnRequestID string - steps := make([][]byte, 0, len(entries)) - seenToolCalls := make(map[string]struct{}) - openToolCalls := make(map[string]struct{}) - for _, entry := range entries { - if turnRequestID == "" { - turnRequestID = strings.TrimSpace(entry.RequestID) - } - switch strings.TrimSpace(entry.Kind) { - case "user_message": - userMessage := &agentv1.UserMessage{} - if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil { - return nil, fmt.Errorf("decode checkpoint user_message: %w", err) - } - payload, err := proto.Marshal(userMessage) - if err != nil { - return nil, err - } - rawUserMessage = payload - case "assistant_text": - var payload assistantTextPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if strings.TrimSpace(payload.Text) == "" && strings.TrimSpace(payload.ReasoningContent) != "" && len(openToolCalls) > 0 { - continue - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - if strings.TrimSpace(payload.Text) == "" { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_AssistantMessage{ - AssistantMessage: &agentv1.AssistantMessage{Text: strings.TrimSpace(payload.Text)}, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - case "tool_call": - var payload toolCallEntryPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - toolCall := &agentv1.ToolCall{} - if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { - return nil, err - } - if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ToolCall{ - ToolCall: toolCall, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { - seenToolCalls[toolCallID] = struct{}{} - openToolCalls[toolCallID] = struct{}{} - } - case "tool_result": - var payload toolResultEntryPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { - if _, ok := seenToolCalls[toolCallID]; ok { - delete(openToolCalls, toolCallID) - continue - } - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - if len(payload.ToolCall) == 0 { - continue - } - toolCall := &agentv1.ToolCall{} - if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { - return nil, err - } - if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ToolCall{ - ToolCall: toolCall, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - } - if len(rawUserMessage) == 0 && len(steps) == 0 { - continue - } - agentTurn := &agentv1.AgentConversationTurnStructure{ - UserMessage: rawUserMessage, - Steps: steps, - } - if turnRequestID != "" { - agentTurn.RequestId = &turnRequestID - } - turnPayload, err := proto.Marshal(&agentv1.ConversationTurnStructure{ - Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ - AgentConversationTurn: agentTurn, - }, - }) - if err != nil { - return nil, err - } - state.Turns = append(state.Turns, turnPayload) + turnIDs, err := projectCheckpointTurnBlobs(conversation, blobs) + if err != nil { + return nil, err } + state.Turns = append(cloneByteSlices(conversation.ImportedTurnIDs), turnIDs...) replayMessages, err := projector.ProjectPromptReplay(conversation) if err != nil { return nil, err @@ -685,15 +594,183 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers return nil, err } state.RootPromptMessagesJson = rootPromptMessages - return state, nil + return &CheckpointProjection{State: state, Blobs: blobs.list()}, nil } -func marshalThinkingStep(text string) ([]byte, error) { - return proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ThinkingMessage{ - ThinkingMessage: &agentv1.ThinkingMessage{Text: text}, - }, - }) +func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpointBlobGraph) ([][]byte, error) { + if conversation == nil || blobs == nil { + return nil, nil + } + grouped := make(map[int64][]HistoryEntry) + order := make([]int64, 0, conversation.NextTurnSeq) + for _, entry := range checkpointProjectionEntries(conversation.Entries) { + if entry.TurnSeq <= 0 { + continue + } + if _, ok := grouped[entry.TurnSeq]; !ok { + order = append(order, entry.TurnSeq) + } + grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry) + } + + turnIDs := make([][]byte, 0, len(order)) + for _, turnSeq := range order { + entries := grouped[turnSeq] + var userMessageID []byte + var turnRequestID string + stepIDs := make([][]byte, 0, len(entries)) + seenToolCalls := make(map[string]struct{}) + openToolCalls := make(map[string]struct{}) + for _, entry := range entries { + if turnRequestID == "" { + turnRequestID = strings.TrimSpace(entry.RequestID) + } + switch strings.TrimSpace(entry.Kind) { + case "user_message": + userMessage := &agentv1.UserMessage{} + if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil { + return nil, fmt.Errorf("decode checkpoint user_message: %w", err) + } + payload, err := proto.Marshal(userMessage) + if err != nil { + return nil, err + } + userMessageID = blobs.add(payload) + case "assistant_text": + var payload assistantTextPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if strings.TrimSpace(payload.Text) == "" && strings.TrimSpace(payload.ReasoningContent) != "" && len(openToolCalls) > 0 { + continue + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + if strings.TrimSpace(payload.Text) == "" { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_AssistantMessage{ + AssistantMessage: &agentv1.AssistantMessage{Text: strings.TrimSpace(payload.Text)}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + case "tool_call": + var payload toolCallEntryPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + toolCall := &agentv1.ToolCall{} + if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { + return nil, err + } + if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ToolCall{ToolCall: toolCall}, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { + seenToolCalls[toolCallID] = struct{}{} + openToolCalls[toolCallID] = struct{}{} + } + case "tool_result": + var payload toolResultEntryPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { + if _, ok := seenToolCalls[toolCallID]; ok { + delete(openToolCalls, toolCallID) + continue + } + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + if len(payload.ToolCall) == 0 { + continue + } + toolCall := &agentv1.ToolCall{} + if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { + return nil, err + } + if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ToolCall{ToolCall: toolCall}, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + } + if len(userMessageID) == 0 && len(stepIDs) == 0 { + continue + } + agentTurn := &agentv1.AgentConversationTurnStructure{ + UserMessage: userMessageID, + Steps: stepIDs, + } + if turnRequestID != "" { + agentTurn.RequestId = &turnRequestID + } + turnPayload, err := proto.Marshal(&agentv1.ConversationTurnStructure{ + Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ + AgentConversationTurn: agentTurn, + }, + }) + if err != nil { + return nil, err + } + turnIDs = append(turnIDs, blobs.add(turnPayload)) + } + return turnIDs, nil +} + +func addCheckpointStepBlob(blobs *checkpointBlobGraph, step *agentv1.ConversationStep) ([]byte, error) { + payload, err := proto.Marshal(step) + if err != nil { + return nil, err + } + return blobs.add(payload), nil } func conversationTokenDetailsUsedTokens(conversation *ConversationFile) uint32 { @@ -1141,60 +1218,6 @@ func shouldPersistToolResultName(toolName string) bool { } } -func filterCheckpointTurns(rawTurns [][]byte) [][]byte { - if len(rawTurns) == 0 { - return nil - } - filtered := make([][]byte, 0, len(rawTurns)) - for _, rawTurn := range rawTurns { - if len(rawTurn) == 0 { - continue - } - turn := &agentv1.ConversationTurnStructure{} - if err := proto.Unmarshal(rawTurn, turn); err != nil { - filtered = append(filtered, append([]byte(nil), rawTurn...)) - continue - } - agentTurn := turn.GetAgentConversationTurn() - if agentTurn == nil { - filtered = append(filtered, append([]byte(nil), rawTurn...)) - continue - } - - nextSteps := make([][]byte, 0, len(agentTurn.GetSteps())) - for _, rawStep := range agentTurn.GetSteps() { - if len(rawStep) == 0 { - continue - } - step := &agentv1.ConversationStep{} - if err := proto.Unmarshal(rawStep, step); err != nil { - continue - } - if toolCall := step.GetToolCall(); toolCall != nil && !shouldPersistToolResultName(inferToolName(toolCall)) { - continue - } - nextSteps = append(nextSteps, append([]byte(nil), rawStep...)) - } - if len(agentTurn.GetUserMessage()) == 0 && len(nextSteps) == 0 { - continue - } - encoded, err := proto.Marshal(&agentv1.ConversationTurnStructure{ - Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ - AgentConversationTurn: &agentv1.AgentConversationTurnStructure{ - UserMessage: append([]byte(nil), agentTurn.GetUserMessage()...), - Steps: nextSteps, - }, - }, - }) - if err != nil { - filtered = append(filtered, append([]byte(nil), rawTurn...)) - continue - } - filtered = append(filtered, encoded) - } - return filtered -} - func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []promptengine.Message { if len(messages) == 0 { return nil @@ -1231,7 +1254,7 @@ func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []pro return filtered } -func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte) []promptengine.Message { +func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte, blobs importedBlobStore) []promptengine.Message { if len(messages) == 0 || len(importedTurns) == 0 { return messages } @@ -1240,16 +1263,16 @@ func restoreImportedReplayUserMessages(messages []promptengine.Message, imported if len(rawTurn) == 0 { continue } - turn := &agentv1.ConversationTurnStructure{} - if err := proto.Unmarshal(rawTurn, turn); err != nil { + turn, _, err := decodeImportedTurn(rawTurn, blobs) + if err != nil || turn == nil { continue } agentTurn := turn.GetAgentConversationTurn() if agentTurn == nil || len(agentTurn.GetUserMessage()) == 0 { continue } - userMessage := &agentv1.UserMessage{} - if err := proto.Unmarshal(agentTurn.GetUserMessage(), userMessage); err != nil { + userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs) + if err != nil { continue } replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage) diff --git a/internal/backend/forwarder/projector_test.go b/internal/backend/forwarder/projector_test.go new file mode 100644 index 0000000..10330ba --- /dev/null +++ b/internal/backend/forwarder/projector_test.go @@ -0,0 +1,433 @@ +package forwarder + +import ( + "crypto/sha256" + "encoding/json" + "fmt" + "strings" + "testing" + + "google.golang.org/protobuf/encoding/protojson" + "google.golang.org/protobuf/proto" + + "cursor/gen/agentv1" + modeladapter "cursor/internal/backend/agent/model" + promptengine "cursor/internal/backend/agent/prompt" +) + +func TestProjectCheckpointProjectionBuildsBlobBackedTurns(t *testing.T) { + toolCall := testEditToolCall(t, "file.txt") + + tests := []struct { + name string + entries []HistoryEntry + }{ + { + name: "no tools", + entries: []HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "hello"), + newAssistantTextEntry(1, "request-1", "hi", "", ""), + }, + }, + { + name: "completed tool call", + entries: []HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "edit the file"), + newToolCallEntry(1, "request-1", "call-1", "Edit", "", "", toolCall), + newToolResultEntry(1, "request-1", "call-1", "Edit", `{"path":"file.txt"}`, "edited", "", toolCall), + newAssistantTextEntry(1, "request-1", "done", "", ""), + }, + }, + { + name: "unfinished tool call", + entries: []HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "edit the file"), + newToolCallEntry(1, "request-1", "call-1", "Edit", "", "", toolCall), + }, + }, + { + name: "orphan tool result", + entries: []HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "edit the file"), + newToolResultEntry(1, "request-1", "call-1", "Edit", `{"path":"file.txt"}`, "edited", "", toolCall), + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + conversation := testConversation(test.entries) + projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + if len(projection.State.GetTurns()) != 1 { + t.Fatalf("ProjectCheckpointProjection() turns = %d, want 1 Blob ID", len(projection.State.GetTurns())) + } + assertCheckpointBlobGraph(t, projection) + messages, err := promptengine.DecodeReplayMessages(projection.State.GetRootPromptMessagesJson()) + if err != nil { + t.Fatalf("DecodeReplayMessages() error = %v", err) + } + if len(messages) == 0 { + t.Fatal("ProjectCheckpointProjection() removed all root prompt replay history") + } + if messages[0].Role != "user" || messages[0].Content == "" { + t.Fatalf("first replay message = %#v, want retained user history", messages[0]) + } + }) + } +} + +func TestProjectLegacyCheckpointLargeModelHistoryUsesRootReplay(t *testing.T) { + entries := make([]HistoryEntry, 0, 400) + for turn := int64(1); turn <= 200; turn++ { + requestID := fmt.Sprintf("request-%d", turn) + entries = append(entries, + testModelMessageEntry(t, turn, requestID, modeladapter.Message{Role: "user", Content: fmt.Sprintf("question %d", turn)}), + testModelMessageEntry(t, turn, requestID, modeladapter.Message{Role: "assistant", Content: fmt.Sprintf("answer %d", turn)}), + ) + } + + state, err := NewHistoryProjector().ProjectLegacyCheckpoint(testConversation(entries)) + if err != nil { + t.Fatalf("ProjectLegacyCheckpoint() error = %v", err) + } + if len(state.GetTurns()) != 0 { + t.Fatalf("ProjectLegacyCheckpoint() model-only turns = %d, want 0", len(state.GetTurns())) + } + messages, err := promptengine.DecodeReplayMessages(state.GetRootPromptMessagesJson()) + if err != nil { + t.Fatalf("DecodeReplayMessages() error = %v", err) + } + if len(messages) != 400 { + t.Fatalf("decoded replay messages = %d, want 400", len(messages)) + } +} + +func TestProjectLegacyCheckpointSnapshotIsIsolatedFromLaterHistory(t *testing.T) { + conversation := testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "first question"), + newAssistantTextEntry(1, "request-1", "first answer", "", ""), + }) + projector := NewHistoryProjector() + midpointProjection, err := projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("midpoint ProjectCheckpointProjection() error = %v", err) + } + midpoint := midpointProjection.State + + appendEntriesInPlace(conversation, []HistoryEntry{ + testUserMessageEntry(t, 2, "request-2", "second question"), + newAssistantTextEntry(2, "request-2", "second answer", "", ""), + }) + latestProjection, err := projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("latest ProjectCheckpointProjection() error = %v", err) + } + latest := latestProjection.State + + midpointMessages, err := promptengine.DecodeReplayMessages(midpoint.GetRootPromptMessagesJson()) + if err != nil { + t.Fatalf("decode midpoint replay: %v", err) + } + latestMessages, err := promptengine.DecodeReplayMessages(latest.GetRootPromptMessagesJson()) + if err != nil { + t.Fatalf("decode latest replay: %v", err) + } + if len(midpointMessages) != 2 { + t.Fatalf("midpoint replay messages = %d, want 2", len(midpointMessages)) + } + if len(latestMessages) != 4 { + t.Fatalf("latest replay messages = %d, want 4", len(latestMessages)) + } + if len(midpoint.GetTurns()) != 1 || len(latest.GetTurns()) != 2 { + t.Fatalf("checkpoint turn counts = (%d, %d), want (1, 2)", len(midpoint.GetTurns()), len(latest.GetTurns())) + } + assertCheckpointBlobGraph(t, midpointProjection) + assertCheckpointBlobGraph(t, latestProjection) +} + +func TestProjectCheckpointProjectionKeepsVisibleTurnsAcrossCompaction(t *testing.T) { + summaryPayload, err := json.Marshal(compactionSummaryEntryPayload{Summary: "first turn summarized"}) + if err != nil { + t.Fatalf("marshal compaction summary: %v", err) + } + conversation := testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "first question"), + newAssistantTextEntry(1, "request-1", "first answer", "", ""), + {TurnSeq: 0, Role: "system", Kind: "compaction_summary", Payload: summaryPayload}, + testUserMessageEntry(t, 2, "request-2", "second question"), + newAssistantTextEntry(2, "request-2", "second answer", "", ""), + }) + + projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + if len(projection.State.GetTurns()) != 2 { + t.Fatalf("visible turns after compaction = %d, want 2", len(projection.State.GetTurns())) + } + assertCheckpointBlobGraph(t, projection) + + messages, err := promptengine.DecodeReplayMessages(projection.State.GetRootPromptMessagesJson()) + if err != nil { + t.Fatalf("DecodeReplayMessages() error = %v", err) + } + if len(messages) != 3 { + t.Fatalf("compacted root replay messages = %d, want summary plus latest turn", len(messages)) + } + if messages[0].Role != "user" || messages[0].Content != "\nfirst turn summarized\n" { + t.Fatalf("first compacted replay message = %#v", messages[0]) + } +} + +func TestImportedConversationStateRejectsBlobTurnIDsWithoutPrefetchedData(t *testing.T) { + turnID := sha256.Sum256([]byte("imported turn")) + state := &agentv1.ConversationStateStructure{Turns: [][]byte{turnID[:]}} + if _, err := importedConversationStateModelMessages(state, nil); err == nil { + t.Fatal("importedConversationStateModelMessages() accepted unresolved Blob turn") + } + conversation := testConversation(nil) + service := &Service{} + if _, err := service.importConversationState(conversation, state, nil); err == nil { + t.Fatal("importConversationState() accepted unresolved Blob turn") + } +} + +func TestImportedConversationStateRestoresBlobOnlyForkFromPrefetchedBlobs(t *testing.T) { + projection, err := NewHistoryProjector().ProjectCheckpointProjection(testConversation([]HistoryEntry{ + testUserMessageEntry(t, 1, "request-1", "parent question"), + newAssistantTextEntry(1, "request-1", "parent answer", "", ""), + })) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + prefetched := make([]*agentv1.PreFetchedBlob, 0, len(projection.Blobs)) + for _, blob := range projection.Blobs { + prefetched = append(prefetched, &agentv1.PreFetchedBlob{Id: blob.ID, Value: blob.Data}) + } + state := proto.Clone(projection.State).(*agentv1.ConversationStateStructure) + state.RootPromptMessagesJson = nil + conversation := testConversation(nil) + entries, err := (&Service{}).importConversationState(conversation, state, prefetched) + if err != nil { + t.Fatalf("importConversationState() error = %v", err) + } + if len(conversation.ImportedTurnIDs) != 1 { + t.Fatalf("ImportedTurnIDs = %d, want 1", len(conversation.ImportedTurnIDs)) + } + if len(entries) != 2 { + t.Fatalf("imported model entries = %d, want user and assistant", len(entries)) + } +} + +func TestImportedInlineTurnWithSHA256LengthIsNotMisclassified(t *testing.T) { + var rawTurn []byte + for size := 1; size <= 128; size++ { + rawUser, err := proto.Marshal(&agentv1.UserMessage{Text: strings.Repeat("x", size), MessageId: "inline"}) + if err != nil { + t.Fatalf("marshal user message: %v", err) + } + rawTurn, err = proto.Marshal(&agentv1.ConversationTurnStructure{ + Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ + AgentConversationTurn: &agentv1.AgentConversationTurnStructure{UserMessage: rawUser}, + }, + }) + if err != nil { + t.Fatalf("marshal turn: %v", err) + } + if len(rawTurn) == sha256.Size { + break + } + } + if len(rawTurn) != sha256.Size { + t.Fatal("test could not construct a 32-byte inline turn") + } + ids, err := importedTurnIDs([][]byte{rawTurn}, nil) + if err != nil { + t.Fatalf("importedTurnIDs() error = %v", err) + } + if len(ids) != 0 { + t.Fatal("32-byte inline turn was misclassified as a Blob ID") + } + messages, err := importedConversationStateModelMessages(&agentv1.ConversationStateStructure{Turns: [][]byte{rawTurn}}, nil) + if err != nil { + t.Fatalf("importedConversationStateModelMessages() error = %v", err) + } + if len(messages) != 1 || messages[0].Role != "user" { + t.Fatalf("inline turn messages = %#v, want one user message", messages) + } +} + +func TestImportedTurnIDsPersistThroughConversationStore(t *testing.T) { + store := NewConversationFileStore(t.TempDir()) + turnID := sha256.Sum256([]byte("parent turn")) + conversation := testConversation(nil) + conversation.ImportedTurnIDs = [][]byte{turnID[:]} + persisted, err := store.SaveConversationWithEntries(conversation.ConversationID, conversation, []HistoryEntry{ + testUserMessageEntry(t, 2, "request-2", "fork question"), + }) + if err != nil { + t.Fatalf("SaveConversationWithEntries() error = %v", err) + } + if len(persisted.ImportedTurnIDs) != 1 || string(persisted.ImportedTurnIDs[0]) != string(turnID[:]) { + t.Fatalf("persisted ImportedTurnIDs = %x, want %x", persisted.ImportedTurnIDs, turnID) + } + loaded, err := store.LoadConversation(conversation.ConversationID) + if err != nil { + t.Fatalf("LoadConversation() error = %v", err) + } + if len(loaded.ImportedTurnIDs) != 1 || string(loaded.ImportedTurnIDs[0]) != string(turnID[:]) { + t.Fatalf("loaded ImportedTurnIDs = %x, want %x", loaded.ImportedTurnIDs, turnID) + } +} + +func TestRewindImportedTurnPrefixUsesClientForkPoint(t *testing.T) { + ids := make([][]byte, 4) + for index := range ids { + digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1))) + ids[index] = digest[:] + } + trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{ + TargetTurnSeq: 4, + HasClientTurnCount: true, + ClientTurnCount: 2, + }) + if len(trimmed) != 2 || string(trimmed[0]) != string(ids[0]) || string(trimmed[1]) != string(ids[1]) { + t.Fatalf("rewindImportedTurnPrefix() = %x, want first two IDs", trimmed) + } +} + +func TestRewindImportedTurnPrefixClearsAllIDsAtClientTurnZero(t *testing.T) { + ids := make([][]byte, 2) + for index := range ids { + digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1))) + ids[index] = digest[:] + } + trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{ + TargetTurnSeq: 3, + HasClientTurnCount: true, + ClientTurnCount: 0, + }) + if trimmed != nil { + t.Fatalf("rewindImportedTurnPrefix() = %x, want nil at client turn zero", trimmed) + } +} + +func TestRewindImportedTurnPrefixUsesTargetWithoutClientCount(t *testing.T) { + ids := make([][]byte, 4) + for index := range ids { + digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1))) + ids[index] = digest[:] + } + trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{TargetTurnSeq: 3}) + if len(trimmed) != 2 || string(trimmed[0]) != string(ids[0]) || string(trimmed[1]) != string(ids[1]) { + t.Fatalf("rewindImportedTurnPrefix() = %x, want target-derived first two IDs", trimmed) + } +} + +func assertCheckpointBlobGraph(t *testing.T, projection *CheckpointProjection) { + t.Helper() + if projection == nil || projection.State == nil { + t.Fatal("checkpoint projection is nil") + } + blobByID := make(map[string][]byte, len(projection.Blobs)) + for _, blob := range projection.Blobs { + if len(blob.ID) != sha256.Size { + t.Fatalf("blob id length = %d, want %d", len(blob.ID), sha256.Size) + } + digest := sha256.Sum256(blob.Data) + if string(blob.ID) != string(digest[:]) { + t.Fatal("blob id does not match SHA-256(data)") + } + blobByID[string(blob.ID)] = blob.Data + } + for _, turnID := range projection.State.GetTurns() { + turnData, ok := blobByID[string(turnID)] + if !ok { + t.Fatal("turn references missing blob") + } + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(turnData, turn); err != nil { + t.Fatalf("decode turn blob: %v", err) + } + agentTurn := turn.GetAgentConversationTurn() + if agentTurn == nil { + continue + } + if userID := agentTurn.GetUserMessage(); len(userID) > 0 { + userData, exists := blobByID[string(userID)] + if !exists { + t.Fatal("turn references missing user message blob") + } + if err := proto.Unmarshal(userData, &agentv1.UserMessage{}); err != nil { + t.Fatalf("decode user message blob: %v", err) + } + } + for _, stepID := range agentTurn.GetSteps() { + stepData, exists := blobByID[string(stepID)] + if !exists { + t.Fatal("turn references missing step blob") + } + if err := proto.Unmarshal(stepData, &agentv1.ConversationStep{}); err != nil { + t.Fatalf("decode conversation step blob: %v", err) + } + } + } +} + +func testConversation(entries []HistoryEntry) *ConversationFile { + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 1, + NextEntrySeq: 1, + Entries: make([]HistoryEntry, 0, len(entries)), + } + appendEntriesInPlace(conversation, entries) + return conversation +} + +func testUserMessageEntry(t *testing.T, turnSeq int64, requestID string, text string) HistoryEntry { + t.Helper() + payload, err := protojson.Marshal(&agentv1.UserMessage{Text: text, MessageId: fmt.Sprintf("message-%d", turnSeq)}) + if err != nil { + t.Fatalf("marshal user message: %v", err) + } + return HistoryEntry{ + TurnSeq: turnSeq, + RequestID: requestID, + Role: "user", + Kind: "user_message", + Payload: payload, + } +} + +func testModelMessageEntry(t *testing.T, turnSeq int64, requestID string, message modeladapter.Message) HistoryEntry { + t.Helper() + entry, ok, err := newModelMessageEntry(turnSeq, requestID, message) + if err != nil { + t.Fatalf("newModelMessageEntry() error = %v", err) + } + if !ok { + t.Fatal("newModelMessageEntry() rejected test message") + } + return entry +} + +func testEditToolCall(t *testing.T, path string) []byte { + t.Helper() + payload, err := protojson.Marshal(&agentv1.ToolCall{ + Tool: &agentv1.ToolCall_EditToolCall{ + EditToolCall: &agentv1.EditToolCall{ + Args: &agentv1.EditArgs{Path: path}, + }, + }, + }) + if err != nil { + t.Fatalf("marshal edit tool call: %v", err) + } + return payload +} diff --git a/internal/backend/forwarder/rewind.go b/internal/backend/forwarder/rewind.go index ba7c851..1f11ee0 100644 --- a/internal/backend/forwarder/rewind.go +++ b/internal/backend/forwarder/rewind.go @@ -36,7 +36,7 @@ type runRewindMatch struct { } func (service *Service) decideRunRewind(intent InboundIntent, conversation *ConversationFile) runRewindDecision { - decision := runRewindDecision{ClientTurnCount: -1} + decision := runRewindDecision{} if !shouldEvaluateRunRewind(intent) { return decision } @@ -121,7 +121,7 @@ func selectRunRewindMatch(matches []runRewindMatch, clientTurnCount int, hasClie if len(matches) == 0 { return runRewindMatch{}, "no_match" } - if hasClientTurnCount && clientTurnCount >= 0 { + if hasClientTurnCount { targetTurnSeq := int64(clientTurnCount) + 1 for _, match := range matches { if match.Entry.TurnSeq == targetTurnSeq { @@ -224,6 +224,7 @@ func (service *Service) applyRunRewindToConversation(conversation *ConversationF conversation.Entries = nil conversation.NextEntrySeq = 1 conversation.NextTurnSeq = 1 + conversation.ImportedTurnIDs = rewindImportedTurnPrefix(conversation.ImportedTurnIDs, decision) appendEntriesInPlace(conversation, appendReplacementRunEntries(decision.PrefixEntries, entries)) applyRunRewindConversationState(conversation, intent, turnSeq) deriveConversationLoopState(conversation) @@ -269,10 +270,30 @@ func applyRunRewindMetadata(conversation *ConversationFile, source *Conversation if source.TokenDetailsMaxTokens > 0 { conversation.TokenDetailsMaxTokens = source.TokenDetailsMaxTokens } + decision := runRewindDecision{TargetTurnSeq: turnSeq} + if intent.ConversationState != nil { + decision.HasClientTurnCount = true + decision.ClientTurnCount = len(intent.ConversationState.GetTurns()) + } + conversation.ImportedTurnIDs = rewindImportedTurnPrefix(source.ImportedTurnIDs, decision) } applyRunRewindConversationState(conversation, intent, turnSeq) } +func rewindImportedTurnPrefix(importedTurnIDs [][]byte, decision runRewindDecision) [][]byte { + keep := decision.TargetTurnSeq - 1 + if decision.HasClientTurnCount { + keep = int64(decision.ClientTurnCount) + } + if keep <= 0 || len(importedTurnIDs) == 0 { + return nil + } + if keep > int64(len(importedTurnIDs)) { + keep = int64(len(importedTurnIDs)) + } + return cloneByteSlices(importedTurnIDs[:keep]) +} + func (service *Service) logRunRewindDecision(requestID string, conversationID string, eventName string, decision runRewindDecision) { if service == nil || !decision.Evaluated { return diff --git a/internal/backend/forwarder/runtime_summary.go b/internal/backend/forwarder/runtime_summary.go index 9abbba3..b302c8a 100644 --- a/internal/backend/forwarder/runtime_summary.go +++ b/internal/backend/forwarder/runtime_summary.go @@ -50,7 +50,7 @@ func (service *Service) bootstrapRuntimeConversation(intent InboundIntent) (*Con } importedEntries := []HistoryEntry(nil) if len(conversation.Entries) == 0 && intent.ConversationState != nil { - importedEntries, err = service.importConversationState(conversation, intent.ConversationState) + importedEntries, err = service.importConversationState(conversation, intent.ConversationState, intent.PreFetchedBlobs) if err != nil { return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err } @@ -138,6 +138,7 @@ func (service *Service) syncConversationRecord(conversationID string, conversati item.AutoCompactionReserveTokens = conversation.AutoCompactionReserveTokens item.AutoCompactionTriggeredAt = conversation.AutoCompactionTriggeredAt item.AutoCompactionSourceModelCallID = conversation.AutoCompactionSourceModelCallID + item.ImportedTurnIDs = cloneByteSlices(conversation.ImportedTurnIDs) item.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix) item.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall) item.CreatedAt = conversation.CreatedAt diff --git a/internal/backend/forwarder/service.go b/internal/backend/forwarder/service.go index 5bb04ca..fb4c848 100644 --- a/internal/backend/forwarder/service.go +++ b/internal/backend/forwarder/service.go @@ -9,6 +9,7 @@ import ( "log" "sort" "strings" + "sync" "time" "connectrpc.com/connect" @@ -261,6 +262,8 @@ type Service struct { execBridge execbridge.ExecBridge interactionBridge interactionbridge.InteractionBridge appendSeq *appendSequenceTracker + checkpointBlobMu sync.Mutex + checkpointBlobs map[string]*checkpointBlobCacheEntry } type agentModelMemory interface { @@ -300,6 +303,7 @@ func NewService(historyRoot string, resolver modeladapter.ChannelResolver) *Serv execBridge: execbridge.NewBridge(), interactionBridge: interactionbridge.NewBridge(), appendSeq: newAppendSequenceTracker(), + checkpointBlobs: make(map[string]*checkpointBlobCacheEntry), } service.startHistoryMaintenance() store.SyncAllCursorTranscriptsBestEffort() @@ -328,6 +332,7 @@ func newServiceWithDependencies(store *ConversationFileStore, projector *History execBridge: execbridge.NewBridge(), interactionBridge: interactionbridge.NewBridge(), appendSeq: newAppendSequenceTracker(), + checkpointBlobs: make(map[string]*checkpointBlobCacheEntry), } } @@ -557,6 +562,7 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A } intent.ConversationID = conversationID intent.ConversationState = runRequest.GetConversationState() + intent.PreFetchedBlobs = runRequest.GetPreFetchedBlobs() intent.UserMessage = extractUserMessage(message) intent.RequestContext = extractRequestContext(message) if service.shouldIgnoreEmptyResumeRunRequest(requestID, runRequest, intent.UserMessage, intent.RequestContext) { @@ -604,6 +610,7 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A intent.ConversationID = conversationID intent.SubagentTypeName = strings.TrimSpace(prewarmRequest.GetSubagentTypeName()) intent.ConversationState = prewarmRequest.GetConversationState() + intent.PreFetchedBlobs = prewarmRequest.GetPreFetchedBlobs() intent.Mode, intent.ModeSource, intent.HasExplicitMode, err = extractPrewarmMode(prewarmRequest) if err != nil { return InboundIntent{}, err @@ -819,11 +826,14 @@ func (service *Service) snapshotVisibleTurns(conversation *ConversationFile) ([] if service == nil || service.projector == nil || conversation == nil { return nil, nil } - state, err := service.projector.ProjectLegacyCheckpoint(conversation) + projection, err := service.projector.ProjectCheckpointProjection(conversation) if err != nil { return nil, err } - return cloneByteSlices(state.GetTurns()), nil + if projection == nil || projection.State == nil { + return nil, fmt.Errorf("checkpoint projection is empty") + } + return cloneByteSlices(projection.State.GetTurns()), nil } // handleCancelIntent 处理取消请求,并向客户端发送执行桥 abort。 @@ -833,42 +843,60 @@ func (service *Service) handleCancelIntent(intent InboundIntent) error { return fmt.Errorf("request is not active: %s", intent.RequestID) } hasCheckpoint := checkpointConversationInitialized(stream) - if hasCheckpoint { - cancelReason := firstNonEmpty(intent.CancelReason, "user aborted") - _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{ - newMetadataEntry(stream.TurnSeq, intent.RequestID, "control", map[string]any{ - "status": "canceled", - "reason": cancelReason, - "replay_policy": cancelReplayPolicyForReason(cancelReason), - }), - }) - if err != nil { - return err - } - } stream.mu.Lock() pendingExecs := make([]runtimecore.PendingExec, 0, len(stream.PendingExecs)) for _, pending := range stream.PendingExecs { pendingExecs = append(pendingExecs, pending) } + if stream.ProviderCancel != nil { + stream.ProviderCancel() + stream.ProviderCancel = nil + } + stream.ProviderActive = false + stream.CurrentProviderToken++ + stream.CurrentCompactionToken++ + stream.PendingProviderAction = providerActionNone + stream.PendingCompaction = nil + stream.UpdatedAt = time.Now().UTC() stream.mu.Unlock() + if hasCheckpoint { + cancelReason := firstNonEmpty(intent.CancelReason, "user aborted") + cancelEntry := newMetadataEntry(stream.TurnSeq, intent.RequestID, "control", map[string]any{ + "status": "canceled", + "reason": cancelReason, + "replay_policy": cancelReplayPolicyForReason(cancelReason), + }) + if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{cancelEntry}); err != nil { + log.Printf("forwarder cancellation metadata persistence failed request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, err) + if memoryErr := service.appendCheckpointEntries(stream, []HistoryEntry{cancelEntry}); memoryErr != nil { + return memoryErr + } + } + } for _, pending := range pendingExecs { _ = service.broker.Publish(intent.RequestID, StreamEvent{ Message: buildExecAbortMessage(pending), }) } - if hasCheckpoint { - if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil { - return err - } - } clearPendingProviderCompletion(stream) + terminalMessage := firstNonEmpty(intent.CancelReason, "[canceled] User aborted request") stream.mu.Lock() - stream.PendingProviderAction = providerActionNone + stream.PendingExecs = make(map[string]runtimecore.PendingExec) + stream.PendingInteractions = make(map[string]runtimecore.PendingInteraction) stream.UpdatedAt = time.Now().UTC() stream.mu.Unlock() - service.setTurnPhase(stream, TurnPhaseCanceled) - return service.broker.Cancel(intent.RequestID, firstNonEmpty(intent.CancelReason, "[canceled] User aborted request")) + service.discardPendingCheckpoint(stream, fmt.Errorf("checkpoint superseded by cancellation")) + if hasCheckpoint { + if err := service.publishCheckpointWithTerminalAction( + stream.RequestID, + stream.ConversationID, + checkpointCancellationAction(terminalMessage), + ); err != nil { + return service.failTerminalCheckpointSync(stream, err) + } + return nil + } + return service.finishCanceledTurnAfterCheckpoint(stream, terminalMessage) } // handleExecResult 处理客户端返回的执行桥结果,并在终态时把 tool_result 写回 history。 @@ -2122,9 +2150,18 @@ func (service *Service) completeSuccessfulTurn(stream *ActiveStream, completion err, ) } - if err := service.publishCheckpoint(requestID, conversationID); err != nil { + if err := service.publishCheckpointWithCompletion(requestID, conversationID, &completion); err != nil { return err } + return nil +} + +func (service *Service) finishSuccessfulTurnAfterCheckpoint(stream *ActiveStream, completion pendingTurnCompletion) error { + if stream == nil { + return nil + } + requestID := firstNonEmpty(strings.TrimSpace(completion.RequestID), strings.TrimSpace(stream.RequestID)) + usage := completion.Usage if err := service.broker.Publish(requestID, StreamEvent{ Message: buildTurnEndedMessage(usage.InputTokens, usage.OutputTokens, usage.CacheReadTokens, usage.CacheWriteTokens), }); err != nil { @@ -2151,7 +2188,15 @@ func (service *Service) failStreamIfNonTerminal(stream *ActiveStream, terminalCo } // publishCheckpoint 按当前内存会话镜像投影出 checkpoint,并广播给所有 RunSSE 订阅者。 -func (service *Service) publishCheckpoint(requestID string, _ string) error { +func (service *Service) publishCheckpoint(requestID string, conversationID string) error { + return service.publishCheckpointWithCompletion(requestID, conversationID, nil) +} + +func (service *Service) publishCheckpointWithCompletion(requestID string, conversationID string, completion *pendingTurnCompletion) error { + return service.publishCheckpointWithTerminalAction(requestID, conversationID, checkpointCompletionAction(completion)) +} + +func (service *Service) publishCheckpointWithTerminalAction(requestID string, conversationID string, terminalAction checkpointTerminalAction) error { stream, ok := service.broker.Get(requestID) if !ok || stream == nil { return fmt.Errorf("request is not active: %s", requestID) @@ -2160,15 +2205,16 @@ func (service *Service) publishCheckpoint(requestID string, _ string) error { if err != nil { return err } - state, err := service.projector.ProjectLegacyCheckpoint(conversation) + projection, err := service.projector.ProjectCheckpointProjection(conversation) if err != nil { return err } - state.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) - service.rewriteCheckpointTokenDetailsForClient(stream, conversation, state) - return service.broker.Publish(requestID, StreamEvent{ - Message: buildCheckpointMessage(state), - }) + if projection == nil || projection.State == nil { + return fmt.Errorf("checkpoint projection is empty") + } + projection.State.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) + service.rewriteCheckpointTokenDetailsForClient(stream, conversation, projection.State) + return service.queueCheckpointProjection(stream, projection, terminalAction) } func (service *Service) rewriteCheckpointTokenDetailsForClient(stream *ActiveStream, conversation *ConversationFile, state *agentv1.ConversationStateStructure) { diff --git a/internal/backend/forwarder/token_usage.go b/internal/backend/forwarder/token_usage.go index ea3e894..20ad0fe 100644 --- a/internal/backend/forwarder/token_usage.go +++ b/internal/backend/forwarder/token_usage.go @@ -45,13 +45,25 @@ func (snapshot turnUsageSnapshot) requestTokensTotal() int64 { return snapshot.promptTokensTotal() + nonNegativeInt64(snapshot.OutputTokens) } -func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure) ([]HistoryEntry, error) { +func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure, prefetchedBlobs []*agentv1.PreFetchedBlob) ([]HistoryEntry, error) { if item == nil || state == nil { return nil, nil } + blobs, err := newImportedBlobStore(prefetchedBlobs) + if err != nil { + return nil, err + } + importedIDs, err := importedTurnIDs(state.GetTurns(), blobs) + if err != nil { + return nil, err + } item.TokenDetailsUsedTokens = state.GetTokenDetails().GetUsedTokens() + item.ImportedTurnIDs = importedIDs + if minimumNextTurnSeq := int64(len(item.ImportedTurnIDs)) + 1; item.NextTurnSeq < minimumNextTurnSeq { + item.NextTurnSeq = minimumNextTurnSeq + } entries := make([]HistoryEntry, 0, 2) - if messages, err := importedConversationStateModelMessages(state); err != nil { + if messages, err := importedConversationStateModelMessages(state, blobs); err != nil { return nil, err } else { for _, message := range messages { @@ -104,7 +116,7 @@ func (service *Service) importConversationState(item *ConversationFile, state *a return entries, nil } -func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure) ([]modeladapter.Message, error) { +func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure, blobs importedBlobStore) ([]modeladapter.Message, error) { if state == nil { return nil, nil } @@ -113,7 +125,7 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru if err != nil { return nil, fmt.Errorf("decode imported replay messages: %w", err) } - decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns()) + decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns(), blobs) decoded = filterLegacyPlainWriteReplay(decoded) decoded = filterInternalPromptContextReplay(decoded) messages := make([]modeladapter.Message, 0, len(decoded)) @@ -133,35 +145,18 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru if len(rawTurn) == 0 { continue } - turn := &agentv1.ConversationTurnStructure{} - if err := proto.Unmarshal(rawTurn, turn); err != nil { - return nil, fmt.Errorf("decode imported turn: %w", err) + turn, turnID, err := decodeImportedTurn(rawTurn, blobs) + if err != nil { + return nil, err } - agentTurn := turn.GetAgentConversationTurn() - if agentTurn == nil { - continue + if turn == nil && len(turnID) > 0 { + return nil, fmt.Errorf("missing prefetched turn blob %x", turnID) } - if rawUser := agentTurn.GetUserMessage(); len(rawUser) > 0 { - userMessage := &agentv1.UserMessage{} - if err := proto.Unmarshal(rawUser, userMessage); err != nil { - return nil, fmt.Errorf("decode imported turn user_message: %w", err) - } - if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok { - messages = append(messages, toModelMessage(replay)) - } - } - for _, rawStep := range agentTurn.GetSteps() { - if len(rawStep) == 0 { - continue - } - step := &agentv1.ConversationStep{} - if err := proto.Unmarshal(rawStep, step); err != nil { - return nil, fmt.Errorf("decode imported turn step: %w", err) - } - for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) { - messages = append(messages, toModelMessage(replay)) - } + turnMessages, err := importedBlobTurnMessages(turn, blobs) + if err != nil { + return nil, err } + messages = append(messages, turnMessages...) } return normalizeReplayMessageSequence(messages), nil } diff --git a/internal/backend/forwarder/types.go b/internal/backend/forwarder/types.go index db78eee..4b5a6ef 100644 --- a/internal/backend/forwarder/types.go +++ b/internal/backend/forwarder/types.go @@ -39,6 +39,7 @@ type ConversationFile struct { CurrentPlanText string `json:"current_plan_text,omitempty"` CurrentPlans map[string]*agentv1.PlanRegistryEntry `json:"current_plans,omitempty"` CurrentTodos []*agentv1.TodoItem `json:"current_todos,omitempty"` + ImportedTurnIDs [][]byte `json:"imported_turn_ids,omitempty"` LatestRequestPrefix *ConversationRequestPrefix `json:"latest_request_prefix,omitempty"` LastProviderCall *ConversationProviderCall `json:"last_provider_call,omitempty"` CreatedAt time.Time `json:"created_at"` @@ -163,6 +164,11 @@ type ActiveStream struct { ProviderUsage turnUsageSnapshot ProviderTerminalToolInvocation bool PendingCompaction *PendingCompaction + PendingCheckpointBlobWrites map[uint32]pendingCheckpointBlobWrite + PendingCheckpointBlobRequests map[string]uint32 + NextCheckpointBlobRequestID uint32 + NextCheckpointRevision uint64 + PendingCheckpoint *pendingCheckpointPublish Backlog []StreamEvent Subscribers map[string]*StreamSubscriber @@ -219,6 +225,32 @@ type pendingTurnCompletion struct { Disposition pendingCompletionDisposition } +type pendingCheckpointBlobWrite struct { + Key string + Revision uint64 +} + +type checkpointTerminalActionKind uint8 + +const ( + checkpointTerminalActionNone checkpointTerminalActionKind = iota + checkpointTerminalActionComplete + checkpointTerminalActionCancel +) + +type checkpointTerminalAction struct { + kind checkpointTerminalActionKind + completion pendingTurnCompletion + cancelMessage string +} + +type pendingCheckpointPublish struct { + Revision uint64 + State *agentv1.ConversationStateStructure + Required map[string]struct{} + TerminalAction checkpointTerminalAction +} + type PendingCompaction struct { Trigger string ContextTokens int64 @@ -418,6 +450,7 @@ type InboundIntent struct { SubagentTypeName string SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection ConversationState *agentv1.ConversationStateStructure + PreFetchedBlobs []*agentv1.PreFetchedBlob UserMessage *agentv1.UserMessage RequestContext *agentv1.RequestContext ClientMessage *agentv1.AgentClientMessage