fix(forwarder): preserve Cursor fork context

Store checkpoint turns as validated blobs and restore prefetched parent turns so forked conversations retain their replay history and rewind prefix.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
DedSecer
2026-08-01 10:51:23 +08:00
co-authored by Cursor
parent 00f52166b2
commit 14b286a466
14 changed files with 1979 additions and 304 deletions
+241 -218
View File
@@ -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,158 +562,11 @@ 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)
turnIDs, err := projectCheckpointTurnBlobs(conversation, blobs)
if err != nil {
return nil, err
}
state.Turns = append(cloneByteSlices(conversation.ImportedTurnIDs), turnIDs...)
replayMessages, err := projector.ProjectPromptReplay(conversation)
if err != nil {
return nil, err
@@ -685,15 +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 marshalThinkingStep(text string) ([]byte, error) {
return proto.Marshal(&agentv1.ConversationStep{
Message: &agentv1.ConversationStep_ThinkingMessage{
ThinkingMessage: &agentv1.ThinkingMessage{Text: text},
},
})
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 {
@@ -1141,60 +1218,6 @@ 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
@@ -1231,7 +1254,7 @@ func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []pro
return filtered
}
func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte) []promptengine.Message {
func restoreImportedReplayUserMessages(messages []promptengine.Message, importedTurns [][]byte, blobs importedBlobStore) []promptengine.Message {
if len(messages) == 0 || len(importedTurns) == 0 {
return messages
}
@@ -1240,16 +1263,16 @@ func restoreImportedReplayUserMessages(messages []promptengine.Message, imported
if len(rawTurn) == 0 {
continue
}
turn := &agentv1.ConversationTurnStructure{}
if err := proto.Unmarshal(rawTurn, turn); err != nil {
turn, _, err := decodeImportedTurn(rawTurn, blobs)
if err != nil || turn == nil {
continue
}
agentTurn := turn.GetAgentConversationTurn()
if agentTurn == nil || len(agentTurn.GetUserMessage()) == 0 {
continue
}
userMessage := &agentv1.UserMessage{}
if err := proto.Unmarshal(agentTurn.GetUserMessage(), userMessage); err != nil {
userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs)
if err != nil {
continue
}
replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage)