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
-
-
-
-
-
-
-
-
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
+- 修复对话中错误可能导致的消失问题