From 4e9335d82f0f7a655e02d459ce4716d19b1d0e46 Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 5 Aug 2026 01:46:47 +0800 Subject: [PATCH 1/2] 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. --- internal/backend/forwarder/projector.go | 164 +----------------- .../backend/forwarder/projector_fork_test.go | 54 ++++++ 2 files changed, 58 insertions(+), 160 deletions(-) create mode 100644 internal/backend/forwarder/projector_fork_test.go diff --git a/internal/backend/forwarder/projector.go b/internal/backend/forwarder/projector.go index 5726f8f..ed6dcb9 100644 --- a/internal/backend/forwarder/projector.go +++ b/internal/backend/forwarder/projector.go @@ -506,158 +506,10 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers if structuredState.HasTodos { state.Todos = encodeConversationTodoBytes(structuredState.Todos) } - grouped := make(map[int64][]HistoryEntry) - order := make([]int64, 0, conversation.NextTurnSeq) - for _, entry := range checkpointProjectionEntries(conversation.Entries) { - if entry.TurnSeq <= 0 { - continue - } - if _, ok := grouped[entry.TurnSeq]; !ok { - order = append(order, entry.TurnSeq) - } - grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry) - } - - for _, turnSeq := range order { - entries := grouped[turnSeq] - var rawUserMessage []byte - var turnRequestID string - steps := make([][]byte, 0, len(entries)) - seenToolCalls := make(map[string]struct{}) - openToolCalls := make(map[string]struct{}) - for _, entry := range entries { - if turnRequestID == "" { - turnRequestID = strings.TrimSpace(entry.RequestID) - } - switch strings.TrimSpace(entry.Kind) { - case "user_message": - userMessage := &agentv1.UserMessage{} - if err := protojson.Unmarshal(entry.Payload, userMessage); err != nil { - return nil, fmt.Errorf("decode checkpoint user_message: %w", err) - } - payload, err := proto.Marshal(userMessage) - if err != nil { - return nil, err - } - rawUserMessage = payload - case "assistant_text": - var payload assistantTextPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if strings.TrimSpace(payload.Text) == "" && strings.TrimSpace(payload.ReasoningContent) != "" && len(openToolCalls) > 0 { - continue - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - if strings.TrimSpace(payload.Text) == "" { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_AssistantMessage{ - AssistantMessage: &agentv1.AssistantMessage{Text: strings.TrimSpace(payload.Text)}, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - case "tool_call": - var payload toolCallEntryPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - toolCall := &agentv1.ToolCall{} - if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { - return nil, err - } - if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ToolCall{ - ToolCall: toolCall, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { - seenToolCalls[toolCallID] = struct{}{} - openToolCalls[toolCallID] = struct{}{} - } - case "tool_result": - var payload toolResultEntryPayload - if err := json.Unmarshal(entry.Payload, &payload); err != nil { - return nil, err - } - if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" { - if _, ok := seenToolCalls[toolCallID]; ok { - delete(openToolCalls, toolCallID) - continue - } - } - if strings.TrimSpace(payload.ReasoningContent) != "" { - stepPayload, err := marshalThinkingStep(payload.ReasoningContent) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - if len(payload.ToolCall) == 0 { - continue - } - toolCall := &agentv1.ToolCall{} - if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil { - return nil, err - } - if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) { - continue - } - stepPayload, err := proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ToolCall{ - ToolCall: toolCall, - }, - }) - if err != nil { - return nil, err - } - steps = append(steps, stepPayload) - } - } - if len(rawUserMessage) == 0 && len(steps) == 0 { - continue - } - agentTurn := &agentv1.AgentConversationTurnStructure{ - UserMessage: rawUserMessage, - Steps: steps, - } - if turnRequestID != "" { - agentTurn.RequestId = &turnRequestID - } - turnPayload, err := proto.Marshal(&agentv1.ConversationTurnStructure{ - Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{ - AgentConversationTurn: agentTurn, - }, - }) - if err != nil { - return nil, err - } - state.Turns = append(state.Turns, turnPayload) - } + // 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. replayMessages, err := projector.ProjectPromptReplay(conversation) if err != nil { return nil, err @@ -688,14 +540,6 @@ func (projector *HistoryProjector) ProjectLegacyCheckpoint(conversation *Convers return state, nil } -func marshalThinkingStep(text string) ([]byte, error) { - return proto.Marshal(&agentv1.ConversationStep{ - Message: &agentv1.ConversationStep_ThinkingMessage{ - ThinkingMessage: &agentv1.ThinkingMessage{Text: text}, - }, - }) -} - func conversationTokenDetailsUsedTokens(conversation *ConversationFile) uint32 { if conversation == nil { return 0 diff --git a/internal/backend/forwarder/projector_fork_test.go b/internal/backend/forwarder/projector_fork_test.go new file mode 100644 index 0000000..dcc3ed4 --- /dev/null +++ b/internal/backend/forwarder/projector_fork_test.go @@ -0,0 +1,54 @@ +package forwarder + +import ( + "strings" + "testing" + + "google.golang.org/protobuf/encoding/protojson" + + "cursor/gen/agentv1" +) + +func TestProjectLegacyCheckpointKeepsForkContextWithoutDanglingTurnBlobs(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", "", ""), + }, + } + + state, err := NewHistoryProjector().ProjectLegacyCheckpoint(conversation) + if err != nil { + t.Fatalf("ProjectLegacyCheckpoint() error = %v", err) + } + if len(state.GetTurns()) != 0 { + t.Fatalf("ProjectLegacyCheckpoint() turns = %d, want 0 dangling inline blobs", len(state.GetTurns())) + } + + 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]) + } +} From d3adfffbd8c0a93a86aa633903b5d7bc7d51d031 Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 5 Aug 2026 02:09:45 +0800 Subject: [PATCH 2/2] 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. --- internal/backend/forwarder/actor.go | 7 + internal/backend/forwarder/broker.go | 8 + .../backend/forwarder/checkpoint_blobs.go | 238 +++++++++++++++++ .../forwarder/checkpoint_blobs_test.go | 208 +++++++++++++++ internal/backend/forwarder/events.go | 16 ++ internal/backend/forwarder/projector.go | 245 +++++++++++++++++- .../backend/forwarder/projector_fork_test.go | 96 ++++++- internal/backend/forwarder/service.go | 33 ++- internal/backend/forwarder/types.go | 10 + 9 files changed, 838 insertions(+), 23 deletions(-) create mode 100644 internal/backend/forwarder/checkpoint_blobs.go create mode 100644 internal/backend/forwarder/checkpoint_blobs_test.go 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