diff --git a/README.md b/README.md index 16d7a2f..84f0be1 100644 --- a/README.md +++ b/README.md @@ -20,7 +20,7 @@ https://t.me/cursor_byok ## 路线图 [正式版路线图](https://github.com/leookun/cursor-byok/discussions/32) -[详细使用教程](https://dcne38qm5vlg.feishu.cn/wiki/JeP7wdGnziBXuikNaF5czWbrn8c) +[详细使用教程](https://docs.leokun.cn) ## 后续 @@ -37,13 +37,3 @@ https://t.me/cursor_byok - -## Star History - - - - - - Star History Chart - - diff --git a/build/config.yml b/build/config.yml index 60eb77c..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.43" + version: "0.0.45" dev_mode: root_path: . diff --git a/build/darwin/Info.dev.plist b/build/darwin/Info.dev.plist index 1aa1487..831b455 100644 --- a/build/darwin/Info.dev.plist +++ b/build/darwin/Info.dev.plist @@ -17,9 +17,9 @@ CFBundlePackageType APPL CFBundleShortVersionString - 0.0.43 + 0.0.45 CFBundleVersion - 0.0.43 + 0.0.45 LSMinimumSystemVersion 12.0.0 LSUIElement diff --git a/build/darwin/Info.plist b/build/darwin/Info.plist index 1363dd9..56b50c2 100644 --- a/build/darwin/Info.plist +++ b/build/darwin/Info.plist @@ -17,9 +17,9 @@ CFBundlePackageType APPL CFBundleShortVersionString - 0.0.43 + 0.0.45 CFBundleVersion - 0.0.43 + 0.0.45 LSMinimumSystemVersion 12.0.0 LSUIElement diff --git a/build/linux/nfpm/nfpm.yaml b/build/linux/nfpm/nfpm.yaml index 96a9d63..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.43" +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 6ad5d43..115ac3e 100644 --- a/build/windows/info.json +++ b/build/windows/info.json @@ -1,10 +1,10 @@ { "fixed": { - "file_version": "0.0.43" + "file_version": "0.0.45" }, "info": { "0000": { - "ProductVersion": "0.0.43", + "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 4c944bb..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.43" + !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 8b50f75..7b30836 100644 --- a/build/windows/wails.exe.manifest +++ b/build/windows/wails.exe.manifest @@ -1,6 +1,6 @@ - + diff --git a/internal/app/runner.go b/internal/app/runner.go index d4ee32f..6ab7666 100644 --- a/internal/app/runner.go +++ b/internal/app/runner.go @@ -8,7 +8,9 @@ import ( "encoding/pem" "io/fs" "net" + "os" goruntime "runtime" + "strconv" "strings" "time" @@ -37,6 +39,8 @@ const ( appName = "Cursor助手" // adRefreshInterval 表示后台广告拉取间隔。 adRefreshInterval = 3 * time.Minute + // disableWebViewSandboxEnv allows affected VDI users to opt out of the WebView2 sandbox. + disableWebViewSandboxEnv = "CURSOR_BYOK_DISABLE_WEBVIEW_SANDBOX" ) // EmbeddedResources 定义了当前模块中的 EmbeddedResources 类型。 @@ -134,6 +138,9 @@ func Run(resources EmbeddedResources) error { Assets: application.AssetOptions{ Handler: application.AssetFileServerFS(resources.Assets), }, + Windows: application.WindowsOptions{ + AdditionalBrowserArgs: windowsAdditionalBrowserArgs(), + }, Mac: application.MacOptions{ ActivationPolicy: application.ActivationPolicyAccessory, ApplicationShouldTerminateAfterLastWindowClosed: false, @@ -415,6 +422,14 @@ func Run(resources EmbeddedResources) error { return app.Run() } +func windowsAdditionalBrowserArgs() []string { + disableSandbox, err := strconv.ParseBool(strings.TrimSpace(os.Getenv(disableWebViewSandboxEnv))) + if err != nil || !disableSandbox { + return nil + } + return []string{"--no-sandbox"} +} + func browserReachableLoopbackBaseURL(listenAddr string) string { host, port, err := net.SplitHostPort(strings.TrimSpace(listenAddr)) if err != nil || strings.TrimSpace(port) == "" { diff --git a/internal/app/runner_test.go b/internal/app/runner_test.go new file mode 100644 index 0000000..1010b38 --- /dev/null +++ b/internal/app/runner_test.go @@ -0,0 +1,29 @@ +package app + +import ( + "reflect" + "testing" +) + +func TestWindowsAdditionalBrowserArgs(t *testing.T) { + tests := []struct { + name string + env string + want []string + }{ + {name: "unset"}, + {name: "enabled with one", env: "1", want: []string{"--no-sandbox"}}, + {name: "enabled with true", env: " true ", want: []string{"--no-sandbox"}}, + {name: "disabled", env: "false"}, + {name: "invalid", env: "yes"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Setenv(disableWebViewSandboxEnv, tt.env) + if got := windowsAdditionalBrowserArgs(); !reflect.DeepEqual(got, tt.want) { + t.Fatalf("windowsAdditionalBrowserArgs() = %v, want %v", got, tt.want) + } + }) + } +} 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..7a0f1ce 100644 --- a/internal/backend/forwarder/broker.go +++ b/internal/backend/forwarder/broker.go @@ -76,36 +76,42 @@ func (broker *StreamBroker) OpenStream(requestID string, conversationID string, if existing.BackgroundShellActions == nil { existing.BackgroundShellActions = make(map[string]time.Time) } + if existing.PendingCheckpointBlobWrites == nil { + existing.PendingCheckpointBlobWrites = make(map[uint32]string) + } + if existing.ConfirmedCheckpointBlobs == nil { + existing.ConfirmedCheckpointBlobs = make(map[string]struct{}) + } existing.UpdatedAt = time.Now().UTC() existing.mu.Unlock() return existing, nil } 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), + PendingCheckpointBlobWrites: make(map[uint32]string), + ConfirmedCheckpointBlobs: make(map[string]struct{}), + CreatedAt: now, + UpdatedAt: now, } broker.streams[normalizedRequestID] = stream return stream, nil diff --git a/internal/backend/forwarder/checkpoint_blobs.go b/internal/backend/forwarder/checkpoint_blobs.go new file mode 100644 index 0000000..c07c07f --- /dev/null +++ b/internal/backend/forwarder/checkpoint_blobs.go @@ -0,0 +1,238 @@ +package forwarder + +import ( + "encoding/hex" + "fmt" + "log" + "strings" + "time" + + "google.golang.org/protobuf/proto" + + "cursor/gen/agentv1" +) + +const checkpointBlobWriteTimeout = 5 * time.Second + +type pendingCheckpointBlobWrite struct { + requestID uint32 + blob CheckpointBlob +} + +func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion { + if completion == nil { + return nil + } + cloned := *completion + return &cloned +} + +func (service *Service) queueCheckpointProjection(stream *ActiveStream, projection *CheckpointProjection, completion *pendingTurnCompletion) 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]string) + } + if stream.ConfirmedCheckpointBlobs == nil { + stream.ConfirmedCheckpointBlobs = make(map[string]struct{}) + } + if completion == nil && stream.PendingCheckpoint != nil { + completion = stream.PendingCheckpoint.Completion + } + required := make(map[string]struct{}, len(projection.Blobs)) + pendingKeys := make(map[string]struct{}, len(stream.PendingCheckpointBlobWrites)) + for _, key := range stream.PendingCheckpointBlobWrites { + pendingKeys[key] = struct{}{} + } + toWrite := make([]pendingCheckpointBlobWrite, 0, len(projection.Blobs)) + for _, blob := range projection.Blobs { + key := string(blob.ID) + if key == "" { + continue + } + required[key] = struct{}{} + if _, confirmed := stream.ConfirmedCheckpointBlobs[key]; confirmed { + continue + } + if _, pending := pendingKeys[key]; pending { + continue + } + stream.NextCheckpointBlobRequestID++ + if stream.NextCheckpointBlobRequestID == 0 { + stream.NextCheckpointBlobRequestID++ + } + requestID := stream.NextCheckpointBlobRequestID + stream.PendingCheckpointBlobWrites[requestID] = key + pendingKeys[key] = struct{}{} + toWrite = append(toWrite, pendingCheckpointBlobWrite{requestID: requestID, blob: blob}) + } + stream.PendingCheckpoint = &pendingCheckpointPublish{ + State: state, + Required: required, + Completion: clonePendingTurnCompletion(completion), + } + if completion != nil { + stream.Phase = TurnPhaseCheckpointing + } + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + + for _, write := range toWrite { + if err := service.broker.Publish(stream.RequestID, StreamEvent{ + Message: buildSetCheckpointBlobMessage(write.requestID, write.blob), + }); err != nil { + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf("publish checkpoint blob: %w", err)) + } + } + if service.checkpointProjectionReady(stream) { + return service.publishReadyCheckpoint(stream) + } + service.scheduleStreamTimer( + stream, + providerTimerKey(streamTimerCheckpointBlobs, ""), + checkpointBlobWriteTimeout, + streamTimerCheckpointBlobs, + "", + 0, + "checkpoint blob write timeout", + ) + return nil +} + +func (service *Service) checkpointProjectionReady(stream *ActiveStream) bool { + if stream == nil { + return false + } + stream.mu.Lock() + defer stream.mu.Unlock() + if stream.PendingCheckpoint == nil { + return false + } + for key := range stream.PendingCheckpoint.Required { + if _, confirmed := stream.ConfirmedCheckpointBlobs[key]; !confirmed { + return false + } + } + return true +} + +func (service *Service) handleCheckpointBlobResult(stream *ActiveStream, message *agentv1.KvClientMessage) error { + if service == nil || stream == nil || message == nil || message.GetSetBlobResult() == nil { + return nil + } + stream.mu.Lock() + key, ok := stream.PendingCheckpointBlobWrites[message.GetId()] + if ok { + delete(stream.PendingCheckpointBlobWrites, message.GetId()) + } + required := false + if ok && stream.PendingCheckpoint != nil { + _, required = stream.PendingCheckpoint.Required[key] + } + if ok && message.GetSetBlobResult().GetError() == nil { + stream.ConfirmedCheckpointBlobs[key] = struct{}{} + } + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + if !ok { + return nil + } + if blobErr := message.GetSetBlobResult().GetError(); blobErr != nil && required { + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf( + "client rejected checkpoint blob %s: %s", + hex.EncodeToString([]byte(key)), + firstNonEmpty(strings.TrimSpace(blobErr.GetMessage()), "unknown error"), + )) + } + if service.checkpointProjectionReady(stream) { + return service.publishReadyCheckpoint(stream) + } + return nil +} + +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 _, confirmed := stream.ConfirmedCheckpointBlobs[key]; !confirmed { + stream.mu.Unlock() + return nil + } + } + stream.PendingCheckpoint = nil + state := pending.State + completion := clonePendingTurnCompletion(pending.Completion) + 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 { + if completion != nil { + log.Printf("forwarder checkpoint publish skipped before successful terminal request_id=%s err=%v", stream.RequestID, err) + return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion) + } + return err + } + if completion != nil { + return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion) + } + return nil +} + +func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error { + if stream == nil { + return nil + } + stream.mu.Lock() + pendingCount := len(stream.PendingCheckpointBlobWrites) + stream.mu.Unlock() + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf("%d checkpoint blob writes timed out", pendingCount)) +} + +func (service *Service) finishAfterCheckpointSyncFailure(stream *ActiveStream, cause error) error { + if stream == nil { + return nil + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + stream.PendingCheckpoint = nil + stream.PendingCheckpointBlobWrites = make(map[uint32]string) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if cause != nil { + log.Printf("forwarder checkpoint blob sync skipped request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause) + } + if pending != nil && pending.Completion != nil { + return service.finishSuccessfulTurnAfterCheckpoint(stream, *pending.Completion) + } + return nil +} + +func (service *Service) discardPendingCheckpoint(stream *ActiveStream, reason string) { + if stream == nil { + return + } + stream.mu.Lock() + stream.PendingCheckpoint = nil + stream.PendingCheckpointBlobWrites = make(map[uint32]string) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if strings.TrimSpace(reason) != "" { + log.Printf("forwarder pending checkpoint discarded request_id=%s conversation_id=%s reason=%s", stream.RequestID, stream.ConversationID, strings.TrimSpace(reason)) + } +} diff --git a/internal/backend/forwarder/checkpoint_blobs_test.go b/internal/backend/forwarder/checkpoint_blobs_test.go new file mode 100644 index 0000000..9ada5d3 --- /dev/null +++ b/internal/backend/forwarder/checkpoint_blobs_test.go @@ -0,0 +1,208 @@ +package forwarder + +import ( + "testing" + + "google.golang.org/protobuf/encoding/protojson" + + "cursor/gen/agentv1" +) + +func TestCheckpointBlobSyncPublishesCheckpointAfterAcknowledgements(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + if err := service.queueCheckpointProjection(stream, projection, nil); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + events := readCheckpointTestEvents(t, service, stream) + 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) + } + } + + acknowledgeCheckpointBlobs(t, service, stream) + events = readCheckpointTestEvents(t, service, stream) + 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 TestCheckpointBlobSyncPublishesCheckpointBeforeSuccessfulTerminal(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + completion := &pendingTurnCompletion{ + RequestID: stream.RequestID, + Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7}, + } + if err := service.queueCheckpointProjection(stream, projection, completion); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + acknowledgeCheckpointBlobs(t, service, stream) + + events := readCheckpointTestEvents(t, service, stream) + 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 TestCheckpointBlobTimeoutDoesNotFailSuccessfulTurn(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + completion := &pendingTurnCompletion{ + RequestID: stream.RequestID, + Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7}, + } + if err := service.queueCheckpointProjection(stream, projection, completion); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + if err := service.handleCheckpointBlobTimeout(stream); err != nil { + t.Fatalf("handleCheckpointBlobTimeout() error = %v", err) + } + + events := readCheckpointTestEvents(t, service, stream) + var checkpoint, turnEnded, successfulEnd bool + for _, event := range events { + checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil + turnEnded = turnEnded || event.Message.GetInteractionUpdate().GetTurnEnded() != nil + successfulEnd = successfulEnd || event.End && event.TerminalErrorCode == "" + } + if checkpoint || !turnEnded || !successfulEnd { + t.Fatalf("timeout events checkpoint=%v turn_ended=%v successful_end=%v", checkpoint, turnEnded, successfulEnd) + } +} + +func TestCancellationDiscardsPendingCheckpointAndIgnoresLateAcknowledgements(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + if err := service.queueCheckpointProjection(stream, projection, nil); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + stream.mu.Lock() + requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites)) + for requestID := range stream.PendingCheckpointBlobWrites { + requestIDs = append(requestIDs, requestID) + } + stream.mu.Unlock() + if err := service.handleCancelIntent(InboundIntent{ + Kind: "cancel", + RequestID: stream.RequestID, + CancelReason: "user stopped", + }); err != nil { + t.Fatalf("handleCancelIntent() error = %v", err) + } + for _, requestID := range requestIDs { + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + t.Fatalf("late ACK %d error = %v", requestID, err) + } + } + + events := readCheckpointTestEvents(t, service, stream) + var checkpoint, canceledEnd bool + for _, event := range events { + checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil + canceledEnd = canceledEnd || event.End && event.TerminalErrorCode == "canceled" + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + stream.mu.Unlock() + if checkpoint || !canceledEnd || pending != nil { + t.Fatalf("cancel events checkpoint=%v canceled_end=%v pending=%v", checkpoint, canceledEnd, pending != nil) + } +} + +func testCheckpointBlobProjection(t *testing.T) (*Service, *ActiveStream, *CheckpointProjection) { + t.Helper() + broker := NewStreamBroker() + service := &Service{ + store: NewConversationFileStore(t.TempDir()), + projector: NewHistoryProjector(), + broker: broker, + } + 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) + } + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 2, + NextEntrySeq: 3, + TokenDetailsMaxTokens: projectedConversationMaxTokens, + Entries: []HistoryEntry{ + testCheckpointUserEntry(t), + newAssistantTextEntry(1, "request-1", "hi", "", ""), + }, + } + projection, err := service.projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + if err := service.replaceCheckpointConversation(stream, conversation); err != nil { + t.Fatalf("replaceCheckpointConversation() error = %v", err) + } + return service, stream, projection +} + +func testCheckpointUserEntry(t *testing.T) HistoryEntry { + t.Helper() + payload, err := protojson.Marshal(&agentv1.UserMessage{Text: "hello", MessageId: "message-1"}) + if err != nil { + t.Fatalf("marshal user message: %v", err) + } + return HistoryEntry{Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: payload} +} + +func acknowledgeCheckpointBlobs(t *testing.T, service *Service, stream *ActiveStream) { + t.Helper() + 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 + } + for _, requestID := range requestIDs { + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + t.Fatalf("handleCheckpointBlobResult(%d) error = %v", requestID, err) + } + } + } +} + +func readCheckpointTestEvents(t *testing.T, service *Service, stream *ActiveStream) []StreamEvent { + t.Helper() + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("ReadFromCursor() error = %v", err) + } + return events +} 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..7095860 100644 --- a/internal/backend/forwarder/projector.go +++ b/internal/backend/forwarder/projector.go @@ -566,7 +566,7 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con if err != nil { return nil, err } - state.Turns = append(cloneByteSlices(conversation.ImportedTurnIDs), turnIDs...) + state.Turns = turnIDs replayMessages, err := projector.ProjectPromptReplay(conversation) if err != nil { return nil, err @@ -1218,6 +1218,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 +1308,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 +1317,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_fork_test.go b/internal/backend/forwarder/projector_fork_test.go new file mode 100644 index 0000000..56fc0af --- /dev/null +++ b/internal/backend/forwarder/projector_fork_test.go @@ -0,0 +1,140 @@ +package forwarder + +import ( + "crypto/sha256" + "strings" + "testing" + + "google.golang.org/protobuf/encoding/protojson" + "google.golang.org/protobuf/proto" + + "cursor/gen/agentv1" +) + +func TestProjectCheckpointProjectionBuildsResolvableForkState(t *testing.T) { + userPayload, err := protojson.Marshal(&agentv1.UserMessage{ + Text: "parent question", + MessageId: "message-1", + }) + if err != nil { + t.Fatalf("marshal user message: %v", err) + } + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 2, + NextEntrySeq: 3, + TokenDetailsMaxTokens: projectedConversationMaxTokens, + Entries: []HistoryEntry{ + {Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: userPayload}, + newAssistantTextEntry(1, "request-1", "parent answer", "", ""), + }, + } + + projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + state := projection.State + if len(state.GetTurns()) != 1 { + t.Fatalf("ProjectCheckpointProjection() turns = %d, want 1 Blob-backed turn", len(state.GetTurns())) + } + blobs := make(map[string][]byte, len(projection.Blobs)) + for _, blob := range projection.Blobs { + digest := sha256.Sum256(blob.Data) + if len(blob.ID) != sha256.Size || string(blob.ID) != string(digest[:]) { + t.Fatalf("invalid content-addressed Blob id=%x", blob.ID) + } + blobs[string(blob.ID)] = blob.Data + } + turnPayload, ok := blobs[string(state.GetTurns()[0])] + if !ok { + t.Fatal("turn references a missing Blob") + } + turn := &agentv1.ConversationTurnStructure{} + if err := proto.Unmarshal(turnPayload, turn); err != nil { + t.Fatalf("decode turn Blob: %v", err) + } + agentTurn := turn.GetAgentConversationTurn() + if agentTurn == nil { + t.Fatal("turn Blob does not contain an agent turn") + } + if _, ok := blobs[string(agentTurn.GetUserMessage())]; !ok { + t.Fatal("turn references a missing user message Blob") + } + for _, stepID := range agentTurn.GetSteps() { + if _, ok := blobs[string(stepID)]; !ok { + t.Fatal("turn references a missing step Blob") + } + } + + messages, err := importedConversationStateModelMessages(state) + if err != nil { + t.Fatalf("importedConversationStateModelMessages() error = %v", err) + } + if len(messages) != 2 { + t.Fatalf("imported messages = %d, want parent user and assistant context", len(messages)) + } + if messages[0].Role != "user" || !strings.Contains(messages[0].Content, "parent question") { + t.Fatalf("first imported message = %#v", messages[0]) + } + if messages[1].Role != "assistant" || messages[1].Content != "parent answer" { + t.Fatalf("second imported message = %#v", messages[1]) + } +} + +func TestProjectCheckpointProjectionKeepsForkPointIsolatedFromLaterHistory(t *testing.T) { + firstUser, err := protojson.Marshal(&agentv1.UserMessage{Text: "first question", MessageId: "message-1"}) + if err != nil { + t.Fatalf("marshal first user message: %v", err) + } + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 2, + NextEntrySeq: 3, + TokenDetailsMaxTokens: projectedConversationMaxTokens, + Entries: []HistoryEntry{ + {Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: firstUser}, + newAssistantTextEntry(1, "request-1", "first answer", "", ""), + }, + } + projector := NewHistoryProjector() + midpoint, err := projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("midpoint projection: %v", err) + } + + secondUser, err := protojson.Marshal(&agentv1.UserMessage{Text: "second question", MessageId: "message-2"}) + if err != nil { + t.Fatalf("marshal second user message: %v", err) + } + appendEntriesInPlace(conversation, []HistoryEntry{ + {TurnSeq: 2, RequestID: "request-2", Role: "user", Kind: "user_message", Payload: secondUser}, + newAssistantTextEntry(2, "request-2", "second answer", "", ""), + }) + latest, err := projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("latest projection: %v", err) + } + + midpointMessages, err := importedConversationStateModelMessages(midpoint.State) + if err != nil { + t.Fatalf("import midpoint messages: %v", err) + } + latestMessages, err := importedConversationStateModelMessages(latest.State) + if err != nil { + t.Fatalf("import latest messages: %v", err) + } + if len(midpoint.State.GetTurns()) != 1 || len(midpointMessages) != 2 { + t.Fatalf("midpoint turns=%d messages=%d, want 1 turn and 2 messages", len(midpoint.State.GetTurns()), len(midpointMessages)) + } + if len(latest.State.GetTurns()) != 2 || len(latestMessages) != 4 { + t.Fatalf("latest turns=%d messages=%d, want 2 turns and 4 messages", len(latest.State.GetTurns()), len(latestMessages)) + } + if midpointMessages[1].Content != "first answer" || latestMessages[3].Content != "second answer" { + t.Fatalf("fork snapshots are not isolated: midpoint=%#v latest=%#v", midpointMessages, latestMessages) + } +} 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..4cf5c5d 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,40 @@ 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 { + service.discardPendingCheckpoint(stream, "checkpoint superseded by cancellation") + } 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,10 +2120,7 @@ func (service *Service) completeSuccessfulTurn(stream *ActiveStream, completion err, ) } - if err := service.publishCheckpointWithCompletion(requestID, conversationID, &completion); err != nil { - return err - } - return nil + return service.publishCheckpointWithCompletion(requestID, conversationID, &completion) } func (service *Service) finishSuccessfulTurnAfterCheckpoint(stream *ActiveStream, completion pendingTurnCompletion) error { @@ -2192,11 +2159,7 @@ func (service *Service) publishCheckpoint(requestID string, conversationID strin 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) publishCheckpointWithCompletion(requestID string, _ string, completion *pendingTurnCompletion) error { stream, ok := service.broker.Get(requestID) if !ok || stream == nil { return fmt.Errorf("request is not active: %s", requestID) @@ -2214,7 +2177,7 @@ func (service *Service) publishCheckpointWithTerminalAction(requestID string, co } projection.State.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) service.rewriteCheckpointTokenDetailsForClient(stream, conversation, projection.State) - return service.queueCheckpointProjection(stream, projection, terminalAction) + return service.queueCheckpointProjection(stream, projection, completion) } 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..7b33d20 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,10 +163,9 @@ type ActiveStream struct { ProviderUsage turnUsageSnapshot ProviderTerminalToolInvocation bool PendingCompaction *PendingCompaction - PendingCheckpointBlobWrites map[uint32]pendingCheckpointBlobWrite - PendingCheckpointBlobRequests map[string]uint32 + PendingCheckpointBlobWrites map[uint32]string + ConfirmedCheckpointBlobs map[string]struct{} NextCheckpointBlobRequestID uint32 - NextCheckpointRevision uint64 PendingCheckpoint *pendingCheckpointPublish Backlog []StreamEvent @@ -225,30 +223,10 @@ 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 + State *agentv1.ConversationStateStructure + Required map[string]struct{} + Completion *pendingTurnCompletion } type PendingCompaction struct { @@ -450,7 +428,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/internal/backend/host.go b/internal/backend/host.go index 974153f..3c84d53 100644 --- a/internal/backend/host.go +++ b/internal/backend/host.go @@ -319,6 +319,16 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error { MockBuilder: upstream.ServerConfigMockBuilder, })), ), + server.POST("/aiserver.v1.ServerConfigService/GetServerConfig", + server.Name("server_config_service_get_server_config"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "server_config_service_get_server_config", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetServerConfigResponse", + MockBuilder: upstream.ServerConfigMockBuilder, + })), + ), server.POST("/aiserver.v1.AiService/AvailableModels", server.Name("available_models"), server.ConnectUnary(), @@ -329,6 +339,36 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error { MockBuilder: upstream.AvailableModelsMockBuilder, })), ), + server.POST("/aiserver.v1.AiService/GetUsableModels", + server.Name("usable_models"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "usable_models", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetUsableModelsResponse", + MockBuilder: upstream.UsableModelsMockBuilder, + })), + ), + server.POST("/aiserver.v1.AiService/GetDefaultModelForCli", + server.Name("default_model_for_cli"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "default_model_for_cli", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetDefaultModelForCliResponse", + MockBuilder: upstream.DefaultModelForCliMockBuilder, + })), + ), + server.POST("/aiserver.v1.AiService/GetDefaultModel", + server.Name("default_model"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "default_model", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetDefaultModelResponse", + MockBuilder: upstream.DefaultModelMockBuilder, + })), + ), server.POST("/aiserver.v1.AiService/GetDefaultModelNudgeData", server.Name("default_model_nudge"), server.ConnectUnary(), @@ -359,6 +399,34 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error { MockBuilder: upstream.FirstWindowStatsigDecisionMockBuilder, })), ), + server.POST("/aiserver.v1.AnalyticsService/SubmitLogs", + server.Name("analytics_submit_logs"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "analytics_submit_logs", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.SubmitLogsResponse", + MockBuilder: upstream.SubmitLogsMockBuilder, + })), + ), + server.POST("/aiserver.v1.AnalyticsService/TrackEvents", + server.Name("analytics_track_events"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "analytics_track_events", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.TrackEventsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + ), + server.POST("/v1/traces", + server.Name("otlp_traces"), + server.HTTP(), + server.Local(upstream.FixedStatusAction(routeDeps, upstream.CompatRouteConfig{ + Name: "otlp_traces", + StatusCode: http.StatusOK, + })), + ), server.POST("/oauth/token", server.Name("oauth_token"), server.HTTP(), @@ -477,6 +545,76 @@ func (host *Host) rebuildLocked(cfg serverconfig.Config) error { }), )), ), + server.POST("/aiserver.v1.DashboardService/GetTeamAdminSettingsOrEmptyIfNotInTeam", + server.Name("dashboard_get_team_admin_settings_or_empty"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_team_admin_settings_or_empty", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetTeamAdminSettingsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + ), + server.POST("/aiserver.v1.DashboardService/GetTeamReposOrEmptyIfNotInTeam", + server.Name("dashboard_get_team_repos_or_empty"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_team_repos_or_empty", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetTeamReposResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + ), + server.POST("/aiserver.v1.DashboardService/ListMarketplaces", + server.Name("dashboard_list_marketplaces"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_list_marketplaces", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.ListMarketplacesResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + ), + server.POST("/aiserver.v1.DashboardService/GetGlobalCommands", + server.Name("dashboard_get_global_commands"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_global_commands", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetGlobalCommandsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + ), + server.POST("/aiserver.v1.DashboardService/GetEffectiveUserPlugins", + server.Name("dashboard_get_effective_user_plugins"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_effective_user_plugins", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetEffectiveUserPluginsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + ), + server.POST("/aiserver.v1.DashboardService/RegisterMarketplaceAndPlugins", + server.Name("dashboard_register_marketplace_and_plugins"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_register_marketplace_and_plugins", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.RegisterMarketplaceAndPluginsResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + ), + server.POST("/aiserver.v1.DashboardService/GetCliDownloadUrl", + server.Name("dashboard_get_cli_download_url"), + server.ConnectUnary(), + server.Local(upstream.MockProtoAction(routeDeps, upstream.CompatRouteConfig{ + Name: "dashboard_get_cli_download_url", + StatusCode: http.StatusOK, + MockProtoType: "aiserver.v1.GetCliDownloadUrlResponse", + MockBuilder: upstream.EmptyMockBuilder, + })), + ), server.POST("/aiserver.v1.DashboardService/GetMe", server.Name("dashboard_get_me"), server.ConnectUnary(), diff --git a/internal/backend/server/upstream/action.go b/internal/backend/server/upstream/action.go index 2f090d3..70614cf 100644 --- a/internal/backend/server/upstream/action.go +++ b/internal/backend/server/upstream/action.go @@ -193,6 +193,18 @@ func DefaultModelNudgeMockBuilder(reqCtx *RequestContext) (map[string]any, error return buildDefaultModelNudgeDataPayload(reqCtx) } +func UsableModelsMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return buildUsableModelsPayload(reqCtx) +} + +func DefaultModelForCliMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return buildDefaultModelForCliPayload(reqCtx) +} + +func DefaultModelMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return buildDefaultModelPayload(reqCtx) +} + func BootstrapStatsigMockBuilder(reqCtx *RequestContext) (map[string]any, error) { return buildBootstrapStatsigPayload(reqCtx) } @@ -213,6 +225,18 @@ func DashboardManagedSkillsMockBuilder(reqCtx *RequestContext) (map[string]any, return buildDashboardManagedSkillsPayload(reqCtx) } +// EmptyMockBuilder возвращает пустой proto-ответ для ручек, где клиенту +// достаточно успешного "пусто": нет team-настроек, нет репозиториев, +// нет маркетплейсов/плагинов/команд, телеметрия принята без обработки. +func EmptyMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return map[string]any{}, nil +} + +// SubmitLogsMockBuilder подтверждает приём логов телеметрии без обработки. +func SubmitLogsMockBuilder(reqCtx *RequestContext) (map[string]any, error) { + return map[string]any{"success": true}, nil +} + func DashboardGetMeMockBuilder(reqCtx *RequestContext) (map[string]any, error) { return buildDashboardGetMePayload(reqCtx) } diff --git a/internal/backend/server/upstream/client.go b/internal/backend/server/upstream/client.go index 552dd76..dfedcd8 100644 --- a/internal/backend/server/upstream/client.go +++ b/internal/backend/server/upstream/client.go @@ -14,6 +14,7 @@ import ( "strings" "time" + "cursor/gen/agentv1" "cursor/gen/aiserverv1" "cursor/internal/logger" "cursor/internal/netproxy" @@ -426,6 +427,30 @@ func newProtoMessage(typeName string) (proto.Message, error) { return &aiserverv1.GetUsageLimitStatusAndActiveGrantsResponse{}, nil case "aiserver.v1.IsOnNewPricingResponse": return &aiserverv1.IsOnNewPricingResponse{}, nil + case "aiserver.v1.GetTeamAdminSettingsResponse": + return &aiserverv1.GetTeamAdminSettingsResponse{}, nil + case "aiserver.v1.GetTeamReposResponse": + return &aiserverv1.GetTeamReposResponse{}, nil + case "aiserver.v1.ListMarketplacesResponse": + return &aiserverv1.ListMarketplacesResponse{}, nil + case "aiserver.v1.GetUsableModelsResponse": + return &agentv1.GetUsableModelsResponse{}, nil + case "aiserver.v1.GetDefaultModelForCliResponse": + return &agentv1.GetDefaultModelForCliResponse{}, nil + case "aiserver.v1.GetDefaultModelResponse": + return &aiserverv1.GetDefaultModelResponse{}, nil + case "aiserver.v1.GetGlobalCommandsResponse": + return &aiserverv1.GetGlobalCommandsResponse{}, nil + case "aiserver.v1.GetEffectiveUserPluginsResponse": + return &aiserverv1.GetEffectiveUserPluginsResponse{}, nil + case "aiserver.v1.RegisterMarketplaceAndPluginsResponse": + return &aiserverv1.RegisterMarketplaceAndPluginsResponse{}, nil + case "aiserver.v1.GetCliDownloadUrlResponse": + return &aiserverv1.GetCliDownloadUrlResponse{}, nil + case "aiserver.v1.SubmitLogsResponse": + return &aiserverv1.SubmitLogsResponse{}, nil + case "aiserver.v1.TrackEventsResponse": + return &aiserverv1.TrackEventsResponse{}, nil default: return nil, fmt.Errorf("unsupported proto message type %q", typeName) } diff --git a/internal/backend/server/upstream/mocks.go b/internal/backend/server/upstream/mocks.go index c33dcfc..cc61898 100644 --- a/internal/backend/server/upstream/mocks.go +++ b/internal/backend/server/upstream/mocks.go @@ -17,6 +17,13 @@ const ( modelRuntimeThinkingEffortParameterID = "thinking_effort" + // localPathEncryptionKey — стабильный ключ шифрования путей для индексации + // репозитория, который cursor-agent CLI запрашивает через GetServerConfig + // (indexingConfig.default{User,Team}PathEncryptionKey). Без него CLI + // не может инициализировать repo identity в git-воркспейсе и вызовы + // файловых инструментов падают с "[unimplemented] HTTP 404". + localPathEncryptionKey = "6f6e63652d6c6f63616c2d706174682d656e6372797074696f6e2d6b6579" + localUltraMembershipType = "ultra" localUltraPaymentID = "local_ultra" localUltraSubscriptionStatus = "active" @@ -424,9 +431,13 @@ func buildServerTimePayload(*RequestContext) (map[string]any, error) { func buildServerConfigPayload(*RequestContext) (map[string]any, error) { return map[string]any{ - "configVersion": "local_cli_sandbox_defaults_disabled_v2", - // "http2Config": "HTTP2_CONFIG_FORCE_ALL_DISABLED", + "configVersion": "local_cli_sandbox_defaults_disabled_v2", + "http2Config": "HTTP2_CONFIG_FORCE_ALL_DISABLED", "cliSandboxDefaultEnabled": true, + "indexingConfig": map[string]any{ + "defaultUserPathEncryptionKey": localPathEncryptionKey, + "defaultTeamPathEncryptionKey": localPathEncryptionKey, + }, }, nil } @@ -440,6 +451,7 @@ func buildAvailableModelsPayload(reqCtx *RequestContext) (map[string]any, error) if len(modelRefs) > 0 { defaultModel = modelRefs[0] } + modelEntries := buildAvailableModelEntries(adapters) return map[string]any{ "backgroundComposerModelConfig": map[string]any{ "bestOfNDefaultModels": append([]string(nil), modelRefs...), @@ -459,7 +471,7 @@ func buildAvailableModelsPayload(reqCtx *RequestContext) (map[string]any, error) "defaultModel": defaultModel, }, "disableUnusedModelsAfterNHours": availableModelsDisableUnusedHours, - "models": buildAvailableModelEntries(adapters), + "models": modelEntries, "planExecutionModelConfig": map[string]any{ "defaultModel": defaultModel, "fallbackModels": append([]string(nil), modelRefs...), @@ -486,6 +498,35 @@ func buildDefaultModelNudgeDataPayload(reqCtx *RequestContext) (map[string]any, }, nil } +func buildUsableModelsPayload(reqCtx *RequestContext) (map[string]any, error) { + adapters, err := loadConfiguredModelAdapters(reqCtx) + if err != nil { + return nil, err + } + return map[string]any{"models": buildCLIModelDetails(adapters)}, nil +} + +func buildDefaultModelForCliPayload(reqCtx *RequestContext) (map[string]any, error) { + adapters, err := loadConfiguredModelAdapters(reqCtx) + if err != nil { + return nil, err + } + models := buildCLIModelDetails(adapters) + if len(models) == 0 { + return map[string]any{"model": map[string]any{}}, nil + } + return map[string]any{"model": models[0]}, nil +} + +func buildDefaultModelPayload(reqCtx *RequestContext) (map[string]any, error) { + adapters, err := loadConfiguredModelAdapters(reqCtx) + if err != nil { + return nil, err + } + defaultModel := firstModelAdapterRef(adapters) + return map[string]any{"model": defaultModel, "thinkingModel": defaultModel}, nil +} + func buildBootstrapStatsigPayload(reqCtx *RequestContext) (map[string]any, error) { generatedAtMs := uint64(time.Now().UnixMilli()) authID := resolveBootstrapStatsigAuthID(reqCtx) @@ -679,6 +720,25 @@ func buildAvailableModelEntries(adapters []legacyruntime.ModelAdapterConfig) []m return output } +func buildCLIModelDetails(adapters []legacyruntime.ModelAdapterConfig) []map[string]any { + models := make([]map[string]any, 0, len(adapters)) + for _, adapter := range adapters { + channelID := strings.TrimSpace(adapter.ID) + if channelID == "" { + continue + } + models = append(models, map[string]any{ + "modelId": channelID, + "displayModelId": channelID, + "apiKeyCredentials": map[string]any{ + "apiKey": strings.TrimSpace(adapter.APIKey), + "baseUrl": strings.TrimSpace(adapter.BaseURL), + }, + }) + } + return models +} + func buildThinkingEffortParameterDefinitions(adapterType string) []map[string]any { values := thinkingEffortValuesForAdapter(adapterType) options := make([]map[string]any, 0, len(values)) @@ -823,6 +883,16 @@ func collectModelAdapterRefs(adapters []legacyruntime.ModelAdapterConfig) []stri return output } +// firstModelAdapterRef возвращает канал первого адаптера или пустую строку, +// если ни один адаптер не сконфигурирован. +func firstModelAdapterRef(adapters []legacyruntime.ModelAdapterConfig) string { + refs := collectModelAdapterRefs(adapters) + if len(refs) == 0 { + return "" + } + return refs[0] +} + func resolveBootstrapStatsigAuthID(reqCtx *RequestContext) string { if reqCtx != nil { if authID := authIDFromBearer(reqCtx.Headers.Get("authorization")); authID != "" { diff --git a/internal/backend/server/upstream/mocks_test.go b/internal/backend/server/upstream/mocks_test.go index 5c878b9..48c4fa6 100644 --- a/internal/backend/server/upstream/mocks_test.go +++ b/internal/backend/server/upstream/mocks_test.go @@ -2,9 +2,55 @@ package upstream import ( "encoding/json" + "reflect" "testing" + + "cursor/gen/agentv1" + legacyruntime "cursor/internal/runtime" + + "google.golang.org/protobuf/proto" ) +func TestBuildCLIModelDetailsPreservesChannelCredentials(t *testing.T) { + adapters := []legacyruntime.ModelAdapterConfig{ + {ID: " channel-a ", ModelID: "model-a", APIKey: "provider-secret-a", BaseURL: "https://provider-a.example/v1"}, + {ID: "channel-b", ModelID: "model-a"}, + {ID: "", ModelID: "model-c"}, + } + + got := buildCLIModelDetails(adapters) + want := []map[string]any{ + {"modelId": "channel-a", "displayModelId": "channel-a", "apiKeyCredentials": map[string]any{"apiKey": "provider-secret-a", "baseUrl": "https://provider-a.example/v1"}}, + {"modelId": "channel-b", "displayModelId": "channel-b", "apiKeyCredentials": map[string]any{"apiKey": "", "baseUrl": ""}}, + } + if !reflect.DeepEqual(got, want) { + t.Fatalf("build CLI model details: got %v, want %v", got, want) + } +} + +func TestEncodeCLIModelsUsesAgentModelDetailsWireFormat(t *testing.T) { + payload := map[string]any{"models": buildCLIModelDetails([]legacyruntime.ModelAdapterConfig{{ID: "channel-a", APIKey: "provider-secret", BaseURL: "https://provider.example/v1"}})} + encoded, err := encodeMockProto("aiserver.v1.GetUsableModelsResponse", payload) + if err != nil { + t.Fatalf("encode CLI models: %v", err) + } + + response := &agentv1.GetUsableModelsResponse{} + if err := proto.Unmarshal(encoded, response); err != nil { + t.Fatalf("decode CLI models with agent proto: %v", err) + } + if len(response.Models) != 1 { + t.Fatalf("decoded model count: got %d, want 1", len(response.Models)) + } + model := response.Models[0] + if model.GetModelId() != "channel-a" || model.GetDisplayModelId() != "channel-a" { + t.Fatalf("decoded channel IDs: model=%q display=%q", model.GetModelId(), model.GetDisplayModelId()) + } + if credentials := model.GetApiKeyCredentials(); credentials == nil || credentials.GetApiKey() != "provider-secret" || credentials.GetBaseUrl() != "https://provider.example/v1" { + t.Fatalf("decoded relay credentials: %#v", credentials) + } +} + func TestBuildBootstrapStatsigConfigJSONDisablesAlwaysLocalDecompositionGate(t *testing.T) { payload, err := buildBootstrapStatsigConfigJSON(12345, "test-auth-id") if err != nil { diff --git a/release-notes.md b/release-notes.md index 1c3a92a..690b827 100644 --- a/release-notes.md +++ b/release-notes.md @@ -8,11 +8,6 @@ QQ交流群: Tg群组: https://t.me/cursor_byok -- 修复commands无法识别的问题 @DedSecer -- 修复 summarize 指令(主动压缩) @DedSecer -- 修复 MiniMax 禁止 thinking 的问题 @Octopus -- 支持Fork message (需要新开对话) @DedSecer -- 支持 @关联对话 @DedSecer -- 支持插件市场(须登录你的任意账号) @aike1202 -- 新增支持 cursor debugger 调试器, 用法请看.agents/skills/coding-guidance/SKILL.md +- 支持cursor-cli +- 修复对话中错误可能导致的消失问题