Compare commits

..
Author SHA1 Message Date
leookun d3adfffbd8 Implement checkpoint blob handling in forwarder service
- Added support for checkpointing phases and blob management in the forwarder.
- Introduced new types and methods for handling checkpoint blobs, including queuing and publishing checkpoints.
- Enhanced the projector to build checkpoint projections with content-addressed blobs.
- Implemented tests to ensure proper checkpoint blob synchronization and handling of cancellation scenarios.
2026-08-05 02:09:45 +08:00
leookun 4e9335d82f Implement tests for ProjectLegacyCheckpoint to ensure proper handling of conversation state and message imports. This includes verifying that no dangling inline blobs are present and that the correct user and assistant messages are imported. 2026-08-05 01:46:47 +08:00
leokunandGitHub 374ff9c217 Merge pull request #259 from leookun/release/0.0.45
Release/0.0.45
2026-08-05 01:26:22 +08:00
leookun 06ab0d8dae chore(release): update version to 0.0.45 and fix conversation disappearance issue
- Updated version number to 0.0.45 across all relevant files.
- Fixed an issue that could cause conversations to disappear.
2026-08-05 01:25:39 +08:00
leookun 622a57fed7 Revert "Merge pull request #239 from DedSecer/fix/cursor-fork-context"
This reverts commit 697fa99245, reversing
changes made to 5742073be5.
2026-08-04 23:49:49 +08:00
leokunandGitHub 639c452a00 Merge pull request #252 from leookun/release/0.0.44
chore(release): update version to 0.0.44 and enhance release notes
2026-08-03 23:18:11 +08:00
23 changed files with 749 additions and 1710 deletions
+1 -1
View File
@@ -8,7 +8,7 @@ info:
description: "Cursor助手"
copyright: "© 2026, Cursor助手"
comments: "Cursor助手"
version: "0.0.44"
version: "0.0.45"
dev_mode:
root_path: .
+2 -2
View File
@@ -17,9 +17,9 @@
<key>CFBundlePackageType</key>
<string>APPL</string>
<key>CFBundleShortVersionString</key>
<string>0.0.44</string>
<string>0.0.45</string>
<key>CFBundleVersion</key>
<string>0.0.44</string>
<string>0.0.45</string>
<key>LSMinimumSystemVersion</key>
<string>12.0.0</string>
<key>LSUIElement</key>
+2 -2
View File
@@ -17,9 +17,9 @@
<key>CFBundlePackageType</key>
<string>APPL</string>
<key>CFBundleShortVersionString</key>
<string>0.0.44</string>
<string>0.0.45</string>
<key>CFBundleVersion</key>
<string>0.0.44</string>
<string>0.0.45</string>
<key>LSMinimumSystemVersion</key>
<string>12.0.0</string>
<key>LSUIElement</key>
+1 -1
View File
@@ -6,7 +6,7 @@
name: "Cursor助手"
arch: ${GOARCH}
platform: "linux"
version: "0.0.44"
version: "0.0.45"
section: "default"
priority: "extra"
maintainer: ${GIT_COMMITTER_NAME} <${GIT_COMMITTER_EMAIL}>
+2 -2
View File
@@ -1,10 +1,10 @@
{
"fixed": {
"file_version": "0.0.44"
"file_version": "0.0.45"
},
"info": {
"0000": {
"ProductVersion": "0.0.44",
"ProductVersion": "0.0.45",
"CompanyName": "Cursor助手",
"FileDescription": "Cursor助手",
"LegalCopyright": "© 2026, Cursor助手",
+1 -1
View File
@@ -14,7 +14,7 @@
!define INFO_PRODUCTNAME "Cursor助手"
!endif
!ifndef INFO_PRODUCTVERSION
!define INFO_PRODUCTVERSION "0.0.44"
!define INFO_PRODUCTVERSION "0.0.45"
!endif
!ifndef INFO_COPYRIGHT
!define INFO_COPYRIGHT "© 2026, Cursor助手"
+1 -1
View File
@@ -1,6 +1,6 @@
<?xml version="1.0" encoding="UTF-8" standalone="yes"?>
<assembly manifestVersion="1.0" xmlns="urn:schemas-microsoft-com:asm.v1" xmlns:asmv3="urn:schemas-microsoft-com:asm.v3">
<assemblyIdentity type="win32" name="com.cursor.wuxianxubei" version="0.0.44" processorArchitecture="*"/>
<assemblyIdentity type="win32" name="com.cursor.wuxianxubei" version="0.0.45" processorArchitecture="*"/>
<dependency>
<dependentAssembly>
<assemblyIdentity type="win32" name="Microsoft.Windows.Common-Controls" version="6.0.0.0" processorArchitecture="*" publicKeyToken="6595b64144ccf1df" language="*"/>
-383
View File
@@ -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))
}
@@ -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
}
+30 -24
View File
@@ -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
@@ -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))
}
}
@@ -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
}
-2
View File
@@ -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...)
@@ -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
}
+60 -6
View File
@@ -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)
@@ -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)
}
}
@@ -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 != "<conversation_summary>\nfirst turn summarized\n</conversation_summary>" {
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
}
+2 -23
View File
@@ -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
@@ -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
+24 -61
View File
@@ -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) {
+30 -25
View File
@@ -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
}
+5 -28
View File
@@ -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
+1
View File
@@ -9,4 +9,5 @@ Tg群组:
https://t.me/cursor_byok
- 支持cursor-cli
- 修复对话中错误可能导致的消失问题