diff --git a/build/config.yml b/build/config.yml index c72bd69..16b98d9 100644 --- a/build/config.yml +++ b/build/config.yml @@ -8,7 +8,7 @@ info: description: "Cursor助手" copyright: "© 2026, Cursor助手" comments: "Cursor助手" - version: "0.0.44" + version: "0.0.45" dev_mode: root_path: . diff --git a/build/darwin/Info.dev.plist b/build/darwin/Info.dev.plist index 4ca9f2e..831b455 100644 --- a/build/darwin/Info.dev.plist +++ b/build/darwin/Info.dev.plist @@ -17,9 +17,9 @@ CFBundlePackageType APPL CFBundleShortVersionString - 0.0.44 + 0.0.45 CFBundleVersion - 0.0.44 + 0.0.45 LSMinimumSystemVersion 12.0.0 LSUIElement diff --git a/build/darwin/Info.plist b/build/darwin/Info.plist index 3591e80..56b50c2 100644 --- a/build/darwin/Info.plist +++ b/build/darwin/Info.plist @@ -17,9 +17,9 @@ CFBundlePackageType APPL CFBundleShortVersionString - 0.0.44 + 0.0.45 CFBundleVersion - 0.0.44 + 0.0.45 LSMinimumSystemVersion 12.0.0 LSUIElement diff --git a/build/linux/nfpm/nfpm.yaml b/build/linux/nfpm/nfpm.yaml index 7389f1f..af16591 100644 --- a/build/linux/nfpm/nfpm.yaml +++ b/build/linux/nfpm/nfpm.yaml @@ -6,7 +6,7 @@ name: "Cursor助手" arch: ${GOARCH} platform: "linux" -version: "0.0.44" +version: "0.0.45" section: "default" priority: "extra" maintainer: ${GIT_COMMITTER_NAME} <${GIT_COMMITTER_EMAIL}> diff --git a/build/windows/info.json b/build/windows/info.json index d8c9a4a..115ac3e 100644 --- a/build/windows/info.json +++ b/build/windows/info.json @@ -1,10 +1,10 @@ { "fixed": { - "file_version": "0.0.44" + "file_version": "0.0.45" }, "info": { "0000": { - "ProductVersion": "0.0.44", + "ProductVersion": "0.0.45", "CompanyName": "Cursor助手", "FileDescription": "Cursor助手", "LegalCopyright": "© 2026, Cursor助手", diff --git a/build/windows/nsis/wails_tools.nsh b/build/windows/nsis/wails_tools.nsh index f011a54..eccd696 100644 --- a/build/windows/nsis/wails_tools.nsh +++ b/build/windows/nsis/wails_tools.nsh @@ -14,7 +14,7 @@ !define INFO_PRODUCTNAME "Cursor助手" !endif !ifndef INFO_PRODUCTVERSION - !define INFO_PRODUCTVERSION "0.0.44" + !define INFO_PRODUCTVERSION "0.0.45" !endif !ifndef INFO_COPYRIGHT !define INFO_COPYRIGHT "© 2026, Cursor助手" diff --git a/build/windows/wails.exe.manifest b/build/windows/wails.exe.manifest index 3ed917e..7b30836 100644 --- a/build/windows/wails.exe.manifest +++ b/build/windows/wails.exe.manifest @@ -1,6 +1,6 @@ - + diff --git a/internal/backend/forwarder/actor.go b/internal/backend/forwarder/actor.go index a712976..2614cc1 100644 --- a/internal/backend/forwarder/actor.go +++ b/internal/backend/forwarder/actor.go @@ -23,7 +23,6 @@ const ( TurnPhaseWaitingExternal TurnPhase = "waiting_external" TurnPhaseAwaitingUser TurnPhase = "awaiting_user" TurnPhaseCompacting TurnPhase = "compacting" - TurnPhaseCheckpointing TurnPhase = "checkpointing" TurnPhaseCompleted TurnPhase = "completed" TurnPhaseFailed TurnPhase = "failed" TurnPhaseCanceled TurnPhase = "canceled" @@ -67,7 +66,6 @@ const ( streamTimerNonStreamingRecovery streamTimerKind = "non_streaming_recovery" streamTimerShellForeground streamTimerKind = "shell_foreground" streamTimerShellTransportClose streamTimerKind = "shell_transport_close" - streamTimerCheckpointBlobs streamTimerKind = "checkpoint_blobs" streamTimerOrphanCancel streamTimerKind = "orphan_cancel" ) @@ -320,9 +318,6 @@ 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) @@ -1008,8 +1003,6 @@ 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 deleted file mode 100644 index eb79b33..0000000 --- a/internal/backend/forwarder/blob_sync.go +++ /dev/null @@ -1,383 +0,0 @@ -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 deleted file mode 100644 index d101a92..0000000 --- a/internal/backend/forwarder/blob_sync_test.go +++ /dev/null @@ -1,546 +0,0 @@ -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 3d704c3..b88442c 100644 --- a/internal/backend/forwarder/broker.go +++ b/internal/backend/forwarder/broker.go @@ -82,30 +82,28 @@ 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), - PendingCheckpointBlobWrites: make(map[uint32]pendingCheckpointBlobWrite), - PendingCheckpointBlobRequests: make(map[string]uint32), - 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), + 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 ef69776..7b02404 100644 --- a/internal/backend/forwarder/events.go +++ b/internal/backend/forwarder/events.go @@ -245,22 +245,6 @@ 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 fb96d1b..0732614 100644 --- a/internal/backend/forwarder/file_store.go +++ b/internal/backend/forwarder/file_store.go @@ -750,7 +750,6 @@ 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)) { @@ -883,7 +882,6 @@ 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 deleted file mode 100644 index fbd17ed..0000000 --- a/internal/backend/forwarder/imported_blobs.go +++ /dev/null @@ -1,167 +0,0 @@ -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 3f612d2..5726f8f 100644 --- a/internal/backend/forwarder/projector.go +++ b/internal/backend/forwarder/projector.go @@ -2,7 +2,6 @@ package forwarder import ( - "crypto/sha256" "encoding/json" "fmt" "strings" @@ -20,51 +19,6 @@ 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{} @@ -521,16 +475,6 @@ 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), @@ -544,7 +488,7 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con if conversation == nil { mode := agentv1.AgentMode_AGENT_MODE_AGENT state.Mode = &mode - return &CheckpointProjection{State: state}, nil + return state, nil } mode, err := parseModeAlias(conversation.Mode) if err != nil { @@ -562,11 +506,158 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con if structuredState.HasTodos { state.Todos = encodeConversationTodoBytes(structuredState.Todos) } - turnIDs, err := projectCheckpointTurnBlobs(conversation, blobs) - if err != nil { - return nil, err + 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) } - state.Turns = append(cloneByteSlices(conversation.ImportedTurnIDs), turnIDs...) replayMessages, err := projector.ProjectPromptReplay(conversation) if err != nil { return nil, err @@ -594,183 +685,15 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con return nil, err } state.RootPromptMessagesJson = rootPromptMessages - return &CheckpointProjection{State: state, Blobs: blobs.list()}, nil + return state, nil } -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 marshalThinkingStep(text string) ([]byte, error) { + return proto.Marshal(&agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: text}, + }, + }) } func conversationTokenDetailsUsedTokens(conversation *ConversationFile) uint32 { @@ -1218,6 +1141,60 @@ 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 @@ -1254,7 +1231,7 @@ func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []pro return filtered } -func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte, blobs importedBlobStore) []promptengine.Message { +func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte) []promptengine.Message { if len(messages) == 0 || len(importedTurns) == 0 { return messages } @@ -1263,16 +1240,16 @@ func restoreImportedReplayUserMessages(messages []promptengine.Message, imported if len(rawTurn) == 0 { continue } - turn, _, err := decodeImportedTurn(rawTurn, blobs) - if err != nil || turn == nil { + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(rawTurn, turn); err != nil { continue } agentTurn := turn.GetAgentConversationTurn() if agentTurn == nil || len(agentTurn.GetUserMessage()) == 0 { continue } - userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs) - if err != nil { + userMessage := &agentv1.UserMessage{} + if err := proto.Unmarshal(agentTurn.GetUserMessage(), userMessage); err != nil { continue } replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage) diff --git a/internal/backend/forwarder/projector_test.go b/internal/backend/forwarder/projector_test.go deleted file mode 100644 index 10330ba..0000000 --- a/internal/backend/forwarder/projector_test.go +++ /dev/null @@ -1,433 +0,0 @@ -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 1f11ee0..ba7c851 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{} + decision := runRewindDecision{ClientTurnCount: -1} 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 { + if hasClientTurnCount && clientTurnCount >= 0 { targetTurnSeq := int64(clientTurnCount) + 1 for _, match := range matches { if match.Entry.TurnSeq == targetTurnSeq { @@ -224,7 +224,6 @@ 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) @@ -270,30 +269,10 @@ 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 b302c8a..9abbba3 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, intent.PreFetchedBlobs) + importedEntries, err = service.importConversationState(conversation, intent.ConversationState) if err != nil { return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err } @@ -138,7 +138,6 @@ 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 fb4c848..5bb04ca 100644 --- a/internal/backend/forwarder/service.go +++ b/internal/backend/forwarder/service.go @@ -9,7 +9,6 @@ import ( "log" "sort" "strings" - "sync" "time" "connectrpc.com/connect" @@ -262,8 +261,6 @@ type Service struct { execBridge execbridge.ExecBridge interactionBridge interactionbridge.InteractionBridge appendSeq *appendSequenceTracker - checkpointBlobMu sync.Mutex - checkpointBlobs map[string]*checkpointBlobCacheEntry } type agentModelMemory interface { @@ -303,7 +300,6 @@ 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() @@ -332,7 +328,6 @@ func newServiceWithDependencies(store *ConversationFileStore, projector *History execBridge: execbridge.NewBridge(), interactionBridge: interactionbridge.NewBridge(), appendSeq: newAppendSequenceTracker(), - checkpointBlobs: make(map[string]*checkpointBlobCacheEntry), } } @@ -562,7 +557,6 @@ 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) { @@ -610,7 +604,6 @@ 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 @@ -826,14 +819,11 @@ func (service *Service) snapshotVisibleTurns(conversation *ConversationFile) ([] if service == nil || service.projector == nil || conversation == nil { return nil, nil } - projection, err := service.projector.ProjectCheckpointProjection(conversation) + state, err := service.projector.ProjectLegacyCheckpoint(conversation) if err != nil { return nil, err } - if projection == nil || projection.State == nil { - return nil, fmt.Errorf("checkpoint projection is empty") - } - return cloneByteSlices(projection.State.GetTurns()), nil + return cloneByteSlices(state.GetTurns()), nil } // handleCancelIntent 处理取消请求,并向客户端发送执行桥 abort。 @@ -843,60 +833,42 @@ 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.PendingExecs = make(map[string]runtimecore.PendingExec) - stream.PendingInteractions = make(map[string]runtimecore.PendingInteraction) + stream.PendingProviderAction = providerActionNone stream.UpdatedAt = time.Now().UTC() stream.mu.Unlock() - 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) + service.setTurnPhase(stream, TurnPhaseCanceled) + return service.broker.Cancel(intent.RequestID, firstNonEmpty(intent.CancelReason, "[canceled] User aborted request")) } // handleExecResult 处理客户端返回的执行桥结果,并在终态时把 tool_result 写回 history。 @@ -2150,18 +2122,9 @@ func (service *Service) completeSuccessfulTurn(stream *ActiveStream, completion err, ) } - if err := service.publishCheckpointWithCompletion(requestID, conversationID, &completion); err != nil { + if err := service.publishCheckpoint(requestID, conversationID); 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 { @@ -2188,15 +2151,7 @@ func (service *Service) failStreamIfNonTerminal(stream *ActiveStream, terminalCo } // publishCheckpoint 按当前内存会话镜像投影出 checkpoint,并广播给所有 RunSSE 订阅者。 -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 { +func (service *Service) publishCheckpoint(requestID string, _ string) error { stream, ok := service.broker.Get(requestID) if !ok || stream == nil { return fmt.Errorf("request is not active: %s", requestID) @@ -2205,16 +2160,15 @@ func (service *Service) publishCheckpointWithTerminalAction(requestID string, co if err != nil { return err } - projection, err := service.projector.ProjectCheckpointProjection(conversation) + state, err := service.projector.ProjectLegacyCheckpoint(conversation) if err != nil { return err } - 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) + state.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) + service.rewriteCheckpointTokenDetailsForClient(stream, conversation, state) + return service.broker.Publish(requestID, StreamEvent{ + Message: buildCheckpointMessage(state), + }) } 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 20ad0fe..ea3e894 100644 --- a/internal/backend/forwarder/token_usage.go +++ b/internal/backend/forwarder/token_usage.go @@ -45,25 +45,13 @@ func (snapshot turnUsageSnapshot) requestTokensTotal() int64 { return snapshot.promptTokensTotal() + nonNegativeInt64(snapshot.OutputTokens) } -func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure, prefetchedBlobs []*agentv1.PreFetchedBlob) ([]HistoryEntry, error) { +func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure) ([]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, blobs); err != nil { + if messages, err := importedConversationStateModelMessages(state); err != nil { return nil, err } else { for _, message := range messages { @@ -116,7 +104,7 @@ func (service *Service) importConversationState(item *ConversationFile, state *a return entries, nil } -func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure, blobs importedBlobStore) ([]modeladapter.Message, error) { +func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure) ([]modeladapter.Message, error) { if state == nil { return nil, nil } @@ -125,7 +113,7 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru if err != nil { return nil, fmt.Errorf("decode imported replay messages: %w", err) } - decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns(), blobs) + decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns()) decoded = filterLegacyPlainWriteReplay(decoded) decoded = filterInternalPromptContextReplay(decoded) messages := make([]modeladapter.Message, 0, len(decoded)) @@ -145,18 +133,35 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru if len(rawTurn) == 0 { continue } - turn, turnID, err := decodeImportedTurn(rawTurn, blobs) - if err != nil { - return nil, err + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(rawTurn, turn); err != nil { + return nil, fmt.Errorf("decode imported turn: %w", err) } - if turn == nil && len(turnID) > 0 { - return nil, fmt.Errorf("missing prefetched turn blob %x", turnID) + agentTurn := turn.GetAgentConversationTurn() + if agentTurn == nil { + continue } - turnMessages, err := importedBlobTurnMessages(turn, blobs) - if err != nil { - return nil, err + 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)) + } } - messages = append(messages, turnMessages...) } return normalizeReplayMessageSequence(messages), nil } diff --git a/internal/backend/forwarder/types.go b/internal/backend/forwarder/types.go index 4b5a6ef..db78eee 100644 --- a/internal/backend/forwarder/types.go +++ b/internal/backend/forwarder/types.go @@ -39,7 +39,6 @@ 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"` @@ -164,11 +163,6 @@ 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 @@ -225,32 +219,6 @@ 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 @@ -450,7 +418,6 @@ 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 diff --git a/release-notes.md b/release-notes.md index 1c11759..690b827 100644 --- a/release-notes.md +++ b/release-notes.md @@ -9,4 +9,5 @@ Tg群组: https://t.me/cursor_byok - 支持cursor-cli +- 修复对话中错误可能导致的消失问题