diff --git a/internal/backend/forwarder/actor.go b/internal/backend/forwarder/actor.go index 2614cc1..a712976 100644 --- a/internal/backend/forwarder/actor.go +++ b/internal/backend/forwarder/actor.go @@ -23,6 +23,7 @@ const ( TurnPhaseWaitingExternal TurnPhase = "waiting_external" TurnPhaseAwaitingUser TurnPhase = "awaiting_user" TurnPhaseCompacting TurnPhase = "compacting" + TurnPhaseCheckpointing TurnPhase = "checkpointing" TurnPhaseCompleted TurnPhase = "completed" TurnPhaseFailed TurnPhase = "failed" TurnPhaseCanceled TurnPhase = "canceled" @@ -66,6 +67,7 @@ const ( streamTimerNonStreamingRecovery streamTimerKind = "non_streaming_recovery" streamTimerShellForeground streamTimerKind = "shell_foreground" streamTimerShellTransportClose streamTimerKind = "shell_transport_close" + streamTimerCheckpointBlobs streamTimerKind = "checkpoint_blobs" streamTimerOrphanCancel streamTimerKind = "orphan_cancel" ) @@ -318,6 +320,9 @@ func (service *Service) handleStreamCommand(stream *ActiveStream, command stream case streamCommandCancel: return service.handleCancelIntent(command.Intent) case streamCommandMetadata: + if strings.TrimSpace(command.Intent.Kind) == "kv_result" { + return service.handleCheckpointBlobResult(stream, command.Intent.KVClientMessage) + } return service.handleMetadataIntent(command.Intent) case streamCommandExecResult: return service.handleExecResult(command.Intent) @@ -1003,6 +1008,8 @@ func (service *Service) handleTimerEvent(stream *ActiveStream, payload *streamTi return nil } return service.recoverShellWithoutTerminal(stream, current, shellRecoveryReasonTransportClosed) + case streamTimerCheckpointBlobs: + return service.handleCheckpointBlobTimeout(stream) case streamTimerOrphanCancel: stream.mu.Lock() subscriberCount := len(stream.Subscribers) diff --git a/internal/backend/forwarder/broker.go b/internal/backend/forwarder/broker.go index b88442c..7a0f1ce 100644 --- a/internal/backend/forwarder/broker.go +++ b/internal/backend/forwarder/broker.go @@ -76,6 +76,12 @@ 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 @@ -102,6 +108,8 @@ func (broker *StreamBroker) OpenStream(requestID string, conversationID string, 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, } diff --git a/internal/backend/forwarder/checkpoint_blobs.go b/internal/backend/forwarder/checkpoint_blobs.go new file mode 100644 index 0000000..c07c07f --- /dev/null +++ b/internal/backend/forwarder/checkpoint_blobs.go @@ -0,0 +1,238 @@ +package forwarder + +import ( + "encoding/hex" + "fmt" + "log" + "strings" + "time" + + "google.golang.org/protobuf/proto" + + "cursor/gen/agentv1" +) + +const checkpointBlobWriteTimeout = 5 * time.Second + +type pendingCheckpointBlobWrite struct { + requestID uint32 + blob CheckpointBlob +} + +func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion { + if completion == nil { + return nil + } + cloned := *completion + return &cloned +} + +func (service *Service) queueCheckpointProjection(stream *ActiveStream, projection *CheckpointProjection, completion *pendingTurnCompletion) error { + if service == nil || stream == nil || projection == nil || projection.State == nil { + return nil + } + state, ok := proto.Clone(projection.State).(*agentv1.ConversationStateStructure) + if !ok || state == nil { + return fmt.Errorf("clone checkpoint state") + } + + stream.mu.Lock() + if stream.PendingCheckpointBlobWrites == nil { + stream.PendingCheckpointBlobWrites = make(map[uint32]string) + } + if stream.ConfirmedCheckpointBlobs == nil { + stream.ConfirmedCheckpointBlobs = make(map[string]struct{}) + } + if completion == nil && stream.PendingCheckpoint != nil { + completion = stream.PendingCheckpoint.Completion + } + required := make(map[string]struct{}, len(projection.Blobs)) + pendingKeys := make(map[string]struct{}, len(stream.PendingCheckpointBlobWrites)) + for _, key := range stream.PendingCheckpointBlobWrites { + pendingKeys[key] = struct{}{} + } + toWrite := make([]pendingCheckpointBlobWrite, 0, len(projection.Blobs)) + for _, blob := range projection.Blobs { + key := string(blob.ID) + if key == "" { + continue + } + required[key] = struct{}{} + if _, confirmed := stream.ConfirmedCheckpointBlobs[key]; confirmed { + continue + } + if _, pending := pendingKeys[key]; pending { + continue + } + stream.NextCheckpointBlobRequestID++ + if stream.NextCheckpointBlobRequestID == 0 { + stream.NextCheckpointBlobRequestID++ + } + requestID := stream.NextCheckpointBlobRequestID + stream.PendingCheckpointBlobWrites[requestID] = key + pendingKeys[key] = struct{}{} + toWrite = append(toWrite, pendingCheckpointBlobWrite{requestID: requestID, blob: blob}) + } + stream.PendingCheckpoint = &pendingCheckpointPublish{ + State: state, + Required: required, + Completion: clonePendingTurnCompletion(completion), + } + if completion != nil { + stream.Phase = TurnPhaseCheckpointing + } + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + + for _, write := range toWrite { + if err := service.broker.Publish(stream.RequestID, StreamEvent{ + Message: buildSetCheckpointBlobMessage(write.requestID, write.blob), + }); err != nil { + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf("publish checkpoint blob: %w", err)) + } + } + if service.checkpointProjectionReady(stream) { + return service.publishReadyCheckpoint(stream) + } + service.scheduleStreamTimer( + stream, + providerTimerKey(streamTimerCheckpointBlobs, ""), + checkpointBlobWriteTimeout, + streamTimerCheckpointBlobs, + "", + 0, + "checkpoint blob write timeout", + ) + return nil +} + +func (service *Service) checkpointProjectionReady(stream *ActiveStream) bool { + if stream == nil { + return false + } + stream.mu.Lock() + defer stream.mu.Unlock() + if stream.PendingCheckpoint == nil { + return false + } + for key := range stream.PendingCheckpoint.Required { + if _, confirmed := stream.ConfirmedCheckpointBlobs[key]; !confirmed { + return false + } + } + return true +} + +func (service *Service) handleCheckpointBlobResult(stream *ActiveStream, message *agentv1.KvClientMessage) error { + if service == nil || stream == nil || message == nil || message.GetSetBlobResult() == nil { + return nil + } + stream.mu.Lock() + key, ok := stream.PendingCheckpointBlobWrites[message.GetId()] + if ok { + delete(stream.PendingCheckpointBlobWrites, message.GetId()) + } + required := false + if ok && stream.PendingCheckpoint != nil { + _, required = stream.PendingCheckpoint.Required[key] + } + if ok && message.GetSetBlobResult().GetError() == nil { + stream.ConfirmedCheckpointBlobs[key] = struct{}{} + } + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + if !ok { + return nil + } + if blobErr := message.GetSetBlobResult().GetError(); blobErr != nil && required { + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf( + "client rejected checkpoint blob %s: %s", + hex.EncodeToString([]byte(key)), + firstNonEmpty(strings.TrimSpace(blobErr.GetMessage()), "unknown error"), + )) + } + if service.checkpointProjectionReady(stream) { + return service.publishReadyCheckpoint(stream) + } + return nil +} + +func (service *Service) publishReadyCheckpoint(stream *ActiveStream) error { + if service == nil || stream == nil { + return nil + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + if pending == nil { + stream.mu.Unlock() + return nil + } + for key := range pending.Required { + if _, confirmed := stream.ConfirmedCheckpointBlobs[key]; !confirmed { + stream.mu.Unlock() + return nil + } + } + stream.PendingCheckpoint = nil + state := pending.State + completion := clonePendingTurnCompletion(pending.Completion) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil { + if completion != nil { + log.Printf("forwarder checkpoint publish skipped before successful terminal request_id=%s err=%v", stream.RequestID, err) + return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion) + } + return err + } + if completion != nil { + return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion) + } + return nil +} + +func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error { + if stream == nil { + return nil + } + stream.mu.Lock() + pendingCount := len(stream.PendingCheckpointBlobWrites) + stream.mu.Unlock() + return service.finishAfterCheckpointSyncFailure(stream, fmt.Errorf("%d checkpoint blob writes timed out", pendingCount)) +} + +func (service *Service) finishAfterCheckpointSyncFailure(stream *ActiveStream, cause error) error { + if stream == nil { + return nil + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + stream.PendingCheckpoint = nil + stream.PendingCheckpointBlobWrites = make(map[uint32]string) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if cause != nil { + log.Printf("forwarder checkpoint blob sync skipped request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause) + } + if pending != nil && pending.Completion != nil { + return service.finishSuccessfulTurnAfterCheckpoint(stream, *pending.Completion) + } + return nil +} + +func (service *Service) discardPendingCheckpoint(stream *ActiveStream, reason string) { + if stream == nil { + return + } + stream.mu.Lock() + stream.PendingCheckpoint = nil + stream.PendingCheckpointBlobWrites = make(map[uint32]string) + stream.UpdatedAt = time.Now().UTC() + stream.mu.Unlock() + clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, "")) + if strings.TrimSpace(reason) != "" { + log.Printf("forwarder pending checkpoint discarded request_id=%s conversation_id=%s reason=%s", stream.RequestID, stream.ConversationID, strings.TrimSpace(reason)) + } +} diff --git a/internal/backend/forwarder/checkpoint_blobs_test.go b/internal/backend/forwarder/checkpoint_blobs_test.go new file mode 100644 index 0000000..9ada5d3 --- /dev/null +++ b/internal/backend/forwarder/checkpoint_blobs_test.go @@ -0,0 +1,208 @@ +package forwarder + +import ( + "testing" + + "google.golang.org/protobuf/encoding/protojson" + + "cursor/gen/agentv1" +) + +func TestCheckpointBlobSyncPublishesCheckpointAfterAcknowledgements(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + if err := service.queueCheckpointProjection(stream, projection, nil); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + events := readCheckpointTestEvents(t, service, stream) + if len(events) != len(projection.Blobs) { + t.Fatalf("events before ACK = %d, want %d Blob writes", len(events), len(projection.Blobs)) + } + for _, event := range events { + if event.Message.GetKvServerMessage().GetSetBlobArgs() == nil { + t.Fatalf("event before ACK = %#v, want set_blob_args", event.Message) + } + } + + acknowledgeCheckpointBlobs(t, service, stream) + events = readCheckpointTestEvents(t, service, stream) + checkpoint := events[len(events)-1].Message.GetConversationCheckpointUpdate() + if checkpoint == nil || len(checkpoint.GetTurns()) != 1 { + t.Fatalf("last event checkpoint = %#v, want one Blob-backed turn", checkpoint) + } +} + +func TestCheckpointBlobSyncPublishesCheckpointBeforeSuccessfulTerminal(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + completion := &pendingTurnCompletion{ + RequestID: stream.RequestID, + Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7}, + } + if err := service.queueCheckpointProjection(stream, projection, completion); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + acknowledgeCheckpointBlobs(t, service, stream) + + events := readCheckpointTestEvents(t, service, stream) + checkpointIndex, turnEndedIndex, endIndex := -1, -1, -1 + for index, event := range events { + switch { + case event.Message.GetConversationCheckpointUpdate() != nil: + checkpointIndex = index + case event.Message.GetInteractionUpdate().GetTurnEnded() != nil: + turnEndedIndex = index + case event.End: + endIndex = index + } + } + if checkpointIndex < 0 || turnEndedIndex <= checkpointIndex || endIndex <= turnEndedIndex { + t.Fatalf("terminal order checkpoint=%d turn_ended=%d end=%d", checkpointIndex, turnEndedIndex, endIndex) + } +} + +func TestCheckpointBlobTimeoutDoesNotFailSuccessfulTurn(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + completion := &pendingTurnCompletion{ + RequestID: stream.RequestID, + Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7}, + } + if err := service.queueCheckpointProjection(stream, projection, completion); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + if err := service.handleCheckpointBlobTimeout(stream); err != nil { + t.Fatalf("handleCheckpointBlobTimeout() error = %v", err) + } + + events := readCheckpointTestEvents(t, service, stream) + var checkpoint, turnEnded, successfulEnd bool + for _, event := range events { + checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil + turnEnded = turnEnded || event.Message.GetInteractionUpdate().GetTurnEnded() != nil + successfulEnd = successfulEnd || event.End && event.TerminalErrorCode == "" + } + if checkpoint || !turnEnded || !successfulEnd { + t.Fatalf("timeout events checkpoint=%v turn_ended=%v successful_end=%v", checkpoint, turnEnded, successfulEnd) + } +} + +func TestCancellationDiscardsPendingCheckpointAndIgnoresLateAcknowledgements(t *testing.T) { + service, stream, projection := testCheckpointBlobProjection(t) + if err := service.queueCheckpointProjection(stream, projection, nil); err != nil { + t.Fatalf("queueCheckpointProjection() error = %v", err) + } + stream.mu.Lock() + requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites)) + for requestID := range stream.PendingCheckpointBlobWrites { + requestIDs = append(requestIDs, requestID) + } + stream.mu.Unlock() + if err := service.handleCancelIntent(InboundIntent{ + Kind: "cancel", + RequestID: stream.RequestID, + CancelReason: "user stopped", + }); err != nil { + t.Fatalf("handleCancelIntent() error = %v", err) + } + for _, requestID := range requestIDs { + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + t.Fatalf("late ACK %d error = %v", requestID, err) + } + } + + events := readCheckpointTestEvents(t, service, stream) + var checkpoint, canceledEnd bool + for _, event := range events { + checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil + canceledEnd = canceledEnd || event.End && event.TerminalErrorCode == "canceled" + } + stream.mu.Lock() + pending := stream.PendingCheckpoint + stream.mu.Unlock() + if checkpoint || !canceledEnd || pending != nil { + t.Fatalf("cancel events checkpoint=%v canceled_end=%v pending=%v", checkpoint, canceledEnd, pending != nil) + } +} + +func testCheckpointBlobProjection(t *testing.T) (*Service, *ActiveStream, *CheckpointProjection) { + t.Helper() + broker := NewStreamBroker() + service := &Service{ + store: NewConversationFileStore(t.TempDir()), + projector: NewHistoryProjector(), + broker: broker, + } + stream, err := broker.OpenStream( + "request-1", "conversation-1", 1, "default", "default", + agentv1.AgentMode_AGENT_MODE_AGENT, "hello", + ) + if err != nil { + t.Fatalf("OpenStream() error = %v", err) + } + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 2, + NextEntrySeq: 3, + TokenDetailsMaxTokens: projectedConversationMaxTokens, + Entries: []HistoryEntry{ + testCheckpointUserEntry(t), + newAssistantTextEntry(1, "request-1", "hi", "", ""), + }, + } + projection, err := service.projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("ProjectCheckpointProjection() error = %v", err) + } + if err := service.replaceCheckpointConversation(stream, conversation); err != nil { + t.Fatalf("replaceCheckpointConversation() error = %v", err) + } + return service, stream, projection +} + +func testCheckpointUserEntry(t *testing.T) HistoryEntry { + t.Helper() + payload, err := protojson.Marshal(&agentv1.UserMessage{Text: "hello", MessageId: "message-1"}) + if err != nil { + t.Fatalf("marshal user message: %v", err) + } + return HistoryEntry{Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: payload} +} + +func acknowledgeCheckpointBlobs(t *testing.T, service *Service, stream *ActiveStream) { + t.Helper() + for { + stream.mu.Lock() + requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites)) + for requestID := range stream.PendingCheckpointBlobWrites { + requestIDs = append(requestIDs, requestID) + } + stream.mu.Unlock() + if len(requestIDs) == 0 { + return + } + for _, requestID := range requestIDs { + if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{ + Id: requestID, + Message: &agentv1.KvClientMessage_SetBlobResult{ + SetBlobResult: &agentv1.SetBlobResult{}, + }, + }); err != nil { + t.Fatalf("handleCheckpointBlobResult(%d) error = %v", requestID, err) + } + } + } +} + +func readCheckpointTestEvents(t *testing.T, service *Service, stream *ActiveStream) []StreamEvent { + t.Helper() + events, err := service.broker.ReadFromCursor(stream.RequestID, 0) + if err != nil { + t.Fatalf("ReadFromCursor() error = %v", err) + } + return events +} diff --git a/internal/backend/forwarder/events.go b/internal/backend/forwarder/events.go index 7b02404..ef69776 100644 --- a/internal/backend/forwarder/events.go +++ b/internal/backend/forwarder/events.go @@ -245,6 +245,22 @@ func buildCheckpointMessage(state *agentv1.ConversationStateStructure) *agentv1. } } +func buildSetCheckpointBlobMessage(id uint32, blob CheckpointBlob) *agentv1.AgentServerMessage { + return &agentv1.AgentServerMessage{ + Message: &agentv1.AgentServerMessage_KvServerMessage{ + KvServerMessage: &agentv1.KvServerMessage{ + Id: id, + Message: &agentv1.KvServerMessage_SetBlobArgs{ + SetBlobArgs: &agentv1.SetBlobArgs{ + BlobId: append([]byte(nil), blob.ID...), + BlobData: append([]byte(nil), blob.Data...), + }, + }, + }, + }, + } +} + // buildExecAbortMessage 构造对客户端执行桥的 abort 控制消息。 func buildExecAbortMessage(pending runtimecore.PendingExec) *agentv1.AgentServerMessage { return &agentv1.AgentServerMessage{ diff --git a/internal/backend/forwarder/projector.go b/internal/backend/forwarder/projector.go index ed6dcb9..7095860 100644 --- a/internal/backend/forwarder/projector.go +++ b/internal/backend/forwarder/projector.go @@ -2,6 +2,7 @@ package forwarder import ( + "crypto/sha256" "encoding/json" "fmt" "strings" @@ -19,6 +20,51 @@ const projectedConversationMaxTokens = 130000 type HistoryProjector struct { } +type CheckpointBlob struct { + ID []byte + Data []byte +} + +type CheckpointProjection struct { + State *agentv1.ConversationStateStructure + Blobs []CheckpointBlob +} + +type checkpointBlobGraph struct { + blobs map[[sha256.Size]byte][]byte + order [][sha256.Size]byte +} + +func newCheckpointBlobGraph() *checkpointBlobGraph { + return &checkpointBlobGraph{blobs: make(map[[sha256.Size]byte][]byte)} +} + +func (graph *checkpointBlobGraph) add(data []byte) []byte { + if graph == nil || len(data) == 0 { + return nil + } + id := sha256.Sum256(data) + if _, exists := graph.blobs[id]; !exists { + graph.blobs[id] = append([]byte(nil), data...) + graph.order = append(graph.order, id) + } + return append([]byte(nil), id[:]...) +} + +func (graph *checkpointBlobGraph) list() []CheckpointBlob { + if graph == nil || len(graph.order) == 0 { + return nil + } + blobs := make([]CheckpointBlob, 0, len(graph.order)) + for _, id := range graph.order { + blobs = append(blobs, CheckpointBlob{ + ID: append([]byte(nil), id[:]...), + Data: append([]byte(nil), graph.blobs[id]...), + }) + } + return blobs +} + // NewHistoryProjector 创建 history 投影器。 func NewHistoryProjector() *HistoryProjector { return &HistoryProjector{} @@ -475,6 +521,16 @@ func isHistoricalReplayToolResult(conversation *ConversationFile, entry HistoryE // ProjectLegacyCheckpoint 按需从 JSON history 投影出兼容旧客户端的 checkpoint 结构。 func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *ConversationFile) (*agentv1.ConversationStateStructure, error) { + projection, err := projector.ProjectCheckpointProjection(conversation) + if err != nil || projection == nil { + return nil, err + } + return projection.State, nil +} + +// ProjectCheckpointProjection 同时返回 checkpoint 状态及其引用的内容寻址 Blob。 +func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *ConversationFile) (*CheckpointProjection, error) { + blobs := newCheckpointBlobGraph() state := &agentv1.ConversationStateStructure{ TokenDetails: &agentv1.ConversationTokenDetails{ UsedTokens: conversationTokenDetailsUsedTokens(conversation), @@ -488,7 +544,7 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers if conversation == nil { mode := agentv1.AgentMode_AGENT_MODE_AGENT state.Mode = &mode - return state, nil + return &CheckpointProjection{State: state}, nil } mode, err := parseModeAlias(conversation.Mode) if err != nil { @@ -506,10 +562,11 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers if structuredState.HasTodos { state.Todos = encodeConversationTodoBytes(structuredState.Todos) } - // Cursor 3.14 treats ConversationStateStructure.turns as content-addressed - // blob IDs. Inline protobuf turns make Fork Chat look up the serialized turn - // itself as an ID and fail with "Missing turn blob". Keep model-visible - // history in root_prompt_messages_json until checkpoint blob sync is safe. + turnIDs, err := projectCheckpointTurnBlobs(conversation, blobs) + if err != nil { + return nil, err + } + state.Turns = turnIDs replayMessages, err := projector.ProjectPromptReplay(conversation) if err != nil { return nil, err @@ -537,7 +594,183 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers return nil, err } state.RootPromptMessagesJson = rootPromptMessages - return state, nil + return &CheckpointProjection{State: state, Blobs: blobs.list()}, nil +} + +func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpointBlobGraph) ([][]byte, error) { + if conversation == nil || blobs == nil { + return nil, nil + } + grouped := make(map[int64][]HistoryEntry) + order := make([]int64, 0, conversation.NextTurnSeq) + for _, entry := range checkpointProjectionEntries(conversation.Entries) { + if entry.TurnSeq <= 0 { + continue + } + if _, ok := grouped[entry.TurnSeq]; !ok { + order = append(order, entry.TurnSeq) + } + grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry) + } + + turnIDs := make([][]byte, 0, len(order)) + for _, turnSeq := range order { + entries := grouped[turnSeq] + var userMessageID []byte + var turnRequestID string + stepIDs := make([][]byte, 0, len(entries)) + seenToolCalls := make(map[string]struct{}) + openToolCalls := make(map[string]struct{}) + for _, entry := range entries { + if turnRequestID == "" { + turnRequestID = strings.TrimSpace(entry.RequestID) + } + switch strings.TrimSpace(entry.Kind) { + case "user_message": + userMessage := &agentv1.UserMessage{} + if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil { + return nil, fmt.Errorf("decode checkpoint user_message: %w", err) + } + payload, err := proto.Marshal(userMessage) + if err != nil { + return nil, err + } + userMessageID = blobs.add(payload) + case "assistant_text": + var payload assistantTextPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if strings.TrimSpace(payload.Text) == "" && strings.TrimSpace(payload.ReasoningContent) != "" && len(openToolCalls) > 0 { + continue + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + if strings.TrimSpace(payload.Text) == "" { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_AssistantMessage{ + AssistantMessage: &agentv1.AssistantMessage{Text: strings.TrimSpace(payload.Text)}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + case "tool_call": + var payload toolCallEntryPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + toolCall := &agentv1.ToolCall{} + if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { + return nil, err + } + if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ToolCall{ToolCall: toolCall}, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { + seenToolCalls[toolCallID] = struct{}{} + openToolCalls[toolCallID] = struct{}{} + } + case "tool_result": + var payload toolResultEntryPayload + if err := json.Unmarshal(entry.Payload, &payload); err != nil { + return nil, err + } + if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { + if _, ok := seenToolCalls[toolCallID]; ok { + delete(openToolCalls, toolCallID) + continue + } + } + if strings.TrimSpace(payload.ReasoningContent) != "" { + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ThinkingMessage{ + ThinkingMessage: &agentv1.ThinkingMessage{Text: payload.ReasoningContent}, + }, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + if len(payload.ToolCall) == 0 { + continue + } + toolCall := &agentv1.ToolCall{} + if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { + return nil, err + } + if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { + continue + } + stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{ + Message: &agentv1.ConversationStep_ToolCall{ToolCall: toolCall}, + }) + if err != nil { + return nil, err + } + stepIDs = append(stepIDs, stepID) + } + } + if len(userMessageID) == 0 && len(stepIDs) == 0 { + continue + } + agentTurn := &agentv1.AgentConversationTurnStructure{ + UserMessage: userMessageID, + Steps: stepIDs, + } + if turnRequestID != "" { + agentTurn.RequestId = &turnRequestID + } + turnPayload, err := proto.Marshal(&agentv1.ConversationTurnStructure{ + Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ + AgentConversationTurn: agentTurn, + }, + }) + if err != nil { + return nil, err + } + turnIDs = append(turnIDs, blobs.add(turnPayload)) + } + return turnIDs, nil +} + +func addCheckpointStepBlob(blobs *checkpointBlobGraph, step *agentv1.ConversationStep) ([]byte, error) { + payload, err := proto.Marshal(step) + if err != nil { + return nil, err + } + return blobs.add(payload), nil } func conversationTokenDetailsUsedTokens(conversation *ConversationFile) uint32 { diff --git a/internal/backend/forwarder/projector_fork_test.go b/internal/backend/forwarder/projector_fork_test.go index dcc3ed4..56fc0af 100644 --- a/internal/backend/forwarder/projector_fork_test.go +++ b/internal/backend/forwarder/projector_fork_test.go @@ -1,15 +1,17 @@ package forwarder import ( + "crypto/sha256" "strings" "testing" "google.golang.org/protobuf/encoding/protojson" + "google.golang.org/protobuf/proto" "cursor/gen/agentv1" ) -func TestProjectLegacyCheckpointKeepsForkContextWithoutDanglingTurnBlobs(t *testing.T) { +func TestProjectCheckpointProjectionBuildsResolvableForkState(t *testing.T) { userPayload, err := protojson.Marshal(&agentv1.UserMessage{ Text: "parent question", MessageId: "message-1", @@ -30,12 +32,41 @@ func TestProjectLegacyCheckpointKeepsForkContextWithoutDanglingTurnBlobs(t *test }, } - state, err := NewHistoryProjector().ProjectLegacyCheckpoint(conversation) + projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation) if err != nil { - t.Fatalf("ProjectLegacyCheckpoint() error = %v", err) + t.Fatalf("ProjectCheckpointProjection() error = %v", err) } - if len(state.GetTurns()) != 0 { - t.Fatalf("ProjectLegacyCheckpoint() turns = %d, want 0 dangling inline blobs", len(state.GetTurns())) + 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) @@ -52,3 +83,58 @@ func TestProjectLegacyCheckpointKeepsForkContextWithoutDanglingTurnBlobs(t *test t.Fatalf("second imported message = %#v", messages[1]) } } + +func TestProjectCheckpointProjectionKeepsForkPointIsolatedFromLaterHistory(t *testing.T) { + firstUser, err := protojson.Marshal(&agentv1.UserMessage{Text: "first question", MessageId: "message-1"}) + if err != nil { + t.Fatalf("marshal first user message: %v", err) + } + conversation := &ConversationFile{ + ConversationID: "conversation-1", + RootConversationID: "conversation-1", + Mode: "agent", + NextTurnSeq: 2, + NextEntrySeq: 3, + TokenDetailsMaxTokens: projectedConversationMaxTokens, + Entries: []HistoryEntry{ + {Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: firstUser}, + newAssistantTextEntry(1, "request-1", "first answer", "", ""), + }, + } + projector := NewHistoryProjector() + midpoint, err := projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("midpoint projection: %v", err) + } + + secondUser, err := protojson.Marshal(&agentv1.UserMessage{Text: "second question", MessageId: "message-2"}) + if err != nil { + t.Fatalf("marshal second user message: %v", err) + } + appendEntriesInPlace(conversation, []HistoryEntry{ + {TurnSeq: 2, RequestID: "request-2", Role: "user", Kind: "user_message", Payload: secondUser}, + newAssistantTextEntry(2, "request-2", "second answer", "", ""), + }) + latest, err := projector.ProjectCheckpointProjection(conversation) + if err != nil { + t.Fatalf("latest projection: %v", err) + } + + midpointMessages, err := importedConversationStateModelMessages(midpoint.State) + if err != nil { + t.Fatalf("import midpoint messages: %v", err) + } + latestMessages, err := importedConversationStateModelMessages(latest.State) + if err != nil { + t.Fatalf("import latest messages: %v", err) + } + if len(midpoint.State.GetTurns()) != 1 || len(midpointMessages) != 2 { + t.Fatalf("midpoint turns=%d messages=%d, want 1 turn and 2 messages", len(midpoint.State.GetTurns()), len(midpointMessages)) + } + if len(latest.State.GetTurns()) != 2 || len(latestMessages) != 4 { + t.Fatalf("latest turns=%d messages=%d, want 2 turns and 4 messages", len(latest.State.GetTurns()), len(latestMessages)) + } + if midpointMessages[1].Content != "first answer" || latestMessages[3].Content != "second answer" { + t.Fatalf("fork snapshots are not isolated: midpoint=%#v latest=%#v", midpointMessages, latestMessages) + } +} diff --git a/internal/backend/forwarder/service.go b/internal/backend/forwarder/service.go index 5bb04ca..4cf5c5d 100644 --- a/internal/backend/forwarder/service.go +++ b/internal/backend/forwarder/service.go @@ -858,9 +858,7 @@ func (service *Service) handleCancelIntent(intent InboundIntent) error { }) } if hasCheckpoint { - if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil { - return err - } + service.discardPendingCheckpoint(stream, "checkpoint superseded by cancellation") } clearPendingProviderCompletion(stream) stream.mu.Lock() @@ -2122,9 +2120,15 @@ func (service *Service) completeSuccessfulTurn(stream *ActiveStream, completion err, ) } - if err := service.publishCheckpoint(requestID, conversationID); err != nil { - return err + return service.publishCheckpointWithCompletion(requestID, conversationID, &completion) +} + +func (service *Service) finishSuccessfulTurnAfterCheckpoint(stream *ActiveStream, completion pendingTurnCompletion) error { + if stream == nil { + return nil } + requestID := firstNonEmpty(strings.TrimSpace(completion.RequestID), strings.TrimSpace(stream.RequestID)) + usage := completion.Usage if err := service.broker.Publish(requestID, StreamEvent{ Message: buildTurnEndedMessage(usage.InputTokens, usage.OutputTokens, usage.CacheReadTokens, usage.CacheWriteTokens), }); err != nil { @@ -2151,7 +2155,11 @@ func (service *Service) failStreamIfNonTerminal(stream *ActiveStream, terminalCo } // publishCheckpoint 按当前内存会话镜像投影出 checkpoint,并广播给所有 RunSSE 订阅者。 -func (service *Service) publishCheckpoint(requestID string, _ string) error { +func (service *Service) publishCheckpoint(requestID string, conversationID string) error { + return service.publishCheckpointWithCompletion(requestID, conversationID, nil) +} + +func (service *Service) publishCheckpointWithCompletion(requestID string, _ string, completion *pendingTurnCompletion) error { stream, ok := service.broker.Get(requestID) if !ok || stream == nil { return fmt.Errorf("request is not active: %s", requestID) @@ -2160,15 +2168,16 @@ func (service *Service) publishCheckpoint(requestID string, _ string) error { if err != nil { return err } - state, err := service.projector.ProjectLegacyCheckpoint(conversation) + projection, err := service.projector.ProjectCheckpointProjection(conversation) if err != nil { return err } - state.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) - service.rewriteCheckpointTokenDetailsForClient(stream, conversation, state) - return service.broker.Publish(requestID, StreamEvent{ - Message: buildCheckpointMessage(state), - }) + if projection == nil || projection.State == nil { + return fmt.Errorf("checkpoint projection is empty") + } + projection.State.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions) + service.rewriteCheckpointTokenDetailsForClient(stream, conversation, projection.State) + return service.queueCheckpointProjection(stream, projection, completion) } func (service *Service) rewriteCheckpointTokenDetailsForClient(stream *ActiveStream, conversation *ConversationFile, state *agentv1.ConversationStateStructure) { diff --git a/internal/backend/forwarder/types.go b/internal/backend/forwarder/types.go index db78eee..7b33d20 100644 --- a/internal/backend/forwarder/types.go +++ b/internal/backend/forwarder/types.go @@ -163,6 +163,10 @@ type ActiveStream struct { ProviderUsage turnUsageSnapshot ProviderTerminalToolInvocation bool PendingCompaction *PendingCompaction + PendingCheckpointBlobWrites map[uint32]string + ConfirmedCheckpointBlobs map[string]struct{} + NextCheckpointBlobRequestID uint32 + PendingCheckpoint *pendingCheckpointPublish Backlog []StreamEvent Subscribers map[string]*StreamSubscriber @@ -219,6 +223,12 @@ type pendingTurnCompletion struct { Disposition pendingCompletionDisposition } +type pendingCheckpointPublish struct { + State *agentv1.ConversationStateStructure + Required map[string]struct{} + Completion *pendingTurnCompletion +} + type PendingCompaction struct { Trigger string ContextTokens int64