mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-17 03:27:02 +08:00
Merge pull request #239 from DedSecer/fix/cursor-fork-context
fix(forwarder): preserve Cursor fork context
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -0,0 +1,383 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
const (
|
||||
checkpointBlobWriteTimeout = 10 * time.Second
|
||||
checkpointBlobCacheIdleTTL = 6 * time.Hour
|
||||
checkpointBlobCacheMaxConversations = 256
|
||||
)
|
||||
|
||||
type checkpointBlobCacheEntry struct {
|
||||
Confirmed map[string]struct{}
|
||||
LastAccess time.Time
|
||||
}
|
||||
|
||||
func checkpointBlobKey(id []byte) string {
|
||||
return string(id)
|
||||
}
|
||||
|
||||
func checkpointBlobHex(key string) string {
|
||||
return hex.EncodeToString([]byte(key))
|
||||
}
|
||||
|
||||
func (service *Service) confirmedCheckpointBlob(conversationID string, key string) bool {
|
||||
if service == nil || key == "" {
|
||||
return false
|
||||
}
|
||||
service.checkpointBlobMu.Lock()
|
||||
defer service.checkpointBlobMu.Unlock()
|
||||
conversationID = strings.TrimSpace(conversationID)
|
||||
entry := service.checkpointBlobs[conversationID]
|
||||
if entry == nil {
|
||||
return false
|
||||
}
|
||||
entry.LastAccess = time.Now().UTC()
|
||||
_, ok := entry.Confirmed[key]
|
||||
return ok
|
||||
}
|
||||
|
||||
func (service *Service) confirmCheckpointBlob(conversationID string, key string) {
|
||||
if service == nil || key == "" {
|
||||
return
|
||||
}
|
||||
service.checkpointBlobMu.Lock()
|
||||
defer service.checkpointBlobMu.Unlock()
|
||||
if service.checkpointBlobs == nil {
|
||||
service.checkpointBlobs = make(map[string]*checkpointBlobCacheEntry)
|
||||
}
|
||||
conversationID = strings.TrimSpace(conversationID)
|
||||
now := time.Now().UTC()
|
||||
entry := service.checkpointBlobs[conversationID]
|
||||
if entry == nil {
|
||||
entry = &checkpointBlobCacheEntry{Confirmed: make(map[string]struct{})}
|
||||
service.checkpointBlobs[conversationID] = entry
|
||||
}
|
||||
entry.Confirmed[key] = struct{}{}
|
||||
entry.LastAccess = now
|
||||
service.pruneCheckpointBlobCacheLocked(now)
|
||||
}
|
||||
|
||||
func (service *Service) pruneCheckpointBlobCacheLocked(now time.Time) {
|
||||
if service == nil || len(service.checkpointBlobs) == 0 {
|
||||
return
|
||||
}
|
||||
cutoff := now.Add(-checkpointBlobCacheIdleTTL)
|
||||
for conversationID, entry := range service.checkpointBlobs {
|
||||
if entry == nil || entry.LastAccess.Before(cutoff) {
|
||||
delete(service.checkpointBlobs, conversationID)
|
||||
}
|
||||
}
|
||||
for len(service.checkpointBlobs) > checkpointBlobCacheMaxConversations {
|
||||
oldestConversationID := ""
|
||||
oldestAccess := now
|
||||
for conversationID, entry := range service.checkpointBlobs {
|
||||
if entry == nil || oldestConversationID == "" || entry.LastAccess.Before(oldestAccess) {
|
||||
oldestConversationID = conversationID
|
||||
if entry != nil {
|
||||
oldestAccess = entry.LastAccess
|
||||
}
|
||||
}
|
||||
}
|
||||
if oldestConversationID == "" {
|
||||
return
|
||||
}
|
||||
delete(service.checkpointBlobs, oldestConversationID)
|
||||
}
|
||||
}
|
||||
|
||||
func checkpointCompletionAction(completion *pendingTurnCompletion) checkpointTerminalAction {
|
||||
if completion == nil {
|
||||
return checkpointTerminalAction{kind: checkpointTerminalActionNone}
|
||||
}
|
||||
return checkpointTerminalAction{
|
||||
kind: checkpointTerminalActionComplete,
|
||||
completion: *clonePendingTurnCompletion(completion),
|
||||
}
|
||||
}
|
||||
|
||||
func checkpointCancellationAction(message string) checkpointTerminalAction {
|
||||
return checkpointTerminalAction{
|
||||
kind: checkpointTerminalActionCancel,
|
||||
cancelMessage: firstNonEmpty(strings.TrimSpace(message), "[canceled] User aborted request"),
|
||||
}
|
||||
}
|
||||
|
||||
func mergeCheckpointTerminalAction(current checkpointTerminalAction, incoming checkpointTerminalAction) checkpointTerminalAction {
|
||||
switch {
|
||||
case incoming.kind == checkpointTerminalActionCancel:
|
||||
return incoming
|
||||
case current.kind == checkpointTerminalActionCancel:
|
||||
return current
|
||||
case incoming.kind == checkpointTerminalActionComplete:
|
||||
return incoming
|
||||
default:
|
||||
return current
|
||||
}
|
||||
}
|
||||
|
||||
func (action checkpointTerminalAction) completionValue() *pendingTurnCompletion {
|
||||
if action.kind != checkpointTerminalActionComplete {
|
||||
return nil
|
||||
}
|
||||
return clonePendingTurnCompletion(&action.completion)
|
||||
}
|
||||
|
||||
func (service *Service) queueCheckpointProjection(stream *ActiveStream, projection *CheckpointProjection, terminalAction checkpointTerminalAction) error {
|
||||
if service == nil || stream == nil || projection == nil || projection.State == nil {
|
||||
return nil
|
||||
}
|
||||
state, ok := proto.Clone(projection.State).(*agentv1.ConversationStateStructure)
|
||||
if !ok || state == nil {
|
||||
return fmt.Errorf("clone checkpoint state")
|
||||
}
|
||||
|
||||
stream.mu.Lock()
|
||||
if stream.PendingCheckpointBlobWrites == nil {
|
||||
stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite)
|
||||
}
|
||||
if stream.PendingCheckpointBlobRequests == nil {
|
||||
stream.PendingCheckpointBlobRequests = make(map[string]uint32)
|
||||
}
|
||||
stream.NextCheckpointRevision++
|
||||
if stream.NextCheckpointRevision == 0 {
|
||||
stream.NextCheckpointRevision++
|
||||
}
|
||||
revision := stream.NextCheckpointRevision
|
||||
if stream.PendingCheckpoint != nil {
|
||||
terminalAction = mergeCheckpointTerminalAction(stream.PendingCheckpoint.TerminalAction, terminalAction)
|
||||
}
|
||||
required := make(map[string]struct{}, len(projection.Blobs))
|
||||
toWrite := make([]struct {
|
||||
requestID uint32
|
||||
blob CheckpointBlob
|
||||
}, 0, len(projection.Blobs))
|
||||
for _, blob := range projection.Blobs {
|
||||
key := checkpointBlobKey(blob.ID)
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
required[key] = struct{}{}
|
||||
if service.confirmedCheckpointBlob(stream.ConversationID, key) {
|
||||
continue
|
||||
}
|
||||
if _, pending := stream.PendingCheckpointBlobRequests[key]; pending {
|
||||
continue
|
||||
}
|
||||
stream.NextCheckpointBlobRequestID++
|
||||
if stream.NextCheckpointBlobRequestID == 0 {
|
||||
stream.NextCheckpointBlobRequestID++
|
||||
}
|
||||
requestID := stream.NextCheckpointBlobRequestID
|
||||
stream.PendingCheckpointBlobWrites[requestID] = pendingCheckpointBlobWrite{
|
||||
Key: key,
|
||||
Revision: revision,
|
||||
}
|
||||
stream.PendingCheckpointBlobRequests[key] = requestID
|
||||
toWrite = append(toWrite, struct {
|
||||
requestID uint32
|
||||
blob CheckpointBlob
|
||||
}{requestID: requestID, blob: blob})
|
||||
}
|
||||
stream.PendingCheckpoint = &pendingCheckpointPublish{
|
||||
Revision: revision,
|
||||
State: state,
|
||||
Required: required,
|
||||
TerminalAction: terminalAction,
|
||||
}
|
||||
if terminalAction.kind != checkpointTerminalActionNone {
|
||||
stream.Phase = TurnPhaseCheckpointing
|
||||
}
|
||||
for requestID, write := range stream.PendingCheckpointBlobWrites {
|
||||
if _, stillRequired := required[write.Key]; stillRequired {
|
||||
continue
|
||||
}
|
||||
delete(stream.PendingCheckpointBlobWrites, requestID)
|
||||
delete(stream.PendingCheckpointBlobRequests, write.Key)
|
||||
}
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
|
||||
for _, item := range toWrite {
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{
|
||||
Message: buildSetCheckpointBlobMessage(item.requestID, item.blob),
|
||||
}); err != nil {
|
||||
service.discardPendingCheckpoint(stream, fmt.Errorf("publish checkpoint blob write: %w", err))
|
||||
return err
|
||||
}
|
||||
}
|
||||
if len(toWrite) > 0 {
|
||||
service.scheduleStreamTimer(
|
||||
stream,
|
||||
providerTimerKey(streamTimerCheckpointBlobs, ""),
|
||||
checkpointBlobWriteTimeout,
|
||||
streamTimerCheckpointBlobs,
|
||||
"",
|
||||
0,
|
||||
"checkpoint blob write timeout",
|
||||
)
|
||||
}
|
||||
return service.publishReadyCheckpoint(stream)
|
||||
}
|
||||
|
||||
func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion {
|
||||
if completion == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *completion
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func (service *Service) handleCheckpointBlobResult(stream *ActiveStream, message *agentv1.KvClientMessage) error {
|
||||
if service == nil || stream == nil || message == nil {
|
||||
return nil
|
||||
}
|
||||
result := message.GetSetBlobResult()
|
||||
if result == nil {
|
||||
return nil
|
||||
}
|
||||
stream.mu.Lock()
|
||||
write, ok := stream.PendingCheckpointBlobWrites[message.GetId()]
|
||||
pendingRequiresBlob := false
|
||||
if ok {
|
||||
delete(stream.PendingCheckpointBlobWrites, message.GetId())
|
||||
delete(stream.PendingCheckpointBlobRequests, write.Key)
|
||||
if stream.PendingCheckpoint != nil {
|
||||
_, pendingRequiresBlob = stream.PendingCheckpoint.Required[write.Key]
|
||||
}
|
||||
}
|
||||
conversationID := stream.ConversationID
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
if result.GetError() != nil {
|
||||
if !pendingRequiresBlob {
|
||||
return service.publishReadyCheckpoint(stream)
|
||||
}
|
||||
return service.abandonPendingCheckpoint(stream, fmt.Errorf(
|
||||
"write checkpoint blob %s: %s",
|
||||
checkpointBlobHex(write.Key),
|
||||
firstNonEmpty(result.GetError().GetMessage(), "client blob store rejected write"),
|
||||
))
|
||||
}
|
||||
service.confirmCheckpointBlob(conversationID, write.Key)
|
||||
return service.publishReadyCheckpoint(stream)
|
||||
}
|
||||
|
||||
func (service *Service) publishReadyCheckpoint(stream *ActiveStream) error {
|
||||
if service == nil || stream == nil {
|
||||
return nil
|
||||
}
|
||||
stream.mu.Lock()
|
||||
pending := stream.PendingCheckpoint
|
||||
if pending == nil {
|
||||
stream.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
for key := range pending.Required {
|
||||
if !service.confirmedCheckpointBlob(stream.ConversationID, key) {
|
||||
stream.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
}
|
||||
stream.PendingCheckpoint = nil
|
||||
state := pending.State
|
||||
terminalAction := pending.TerminalAction
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, ""))
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil {
|
||||
return err
|
||||
}
|
||||
switch terminalAction.kind {
|
||||
case checkpointTerminalActionComplete:
|
||||
if completion := terminalAction.completionValue(); completion != nil {
|
||||
return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion)
|
||||
}
|
||||
case checkpointTerminalActionCancel:
|
||||
return service.finishCanceledTurnAfterCheckpoint(stream, terminalAction.cancelMessage)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (service *Service) discardPendingCheckpoint(stream *ActiveStream, cause error) {
|
||||
if service == nil || stream == nil {
|
||||
return
|
||||
}
|
||||
stream.mu.Lock()
|
||||
stream.PendingCheckpoint = nil
|
||||
stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite)
|
||||
stream.PendingCheckpointBlobRequests = make(map[string]uint32)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, ""))
|
||||
if cause != nil {
|
||||
log.Printf("forwarder pending checkpoint discarded request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause)
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) abandonPendingCheckpoint(stream *ActiveStream, cause error) error {
|
||||
if service == nil || stream == nil {
|
||||
return nil
|
||||
}
|
||||
stream.mu.Lock()
|
||||
pending := stream.PendingCheckpoint
|
||||
stream.PendingCheckpoint = nil
|
||||
stream.PendingCheckpointBlobWrites = make(map[uint32]pendingCheckpointBlobWrite)
|
||||
stream.PendingCheckpointBlobRequests = make(map[string]uint32)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, ""))
|
||||
if cause != nil {
|
||||
log.Printf("forwarder checkpoint blob sync abandoned request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause)
|
||||
}
|
||||
if pending != nil {
|
||||
return service.failTerminalCheckpointSync(stream, cause)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (service *Service) finishCanceledTurnAfterCheckpoint(stream *ActiveStream, message string) error {
|
||||
if stream == nil {
|
||||
return nil
|
||||
}
|
||||
service.setTurnPhase(stream, TurnPhaseCanceled)
|
||||
return service.broker.Cancel(stream.RequestID, firstNonEmpty(strings.TrimSpace(message), "[canceled] User aborted request"))
|
||||
}
|
||||
|
||||
func (service *Service) failTerminalCheckpointSync(stream *ActiveStream, cause error) error {
|
||||
if stream == nil {
|
||||
return nil
|
||||
}
|
||||
message := "checkpoint synchronization failed"
|
||||
if cause != nil && strings.TrimSpace(cause.Error()) != "" {
|
||||
message = strings.TrimSpace(cause.Error())
|
||||
}
|
||||
service.setTurnPhase(stream, TurnPhaseFailed)
|
||||
return service.broker.Fail(stream.RequestID, "checkpoint_sync_error", message)
|
||||
}
|
||||
|
||||
func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error {
|
||||
if stream == nil {
|
||||
return nil
|
||||
}
|
||||
stream.mu.Lock()
|
||||
pendingCount := len(stream.PendingCheckpointBlobWrites)
|
||||
stream.mu.Unlock()
|
||||
if pendingCount == 0 {
|
||||
return service.publishReadyCheckpoint(stream)
|
||||
}
|
||||
return service.abandonPendingCheckpoint(stream, fmt.Errorf("%d checkpoint blob writes timed out", pendingCount))
|
||||
}
|
||||
@@ -0,0 +1,546 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
func TestCheckpointBlobSyncPublishesCheckpointAfterAllWrites(t *testing.T) {
|
||||
service, stream := testCheckpointBlobService(t)
|
||||
projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
newAssistantTextEntry(1, "request-1", "hi", "", ""),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
|
||||
if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queueCheckpointProjection() error = %v", err)
|
||||
}
|
||||
events, err := service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFromCursor() error = %v", err)
|
||||
}
|
||||
if len(events) != len(projection.Blobs) {
|
||||
t.Fatalf("events before ACK = %d, want %d blob writes", len(events), len(projection.Blobs))
|
||||
}
|
||||
for _, event := range events {
|
||||
if event.Message.GetKvServerMessage().GetSetBlobArgs() == nil {
|
||||
t.Fatalf("event before ACK = %#v, want set_blob_args", event.Message)
|
||||
}
|
||||
}
|
||||
|
||||
for index, event := range events {
|
||||
requestID := event.Message.GetKvServerMessage().GetId()
|
||||
if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{
|
||||
Id: requestID,
|
||||
Message: &agentv1.KvClientMessage_SetBlobResult{
|
||||
SetBlobResult: &agentv1.SetBlobResult{},
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("handleCheckpointBlobResult(%d) error = %v", index, err)
|
||||
}
|
||||
}
|
||||
events, err = service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFromCursor() after ACK error = %v", err)
|
||||
}
|
||||
if len(events) != len(projection.Blobs)+1 {
|
||||
t.Fatalf("events after ACK = %d, want %d", len(events), len(projection.Blobs)+1)
|
||||
}
|
||||
checkpoint := events[len(events)-1].Message.GetConversationCheckpointUpdate()
|
||||
if checkpoint == nil || len(checkpoint.GetTurns()) != 1 {
|
||||
t.Fatalf("last event checkpoint = %#v, want one Blob-backed turn", checkpoint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckpointBlobSyncRejectDoesNotPublishDanglingCheckpoint(t *testing.T) {
|
||||
service, stream := testCheckpointBlobService(t)
|
||||
projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queueCheckpointProjection() error = %v", err)
|
||||
}
|
||||
events, err := service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil || len(events) == 0 {
|
||||
t.Fatalf("blob write events = %d, err = %v", len(events), err)
|
||||
}
|
||||
requestID := events[0].Message.GetKvServerMessage().GetId()
|
||||
if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{
|
||||
Id: requestID,
|
||||
Message: &agentv1.KvClientMessage_SetBlobResult{
|
||||
SetBlobResult: &agentv1.SetBlobResult{Error: &agentv1.Error{Message: "disk full"}},
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("handleCheckpointBlobResult() error = %v", err)
|
||||
}
|
||||
events, err = service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFromCursor() after rejection error = %v", err)
|
||||
}
|
||||
for _, event := range events {
|
||||
if event.Message.GetConversationCheckpointUpdate() != nil {
|
||||
t.Fatal("rejected Blob write published a dangling checkpoint")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTerminalCheckpointRejectionFailsInsteadOfCompleting(t *testing.T) {
|
||||
service, stream := testCheckpointBlobService(t)
|
||||
projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, stream.RequestID, "hello"),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("projection: %v", err)
|
||||
}
|
||||
completion := &pendingTurnCompletion{RequestID: stream.RequestID}
|
||||
if err := service.queueCheckpointProjection(stream, projection, checkpointCompletionAction(completion)); err != nil {
|
||||
t.Fatalf("queue checkpoint: %v", err)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
var requestID uint32
|
||||
for pendingID := range stream.PendingCheckpointBlobWrites {
|
||||
requestID = pendingID
|
||||
break
|
||||
}
|
||||
stream.mu.Unlock()
|
||||
if requestID == 0 {
|
||||
t.Fatal("test did not queue a Blob write")
|
||||
}
|
||||
if err := rejectCheckpointBlob(service, stream, requestID, "disk full"); err != nil {
|
||||
t.Fatalf("reject terminal checkpoint: %v", err)
|
||||
}
|
||||
events, err := service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFromCursor() error = %v", err)
|
||||
}
|
||||
var failed, completed bool
|
||||
for _, event := range events {
|
||||
if event.End && event.TerminalErrorCode == "checkpoint_sync_error" {
|
||||
failed = true
|
||||
}
|
||||
if event.Message.GetInteractionUpdate().GetTurnEnded() != nil {
|
||||
completed = true
|
||||
}
|
||||
}
|
||||
if !failed || completed {
|
||||
t.Fatalf("terminal checkpoint events failed=%v completed=%v", failed, completed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckpointBlobSyncMergesRevisionsWithoutObsoleteFailure(t *testing.T) {
|
||||
service, stream := testCheckpointBlobService(t)
|
||||
first, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("first ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
latest, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
newAssistantTextEntry(1, "request-1", "latest answer", "", ""),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("latest ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if err := service.queueCheckpointProjection(stream, first, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queue first projection: %v", err)
|
||||
}
|
||||
firstEvents, err := service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("read first events: %v", err)
|
||||
}
|
||||
if err := service.queueCheckpointProjection(stream, latest, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queue latest projection: %v", err)
|
||||
}
|
||||
|
||||
latestRequired := make(map[string]struct{}, len(latest.Blobs))
|
||||
for _, blob := range latest.Blobs {
|
||||
latestRequired[string(blob.ID)] = struct{}{}
|
||||
}
|
||||
var obsoleteRequestID uint32
|
||||
for _, event := range firstEvents {
|
||||
message := event.Message.GetKvServerMessage()
|
||||
if message == nil || message.GetSetBlobArgs() == nil {
|
||||
continue
|
||||
}
|
||||
if _, required := latestRequired[string(message.GetSetBlobArgs().GetBlobId())]; !required {
|
||||
obsoleteRequestID = message.GetId()
|
||||
break
|
||||
}
|
||||
}
|
||||
if obsoleteRequestID == 0 {
|
||||
t.Fatal("test did not find an obsolete first-revision Blob write")
|
||||
}
|
||||
if err := rejectCheckpointBlob(service, stream, obsoleteRequestID, "obsolete write rejected"); err != nil {
|
||||
t.Fatalf("reject obsolete Blob: %v", err)
|
||||
}
|
||||
if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil {
|
||||
t.Fatalf("acknowledge latest Blob writes: %v", err)
|
||||
}
|
||||
|
||||
events, err := service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("read merged events: %v", err)
|
||||
}
|
||||
checkpoints := 0
|
||||
for _, event := range events {
|
||||
if checkpoint := event.Message.GetConversationCheckpointUpdate(); checkpoint != nil {
|
||||
checkpoints++
|
||||
if len(checkpoint.GetTurns()) != len(latest.State.GetTurns()) || string(checkpoint.GetTurns()[0]) != string(latest.State.GetTurns()[0]) {
|
||||
t.Fatalf("published checkpoint is not the latest revision: %#v", checkpoint)
|
||||
}
|
||||
}
|
||||
}
|
||||
if checkpoints != 1 {
|
||||
t.Fatalf("published checkpoints = %d, want exactly latest revision", checkpoints)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckpointBlobSyncCarriesCompletionIntoLatestRevision(t *testing.T) {
|
||||
service, stream := testCheckpointBlobService(t)
|
||||
first, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("first projection: %v", err)
|
||||
}
|
||||
latest, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
newAssistantTextEntry(1, "request-1", "done", "", ""),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("latest projection: %v", err)
|
||||
}
|
||||
completion := &pendingTurnCompletion{
|
||||
RequestID: stream.RequestID,
|
||||
Usage: turnUsageSnapshot{InputTokens: 11, OutputTokens: 7},
|
||||
}
|
||||
if err := service.queueCheckpointProjection(stream, first, checkpointCompletionAction(completion)); err != nil {
|
||||
t.Fatalf("queue completion projection: %v", err)
|
||||
}
|
||||
if err := service.queueCheckpointProjection(stream, latest, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queue latest projection: %v", err)
|
||||
}
|
||||
if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil {
|
||||
t.Fatalf("acknowledge latest Blob writes: %v", err)
|
||||
}
|
||||
|
||||
events, err := service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("read completion events: %v", err)
|
||||
}
|
||||
checkpointIndex, turnEndedIndex, endIndex := -1, -1, -1
|
||||
for index, event := range events {
|
||||
switch {
|
||||
case event.Message.GetConversationCheckpointUpdate() != nil:
|
||||
checkpointIndex = index
|
||||
case event.Message.GetInteractionUpdate().GetTurnEnded() != nil:
|
||||
turnEndedIndex = index
|
||||
case event.End:
|
||||
endIndex = index
|
||||
}
|
||||
}
|
||||
if checkpointIndex < 0 || turnEndedIndex <= checkpointIndex || endIndex <= turnEndedIndex {
|
||||
t.Fatalf("terminal order checkpoint=%d turn_ended=%d end=%d", checkpointIndex, turnEndedIndex, endIndex)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckpointBlobSyncTimeoutFailsStreamWithoutPublishingDanglingCheckpoint(t *testing.T) {
|
||||
service, stream := testCheckpointBlobService(t)
|
||||
projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("projection: %v", err)
|
||||
}
|
||||
if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queue checkpoint: %v", err)
|
||||
}
|
||||
if err := service.handleCheckpointBlobTimeout(stream); err != nil {
|
||||
t.Fatalf("timeout checkpoint: %v", err)
|
||||
}
|
||||
events, err := service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("read timeout events: %v", err)
|
||||
}
|
||||
var failed bool
|
||||
for _, event := range events {
|
||||
if event.Message.GetConversationCheckpointUpdate() != nil {
|
||||
t.Fatal("timed-out Blob dependency published a dangling checkpoint")
|
||||
}
|
||||
if event.End && event.TerminalErrorCode == "checkpoint_sync_error" {
|
||||
failed = true
|
||||
}
|
||||
}
|
||||
if !failed {
|
||||
t.Fatal("timed-out checkpoint did not fail the stream explicitly")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckpointBlobSyncReusesConversationCacheAcrossRequests(t *testing.T) {
|
||||
service, firstStream := testCheckpointBlobService(t)
|
||||
projection, err := service.projector.ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("projection: %v", err)
|
||||
}
|
||||
if err := service.queueCheckpointProjection(firstStream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queue first request: %v", err)
|
||||
}
|
||||
if err := acknowledgePendingCheckpointBlobs(service, firstStream); err != nil {
|
||||
t.Fatalf("acknowledge first request: %v", err)
|
||||
}
|
||||
|
||||
secondStream, err := service.broker.OpenStream(
|
||||
"request-2", firstStream.ConversationID, 2, "default", "default",
|
||||
agentv1.AgentMode_AGENT_MODE_AGENT, "continue",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("OpenStream() second request error = %v", err)
|
||||
}
|
||||
if err := service.queueCheckpointProjection(secondStream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queue second request: %v", err)
|
||||
}
|
||||
events, err := service.broker.ReadFromCursor(secondStream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("read second request events: %v", err)
|
||||
}
|
||||
if len(events) != 1 || events[0].Message.GetConversationCheckpointUpdate() == nil {
|
||||
t.Fatalf("second request events = %#v, want cached immediate checkpoint", events)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCancellationReplacesUnconfirmedCheckpointBeforeEnding(t *testing.T) {
|
||||
service, stream := testCheckpointBlobService(t)
|
||||
conversation := testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, stream.RequestID, "hello"),
|
||||
})
|
||||
if err := service.replaceCheckpointConversation(stream, conversation); err != nil {
|
||||
t.Fatalf("replaceCheckpointConversation() error = %v", err)
|
||||
}
|
||||
projection, err := service.projector.ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if err := service.queueCheckpointProjection(stream, projection, checkpointTerminalAction{kind: checkpointTerminalActionNone}); err != nil {
|
||||
t.Fatalf("queueCheckpointProjection() error = %v", err)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
var staleRequestID uint32
|
||||
for requestID := range stream.PendingCheckpointBlobWrites {
|
||||
staleRequestID = requestID
|
||||
break
|
||||
}
|
||||
stream.mu.Unlock()
|
||||
if staleRequestID == 0 {
|
||||
t.Fatal("test did not queue an unconfirmed Blob write")
|
||||
}
|
||||
|
||||
if err := service.handleCancelIntent(InboundIntent{
|
||||
Kind: "cancel",
|
||||
RequestID: stream.RequestID,
|
||||
CancelReason: "user stopped",
|
||||
}); err != nil {
|
||||
t.Fatalf("handleCancelIntent() error = %v", err)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
phase := stream.Phase
|
||||
status := stream.Status
|
||||
pendingCheckpoint := stream.PendingCheckpoint
|
||||
pendingWrites := len(stream.PendingCheckpointBlobWrites)
|
||||
stream.mu.Unlock()
|
||||
if phase != TurnPhaseCheckpointing || status != StreamStatusCreated {
|
||||
t.Fatalf("before checkpoint ACK phase=%s status=%s, want checkpointing/created", phase, status)
|
||||
}
|
||||
if pendingCheckpoint == nil || pendingWrites == 0 {
|
||||
t.Fatalf("before checkpoint ACK pending_checkpoint=%v pending_writes=%d", pendingCheckpoint != nil, pendingWrites)
|
||||
}
|
||||
if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{
|
||||
Id: staleRequestID,
|
||||
Message: &agentv1.KvClientMessage_SetBlobResult{
|
||||
SetBlobResult: &agentv1.SetBlobResult{},
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("stale Blob ACK error = %v", err)
|
||||
}
|
||||
if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil {
|
||||
t.Fatalf("acknowledge cancellation checkpoint: %v", err)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
phase = stream.Phase
|
||||
status = stream.Status
|
||||
stream.mu.Unlock()
|
||||
if phase != TurnPhaseCanceled || status != StreamStatusCanceled {
|
||||
t.Fatalf("after checkpoint ACK phase=%s status=%s, want canceled", phase, status)
|
||||
}
|
||||
assertCanceledEndEvent(t, service, stream)
|
||||
}
|
||||
|
||||
func TestCancellationMetadataFailureStillEndsStream(t *testing.T) {
|
||||
service, stream := testCheckpointBlobService(t)
|
||||
blockingPath := filepath.Join(t.TempDir(), "not-a-directory")
|
||||
if err := os.WriteFile(blockingPath, []byte("block child creation"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
service.store = NewConversationFileStore(blockingPath)
|
||||
if err := service.replaceCheckpointConversation(stream, testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, stream.RequestID, "hello"),
|
||||
})); err != nil {
|
||||
t.Fatalf("replaceCheckpointConversation() error = %v", err)
|
||||
}
|
||||
|
||||
if err := service.handleCancelIntent(InboundIntent{Kind: "cancel", RequestID: stream.RequestID}); err != nil {
|
||||
t.Fatalf("handleCancelIntent() error = %v", err)
|
||||
}
|
||||
if err := acknowledgePendingCheckpointBlobs(service, stream); err != nil {
|
||||
t.Fatalf("acknowledge cancellation checkpoint: %v", err)
|
||||
}
|
||||
assertCanceledEndEvent(t, service, stream)
|
||||
}
|
||||
|
||||
func TestCheckpointTerminalActionMergePriority(t *testing.T) {
|
||||
complete := checkpointCompletionAction(&pendingTurnCompletion{RequestID: "complete"})
|
||||
cancel := checkpointCancellationAction("user canceled")
|
||||
none := checkpointTerminalAction{kind: checkpointTerminalActionNone}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
current checkpointTerminalAction
|
||||
incoming checkpointTerminalAction
|
||||
wantKind checkpointTerminalActionKind
|
||||
wantID string
|
||||
}{
|
||||
{name: "none then complete", current: none, incoming: complete, wantKind: checkpointTerminalActionComplete, wantID: "complete"},
|
||||
{name: "complete then none", current: complete, incoming: none, wantKind: checkpointTerminalActionComplete, wantID: "complete"},
|
||||
{name: "complete then cancel", current: complete, incoming: cancel, wantKind: checkpointTerminalActionCancel},
|
||||
{name: "cancel then complete", current: cancel, incoming: complete, wantKind: checkpointTerminalActionCancel},
|
||||
{name: "cancel then none", current: cancel, incoming: none, wantKind: checkpointTerminalActionCancel},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
merged := mergeCheckpointTerminalAction(test.current, test.incoming)
|
||||
if merged.kind != test.wantKind {
|
||||
t.Fatalf("merged kind = %d, want %d", merged.kind, test.wantKind)
|
||||
}
|
||||
if test.wantID != "" {
|
||||
completion := merged.completionValue()
|
||||
if completion == nil || completion.RequestID != test.wantID {
|
||||
t.Fatalf("merged completion = %#v, want request_id=%s", completion, test.wantID)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckpointTerminalActionIsMutuallyExclusive(t *testing.T) {
|
||||
completion := &pendingTurnCompletion{RequestID: "request-1"}
|
||||
action := checkpointCompletionAction(completion)
|
||||
if action.kind != checkpointTerminalActionComplete || action.completionValue() == nil {
|
||||
t.Fatalf("completion action = %#v", action)
|
||||
}
|
||||
empty := checkpointCompletionAction(nil)
|
||||
if empty.kind != checkpointTerminalActionNone || empty.completionValue() != nil {
|
||||
t.Fatalf("empty action = %#v", empty)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCanceledEndEvent(t *testing.T, service *Service, stream *ActiveStream) {
|
||||
t.Helper()
|
||||
events, err := service.broker.ReadFromCursor(stream.RequestID, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFromCursor() error = %v", err)
|
||||
}
|
||||
for _, event := range events {
|
||||
if event.End && event.TerminalErrorCode == "canceled" {
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatal("cancellation did not publish canceled end event")
|
||||
}
|
||||
|
||||
func acknowledgePendingCheckpointBlobs(service *Service, stream *ActiveStream) error {
|
||||
for {
|
||||
stream.mu.Lock()
|
||||
requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites))
|
||||
for requestID := range stream.PendingCheckpointBlobWrites {
|
||||
requestIDs = append(requestIDs, requestID)
|
||||
}
|
||||
stream.mu.Unlock()
|
||||
if len(requestIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
for _, requestID := range requestIDs {
|
||||
if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{
|
||||
Id: requestID,
|
||||
Message: &agentv1.KvClientMessage_SetBlobResult{
|
||||
SetBlobResult: &agentv1.SetBlobResult{},
|
||||
},
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func rejectCheckpointBlob(service *Service, stream *ActiveStream, requestID uint32, message string) error {
|
||||
return service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{
|
||||
Id: requestID,
|
||||
Message: &agentv1.KvClientMessage_SetBlobResult{
|
||||
SetBlobResult: &agentv1.SetBlobResult{Error: &agentv1.Error{Message: message}},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func TestImportedTurnIDsRemainCheckpointPrefix(t *testing.T) {
|
||||
importedID := make([]byte, 32)
|
||||
for index := range importedID {
|
||||
importedID[index] = byte(index + 1)
|
||||
}
|
||||
conversation := testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 2, "request-2", "continued question"),
|
||||
})
|
||||
conversation.ImportedTurnIDs = [][]byte{importedID}
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if len(projection.State.GetTurns()) != 2 {
|
||||
t.Fatalf("turns = %d, want imported prefix plus projected turn", len(projection.State.GetTurns()))
|
||||
}
|
||||
if string(projection.State.GetTurns()[0]) != string(importedID) {
|
||||
t.Fatal("imported turn ID was not preserved as the checkpoint prefix")
|
||||
}
|
||||
}
|
||||
|
||||
func testCheckpointBlobService(t *testing.T) (*Service, *ActiveStream) {
|
||||
t.Helper()
|
||||
broker := NewStreamBroker()
|
||||
service := &Service{
|
||||
projector: NewHistoryProjector(),
|
||||
broker: broker,
|
||||
checkpointBlobs: make(map[string]*checkpointBlobCacheEntry),
|
||||
}
|
||||
stream, err := broker.OpenStream(
|
||||
"request-1",
|
||||
"conversation-1",
|
||||
1,
|
||||
"default",
|
||||
"default",
|
||||
agentv1.AgentMode_AGENT_MODE_AGENT,
|
||||
"hello",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("OpenStream() error = %v", err)
|
||||
}
|
||||
return service, stream
|
||||
}
|
||||
@@ -82,28 +82,30 @@ func (broker *StreamBroker) OpenStream(requestID string, conversationID string,
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
stream := &ActiveStream{
|
||||
RequestID: normalizedRequestID,
|
||||
ConversationID: strings.TrimSpace(conversationID),
|
||||
TurnSeq: turnSeq,
|
||||
ModelID: strings.TrimSpace(modelID),
|
||||
ModelName: strings.TrimSpace(modelName),
|
||||
Mode: normalizedMode,
|
||||
LatestUserText: strings.TrimSpace(latestUserText),
|
||||
Status: StreamStatusCreated,
|
||||
Backlog: make([]StreamEvent, 0, 64),
|
||||
Subscribers: make(map[string]*StreamSubscriber),
|
||||
PendingExecs: make(map[string]runtimecore.PendingExec),
|
||||
PendingInteractions: make(map[string]runtimecore.PendingInteraction),
|
||||
PartialToolCallIDs: make(map[string]struct{}),
|
||||
PatchEditQueues: make(map[string][]queuedPatchEditOperation),
|
||||
MCPToolServers: make(map[string]string),
|
||||
RecentCompletedExecs: make(map[uint32]time.Time),
|
||||
BackgroundShells: make(map[string]*BackgroundShellState),
|
||||
BackgroundShellsByMessageID: make(map[uint32]string),
|
||||
BackgroundShellsByExecID: make(map[string]string),
|
||||
BackgroundShellActions: make(map[string]time.Time),
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
RequestID: normalizedRequestID,
|
||||
ConversationID: strings.TrimSpace(conversationID),
|
||||
TurnSeq: turnSeq,
|
||||
ModelID: strings.TrimSpace(modelID),
|
||||
ModelName: strings.TrimSpace(modelName),
|
||||
Mode: normalizedMode,
|
||||
LatestUserText: strings.TrimSpace(latestUserText),
|
||||
Status: StreamStatusCreated,
|
||||
Backlog: make([]StreamEvent, 0, 64),
|
||||
Subscribers: make(map[string]*StreamSubscriber),
|
||||
PendingExecs: make(map[string]runtimecore.PendingExec),
|
||||
PendingInteractions: make(map[string]runtimecore.PendingInteraction),
|
||||
PartialToolCallIDs: make(map[string]struct{}),
|
||||
PatchEditQueues: make(map[string][]queuedPatchEditOperation),
|
||||
MCPToolServers: make(map[string]string),
|
||||
RecentCompletedExecs: make(map[uint32]time.Time),
|
||||
BackgroundShells: make(map[string]*BackgroundShellState),
|
||||
BackgroundShellsByMessageID: make(map[uint32]string),
|
||||
BackgroundShellsByExecID: make(map[string]string),
|
||||
BackgroundShellActions: make(map[string]time.Time),
|
||||
PendingCheckpointBlobWrites: make(map[uint32]pendingCheckpointBlobWrite),
|
||||
PendingCheckpointBlobRequests: make(map[string]uint32),
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
broker.streams[normalizedRequestID] = stream
|
||||
return stream, nil
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -750,6 +750,7 @@ func mergeConversationMetadata(target *ConversationFile, source *ConversationFil
|
||||
target.CurrentPlanText = source.CurrentPlanText
|
||||
target.CurrentPlans = clonePlanRegistryEntries(source.CurrentPlans)
|
||||
target.CurrentTodos = cloneTodoItems(source.CurrentTodos)
|
||||
target.ImportedTurnIDs = cloneByteSlices(source.ImportedTurnIDs)
|
||||
target.LatestRequestPrefix = cloneConversationRequestPrefix(source.LatestRequestPrefix)
|
||||
target.LastProviderCall = cloneConversationProviderCall(source.LastProviderCall)
|
||||
if !source.CreatedAt.IsZero() && (target.CreatedAt.IsZero() || source.CreatedAt.Before(target.CreatedAt)) {
|
||||
@@ -882,6 +883,7 @@ func cloneConversationFile(conversation *ConversationFile) *ConversationFile {
|
||||
cloned := *conversation
|
||||
cloned.CurrentPlans = clonePlanRegistryEntries(conversation.CurrentPlans)
|
||||
cloned.CurrentTodos = cloneTodoItems(conversation.CurrentTodos)
|
||||
cloned.ImportedTurnIDs = cloneByteSlices(conversation.ImportedTurnIDs)
|
||||
cloned.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix)
|
||||
cloned.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall)
|
||||
cloned.Entries = append([]HistoryEntry(nil), conversation.Entries...)
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"fmt"
|
||||
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptengine "cursor/internal/backend/agent/prompt"
|
||||
)
|
||||
|
||||
type importedBlobStore map[string][]byte
|
||||
|
||||
func newImportedBlobStore(items []*agentv1.PreFetchedBlob) (importedBlobStore, error) {
|
||||
if len(items) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
store := make(importedBlobStore, len(items))
|
||||
for _, item := range items {
|
||||
if item == nil || len(item.GetId()) == 0 {
|
||||
continue
|
||||
}
|
||||
if len(item.GetId()) != sha256.Size {
|
||||
return nil, fmt.Errorf("prefetched blob id length %d, want %d", len(item.GetId()), sha256.Size)
|
||||
}
|
||||
digest := sha256.Sum256(item.GetValue())
|
||||
if string(digest[:]) != string(item.GetId()) {
|
||||
return nil, fmt.Errorf("prefetched blob %x failed SHA-256 validation", item.GetId())
|
||||
}
|
||||
store[string(item.GetId())] = append([]byte(nil), item.GetValue()...)
|
||||
}
|
||||
return store, nil
|
||||
}
|
||||
|
||||
func (store importedBlobStore) resolve(id []byte) ([]byte, bool) {
|
||||
if len(id) == 0 || len(store) == 0 {
|
||||
return nil, false
|
||||
}
|
||||
value, ok := store[string(id)]
|
||||
return append([]byte(nil), value...), ok
|
||||
}
|
||||
|
||||
func decodeImportedTurn(raw []byte, blobs importedBlobStore) (*agentv1.ConversationTurnStructure, []byte, error) {
|
||||
if data, ok := blobs.resolve(raw); ok {
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(data, turn); err != nil || turn.GetTurn() == nil {
|
||||
return nil, nil, fmt.Errorf("decode imported turn blob %x: %w", raw, firstNonNilError(err, fmt.Errorf("turn payload is empty")))
|
||||
}
|
||||
return turn, append([]byte(nil), raw...), nil
|
||||
}
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(raw, turn); err == nil && turn.GetTurn() != nil {
|
||||
return turn, nil, nil
|
||||
}
|
||||
if len(raw) == sha256.Size {
|
||||
return nil, append([]byte(nil), raw...), nil
|
||||
}
|
||||
return nil, nil, fmt.Errorf("decode imported inline turn")
|
||||
}
|
||||
|
||||
func decodeImportedUserMessage(raw []byte, blobs importedBlobStore) (*agentv1.UserMessage, error) {
|
||||
data := raw
|
||||
if resolved, ok := blobs.resolve(raw); ok {
|
||||
data = resolved
|
||||
} else if len(raw) == sha256.Size {
|
||||
candidate := &agentv1.UserMessage{}
|
||||
if err := proto.Unmarshal(raw, candidate); err != nil || !hasKnownUserMessageContent(candidate) {
|
||||
return nil, fmt.Errorf("missing prefetched user message blob %x", raw)
|
||||
}
|
||||
return candidate, nil
|
||||
}
|
||||
message := &agentv1.UserMessage{}
|
||||
if err := proto.Unmarshal(data, message); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn user_message: %w", err)
|
||||
}
|
||||
return message, nil
|
||||
}
|
||||
|
||||
func decodeImportedStep(raw []byte, blobs importedBlobStore) (*agentv1.ConversationStep, error) {
|
||||
data := raw
|
||||
if resolved, ok := blobs.resolve(raw); ok {
|
||||
data = resolved
|
||||
} else if len(raw) == sha256.Size {
|
||||
candidate := &agentv1.ConversationStep{}
|
||||
if err := proto.Unmarshal(raw, candidate); err != nil || candidate.GetMessage() == nil {
|
||||
return nil, fmt.Errorf("missing prefetched conversation step blob %x", raw)
|
||||
}
|
||||
return candidate, nil
|
||||
}
|
||||
step := &agentv1.ConversationStep{}
|
||||
if err := proto.Unmarshal(data, step); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn step: %w", err)
|
||||
}
|
||||
if step.GetMessage() == nil {
|
||||
return nil, fmt.Errorf("decode imported turn step: payload is empty")
|
||||
}
|
||||
return step, nil
|
||||
}
|
||||
|
||||
func importedBlobTurnMessages(turn *agentv1.ConversationTurnStructure, blobs importedBlobStore) ([]modeladapter.Message, error) {
|
||||
if turn == nil || turn.GetAgentConversationTurn() == nil {
|
||||
return nil, nil
|
||||
}
|
||||
agentTurn := turn.GetAgentConversationTurn()
|
||||
messages := make([]modeladapter.Message, 0, 1+len(agentTurn.GetSteps()))
|
||||
if len(agentTurn.GetUserMessage()) > 0 {
|
||||
userMessage, err := decodeImportedUserMessage(agentTurn.GetUserMessage(), blobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok {
|
||||
messages = append(messages, toModelMessage(replay))
|
||||
}
|
||||
}
|
||||
for _, rawStep := range agentTurn.GetSteps() {
|
||||
if len(rawStep) == 0 {
|
||||
continue
|
||||
}
|
||||
step, err := decodeImportedStep(rawStep, blobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) {
|
||||
messages = append(messages, toModelMessage(replay))
|
||||
}
|
||||
}
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
func importedTurnIDs(turns [][]byte, blobs importedBlobStore) ([][]byte, error) {
|
||||
ids := make([][]byte, 0, len(turns))
|
||||
for _, raw := range turns {
|
||||
if len(raw) == 0 {
|
||||
continue
|
||||
}
|
||||
_, id, err := decodeImportedTurn(raw, blobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(id) > 0 {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func hasKnownUserMessageContent(message *agentv1.UserMessage) bool {
|
||||
if message == nil {
|
||||
return false
|
||||
}
|
||||
return message.GetText() != "" ||
|
||||
message.GetMessageId() != "" ||
|
||||
message.GetSelectedContext() != nil ||
|
||||
message.GetRichText() != "" ||
|
||||
len(message.GetConversationStateBlobId()) > 0 ||
|
||||
len(message.GetTextBlobId()) > 0 ||
|
||||
len(message.GetRichTextBlobId()) > 0
|
||||
}
|
||||
|
||||
func firstNonNilError(err error, fallback error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,433 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
modeladapter "cursor/internal/backend/agent/model"
|
||||
promptengine "cursor/internal/backend/agent/prompt"
|
||||
)
|
||||
|
||||
func TestProjectCheckpointProjectionBuildsBlobBackedTurns(t *testing.T) {
|
||||
toolCall := testEditToolCall(t, "file.txt")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
entries []HistoryEntry
|
||||
}{
|
||||
{
|
||||
name: "no tools",
|
||||
entries: []HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "hello"),
|
||||
newAssistantTextEntry(1, "request-1", "hi", "", ""),
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "completed tool call",
|
||||
entries: []HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "edit the file"),
|
||||
newToolCallEntry(1, "request-1", "call-1", "Edit", "", "", toolCall),
|
||||
newToolResultEntry(1, "request-1", "call-1", "Edit", `{"path":"file.txt"}`, "edited", "", toolCall),
|
||||
newAssistantTextEntry(1, "request-1", "done", "", ""),
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "unfinished tool call",
|
||||
entries: []HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "edit the file"),
|
||||
newToolCallEntry(1, "request-1", "call-1", "Edit", "", "", toolCall),
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "orphan tool result",
|
||||
entries: []HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "edit the file"),
|
||||
newToolResultEntry(1, "request-1", "call-1", "Edit", `{"path":"file.txt"}`, "edited", "", toolCall),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
conversation := testConversation(test.entries)
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if len(projection.State.GetTurns()) != 1 {
|
||||
t.Fatalf("ProjectCheckpointProjection() turns = %d, want 1 Blob ID", len(projection.State.GetTurns()))
|
||||
}
|
||||
assertCheckpointBlobGraph(t, projection)
|
||||
messages, err := promptengine.DecodeReplayMessages(projection.State.GetRootPromptMessagesJson())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeReplayMessages() error = %v", err)
|
||||
}
|
||||
if len(messages) == 0 {
|
||||
t.Fatal("ProjectCheckpointProjection() removed all root prompt replay history")
|
||||
}
|
||||
if messages[0].Role != "user" || messages[0].Content == "" {
|
||||
t.Fatalf("first replay message = %#v, want retained user history", messages[0])
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectLegacyCheckpointLargeModelHistoryUsesRootReplay(t *testing.T) {
|
||||
entries := make([]HistoryEntry, 0, 400)
|
||||
for turn := int64(1); turn <= 200; turn++ {
|
||||
requestID := fmt.Sprintf("request-%d", turn)
|
||||
entries = append(entries,
|
||||
testModelMessageEntry(t, turn, requestID, modeladapter.Message{Role: "user", Content: fmt.Sprintf("question %d", turn)}),
|
||||
testModelMessageEntry(t, turn, requestID, modeladapter.Message{Role: "assistant", Content: fmt.Sprintf("answer %d", turn)}),
|
||||
)
|
||||
}
|
||||
|
||||
state, err := NewHistoryProjector().ProjectLegacyCheckpoint(testConversation(entries))
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectLegacyCheckpoint() error = %v", err)
|
||||
}
|
||||
if len(state.GetTurns()) != 0 {
|
||||
t.Fatalf("ProjectLegacyCheckpoint() model-only turns = %d, want 0", len(state.GetTurns()))
|
||||
}
|
||||
messages, err := promptengine.DecodeReplayMessages(state.GetRootPromptMessagesJson())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeReplayMessages() error = %v", err)
|
||||
}
|
||||
if len(messages) != 400 {
|
||||
t.Fatalf("decoded replay messages = %d, want 400", len(messages))
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectLegacyCheckpointSnapshotIsIsolatedFromLaterHistory(t *testing.T) {
|
||||
conversation := testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "first question"),
|
||||
newAssistantTextEntry(1, "request-1", "first answer", "", ""),
|
||||
})
|
||||
projector := NewHistoryProjector()
|
||||
midpointProjection, err := projector.ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("midpoint ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
midpoint := midpointProjection.State
|
||||
|
||||
appendEntriesInPlace(conversation, []HistoryEntry{
|
||||
testUserMessageEntry(t, 2, "request-2", "second question"),
|
||||
newAssistantTextEntry(2, "request-2", "second answer", "", ""),
|
||||
})
|
||||
latestProjection, err := projector.ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("latest ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
latest := latestProjection.State
|
||||
|
||||
midpointMessages, err := promptengine.DecodeReplayMessages(midpoint.GetRootPromptMessagesJson())
|
||||
if err != nil {
|
||||
t.Fatalf("decode midpoint replay: %v", err)
|
||||
}
|
||||
latestMessages, err := promptengine.DecodeReplayMessages(latest.GetRootPromptMessagesJson())
|
||||
if err != nil {
|
||||
t.Fatalf("decode latest replay: %v", err)
|
||||
}
|
||||
if len(midpointMessages) != 2 {
|
||||
t.Fatalf("midpoint replay messages = %d, want 2", len(midpointMessages))
|
||||
}
|
||||
if len(latestMessages) != 4 {
|
||||
t.Fatalf("latest replay messages = %d, want 4", len(latestMessages))
|
||||
}
|
||||
if len(midpoint.GetTurns()) != 1 || len(latest.GetTurns()) != 2 {
|
||||
t.Fatalf("checkpoint turn counts = (%d, %d), want (1, 2)", len(midpoint.GetTurns()), len(latest.GetTurns()))
|
||||
}
|
||||
assertCheckpointBlobGraph(t, midpointProjection)
|
||||
assertCheckpointBlobGraph(t, latestProjection)
|
||||
}
|
||||
|
||||
func TestProjectCheckpointProjectionKeepsVisibleTurnsAcrossCompaction(t *testing.T) {
|
||||
summaryPayload, err := json.Marshal(compactionSummaryEntryPayload{Summary: "first turn summarized"})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal compaction summary: %v", err)
|
||||
}
|
||||
conversation := testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "first question"),
|
||||
newAssistantTextEntry(1, "request-1", "first answer", "", ""),
|
||||
{TurnSeq: 0, Role: "system", Kind: "compaction_summary", Payload: summaryPayload},
|
||||
testUserMessageEntry(t, 2, "request-2", "second question"),
|
||||
newAssistantTextEntry(2, "request-2", "second answer", "", ""),
|
||||
})
|
||||
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if len(projection.State.GetTurns()) != 2 {
|
||||
t.Fatalf("visible turns after compaction = %d, want 2", len(projection.State.GetTurns()))
|
||||
}
|
||||
assertCheckpointBlobGraph(t, projection)
|
||||
|
||||
messages, err := promptengine.DecodeReplayMessages(projection.State.GetRootPromptMessagesJson())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeReplayMessages() error = %v", err)
|
||||
}
|
||||
if len(messages) != 3 {
|
||||
t.Fatalf("compacted root replay messages = %d, want summary plus latest turn", len(messages))
|
||||
}
|
||||
if messages[0].Role != "user" || messages[0].Content != "<conversation_summary>\nfirst turn summarized\n</conversation_summary>" {
|
||||
t.Fatalf("first compacted replay message = %#v", messages[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportedConversationStateRejectsBlobTurnIDsWithoutPrefetchedData(t *testing.T) {
|
||||
turnID := sha256.Sum256([]byte("imported turn"))
|
||||
state := &agentv1.ConversationStateStructure{Turns: [][]byte{turnID[:]}}
|
||||
if _, err := importedConversationStateModelMessages(state, nil); err == nil {
|
||||
t.Fatal("importedConversationStateModelMessages() accepted unresolved Blob turn")
|
||||
}
|
||||
conversation := testConversation(nil)
|
||||
service := &Service{}
|
||||
if _, err := service.importConversationState(conversation, state, nil); err == nil {
|
||||
t.Fatal("importConversationState() accepted unresolved Blob turn")
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportedConversationStateRestoresBlobOnlyForkFromPrefetchedBlobs(t *testing.T) {
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(testConversation([]HistoryEntry{
|
||||
testUserMessageEntry(t, 1, "request-1", "parent question"),
|
||||
newAssistantTextEntry(1, "request-1", "parent answer", "", ""),
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
prefetched := make([]*agentv1.PreFetchedBlob, 0, len(projection.Blobs))
|
||||
for _, blob := range projection.Blobs {
|
||||
prefetched = append(prefetched, &agentv1.PreFetchedBlob{Id: blob.ID, Value: blob.Data})
|
||||
}
|
||||
state := proto.Clone(projection.State).(*agentv1.ConversationStateStructure)
|
||||
state.RootPromptMessagesJson = nil
|
||||
conversation := testConversation(nil)
|
||||
entries, err := (&Service{}).importConversationState(conversation, state, prefetched)
|
||||
if err != nil {
|
||||
t.Fatalf("importConversationState() error = %v", err)
|
||||
}
|
||||
if len(conversation.ImportedTurnIDs) != 1 {
|
||||
t.Fatalf("ImportedTurnIDs = %d, want 1", len(conversation.ImportedTurnIDs))
|
||||
}
|
||||
if len(entries) != 2 {
|
||||
t.Fatalf("imported model entries = %d, want user and assistant", len(entries))
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportedInlineTurnWithSHA256LengthIsNotMisclassified(t *testing.T) {
|
||||
var rawTurn []byte
|
||||
for size := 1; size <= 128; size++ {
|
||||
rawUser, err := proto.Marshal(&agentv1.UserMessage{Text: strings.Repeat("x", size), MessageId: "inline"})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal user message: %v", err)
|
||||
}
|
||||
rawTurn, err = proto.Marshal(&agentv1.ConversationTurnStructure{
|
||||
Turn: &agentv1.ConversationTurnStructure_AgentConversationTurn{
|
||||
AgentConversationTurn: &agentv1.AgentConversationTurnStructure{UserMessage: rawUser},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal turn: %v", err)
|
||||
}
|
||||
if len(rawTurn) == sha256.Size {
|
||||
break
|
||||
}
|
||||
}
|
||||
if len(rawTurn) != sha256.Size {
|
||||
t.Fatal("test could not construct a 32-byte inline turn")
|
||||
}
|
||||
ids, err := importedTurnIDs([][]byte{rawTurn}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("importedTurnIDs() error = %v", err)
|
||||
}
|
||||
if len(ids) != 0 {
|
||||
t.Fatal("32-byte inline turn was misclassified as a Blob ID")
|
||||
}
|
||||
messages, err := importedConversationStateModelMessages(&agentv1.ConversationStateStructure{Turns: [][]byte{rawTurn}}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("importedConversationStateModelMessages() error = %v", err)
|
||||
}
|
||||
if len(messages) != 1 || messages[0].Role != "user" {
|
||||
t.Fatalf("inline turn messages = %#v, want one user message", messages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportedTurnIDsPersistThroughConversationStore(t *testing.T) {
|
||||
store := NewConversationFileStore(t.TempDir())
|
||||
turnID := sha256.Sum256([]byte("parent turn"))
|
||||
conversation := testConversation(nil)
|
||||
conversation.ImportedTurnIDs = [][]byte{turnID[:]}
|
||||
persisted, err := store.SaveConversationWithEntries(conversation.ConversationID, conversation, []HistoryEntry{
|
||||
testUserMessageEntry(t, 2, "request-2", "fork question"),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SaveConversationWithEntries() error = %v", err)
|
||||
}
|
||||
if len(persisted.ImportedTurnIDs) != 1 || string(persisted.ImportedTurnIDs[0]) != string(turnID[:]) {
|
||||
t.Fatalf("persisted ImportedTurnIDs = %x, want %x", persisted.ImportedTurnIDs, turnID)
|
||||
}
|
||||
loaded, err := store.LoadConversation(conversation.ConversationID)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConversation() error = %v", err)
|
||||
}
|
||||
if len(loaded.ImportedTurnIDs) != 1 || string(loaded.ImportedTurnIDs[0]) != string(turnID[:]) {
|
||||
t.Fatalf("loaded ImportedTurnIDs = %x, want %x", loaded.ImportedTurnIDs, turnID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewindImportedTurnPrefixUsesClientForkPoint(t *testing.T) {
|
||||
ids := make([][]byte, 4)
|
||||
for index := range ids {
|
||||
digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1)))
|
||||
ids[index] = digest[:]
|
||||
}
|
||||
trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{
|
||||
TargetTurnSeq: 4,
|
||||
HasClientTurnCount: true,
|
||||
ClientTurnCount: 2,
|
||||
})
|
||||
if len(trimmed) != 2 || string(trimmed[0]) != string(ids[0]) || string(trimmed[1]) != string(ids[1]) {
|
||||
t.Fatalf("rewindImportedTurnPrefix() = %x, want first two IDs", trimmed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewindImportedTurnPrefixClearsAllIDsAtClientTurnZero(t *testing.T) {
|
||||
ids := make([][]byte, 2)
|
||||
for index := range ids {
|
||||
digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1)))
|
||||
ids[index] = digest[:]
|
||||
}
|
||||
trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{
|
||||
TargetTurnSeq: 3,
|
||||
HasClientTurnCount: true,
|
||||
ClientTurnCount: 0,
|
||||
})
|
||||
if trimmed != nil {
|
||||
t.Fatalf("rewindImportedTurnPrefix() = %x, want nil at client turn zero", trimmed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRewindImportedTurnPrefixUsesTargetWithoutClientCount(t *testing.T) {
|
||||
ids := make([][]byte, 4)
|
||||
for index := range ids {
|
||||
digest := sha256.Sum256([]byte(fmt.Sprintf("turn-%d", index+1)))
|
||||
ids[index] = digest[:]
|
||||
}
|
||||
trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{TargetTurnSeq: 3})
|
||||
if len(trimmed) != 2 || string(trimmed[0]) != string(ids[0]) || string(trimmed[1]) != string(ids[1]) {
|
||||
t.Fatalf("rewindImportedTurnPrefix() = %x, want target-derived first two IDs", trimmed)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCheckpointBlobGraph(t *testing.T, projection *CheckpointProjection) {
|
||||
t.Helper()
|
||||
if projection == nil || projection.State == nil {
|
||||
t.Fatal("checkpoint projection is nil")
|
||||
}
|
||||
blobByID := make(map[string][]byte, len(projection.Blobs))
|
||||
for _, blob := range projection.Blobs {
|
||||
if len(blob.ID) != sha256.Size {
|
||||
t.Fatalf("blob id length = %d, want %d", len(blob.ID), sha256.Size)
|
||||
}
|
||||
digest := sha256.Sum256(blob.Data)
|
||||
if string(blob.ID) != string(digest[:]) {
|
||||
t.Fatal("blob id does not match SHA-256(data)")
|
||||
}
|
||||
blobByID[string(blob.ID)] = blob.Data
|
||||
}
|
||||
for _, turnID := range projection.State.GetTurns() {
|
||||
turnData, ok := blobByID[string(turnID)]
|
||||
if !ok {
|
||||
t.Fatal("turn references missing blob")
|
||||
}
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(turnData, turn); err != nil {
|
||||
t.Fatalf("decode turn blob: %v", err)
|
||||
}
|
||||
agentTurn := turn.GetAgentConversationTurn()
|
||||
if agentTurn == nil {
|
||||
continue
|
||||
}
|
||||
if userID := agentTurn.GetUserMessage(); len(userID) > 0 {
|
||||
userData, exists := blobByID[string(userID)]
|
||||
if !exists {
|
||||
t.Fatal("turn references missing user message blob")
|
||||
}
|
||||
if err := proto.Unmarshal(userData, &agentv1.UserMessage{}); err != nil {
|
||||
t.Fatalf("decode user message blob: %v", err)
|
||||
}
|
||||
}
|
||||
for _, stepID := range agentTurn.GetSteps() {
|
||||
stepData, exists := blobByID[string(stepID)]
|
||||
if !exists {
|
||||
t.Fatal("turn references missing step blob")
|
||||
}
|
||||
if err := proto.Unmarshal(stepData, &agentv1.ConversationStep{}); err != nil {
|
||||
t.Fatalf("decode conversation step blob: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func testConversation(entries []HistoryEntry) *ConversationFile {
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
RootConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 1,
|
||||
NextEntrySeq: 1,
|
||||
Entries: make([]HistoryEntry, 0, len(entries)),
|
||||
}
|
||||
appendEntriesInPlace(conversation, entries)
|
||||
return conversation
|
||||
}
|
||||
|
||||
func testUserMessageEntry(t *testing.T, turnSeq int64, requestID string, text string) HistoryEntry {
|
||||
t.Helper()
|
||||
payload, err := protojson.Marshal(&agentv1.UserMessage{Text: text, MessageId: fmt.Sprintf("message-%d", turnSeq)})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal user message: %v", err)
|
||||
}
|
||||
return HistoryEntry{
|
||||
TurnSeq: turnSeq,
|
||||
RequestID: requestID,
|
||||
Role: "user",
|
||||
Kind: "user_message",
|
||||
Payload: payload,
|
||||
}
|
||||
}
|
||||
|
||||
func testModelMessageEntry(t *testing.T, turnSeq int64, requestID string, message modeladapter.Message) HistoryEntry {
|
||||
t.Helper()
|
||||
entry, ok, err := newModelMessageEntry(turnSeq, requestID, message)
|
||||
if err != nil {
|
||||
t.Fatalf("newModelMessageEntry() error = %v", err)
|
||||
}
|
||||
if !ok {
|
||||
t.Fatal("newModelMessageEntry() rejected test message")
|
||||
}
|
||||
return entry
|
||||
}
|
||||
|
||||
func testEditToolCall(t *testing.T, path string) []byte {
|
||||
t.Helper()
|
||||
payload, err := protojson.Marshal(&agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_EditToolCall{
|
||||
EditToolCall: &agentv1.EditToolCall{
|
||||
Args: &agentv1.EditArgs{Path: path},
|
||||
},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal edit tool call: %v", err)
|
||||
}
|
||||
return payload
|
||||
}
|
||||
@@ -36,7 +36,7 @@ type runRewindMatch struct {
|
||||
}
|
||||
|
||||
func (service *Service) decideRunRewind(intent InboundIntent, conversation *ConversationFile) runRewindDecision {
|
||||
decision := runRewindDecision{ClientTurnCount: -1}
|
||||
decision := runRewindDecision{}
|
||||
if !shouldEvaluateRunRewind(intent) {
|
||||
return decision
|
||||
}
|
||||
@@ -121,7 +121,7 @@ func selectRunRewindMatch(matches []runRewindMatch, clientTurnCount int, hasClie
|
||||
if len(matches) == 0 {
|
||||
return runRewindMatch{}, "no_match"
|
||||
}
|
||||
if hasClientTurnCount && clientTurnCount >= 0 {
|
||||
if hasClientTurnCount {
|
||||
targetTurnSeq := int64(clientTurnCount) + 1
|
||||
for _, match := range matches {
|
||||
if match.Entry.TurnSeq == targetTurnSeq {
|
||||
@@ -224,6 +224,7 @@ func (service *Service) applyRunRewindToConversation(conversation *ConversationF
|
||||
conversation.Entries = nil
|
||||
conversation.NextEntrySeq = 1
|
||||
conversation.NextTurnSeq = 1
|
||||
conversation.ImportedTurnIDs = rewindImportedTurnPrefix(conversation.ImportedTurnIDs, decision)
|
||||
appendEntriesInPlace(conversation, appendReplacementRunEntries(decision.PrefixEntries, entries))
|
||||
applyRunRewindConversationState(conversation, intent, turnSeq)
|
||||
deriveConversationLoopState(conversation)
|
||||
@@ -269,10 +270,30 @@ func applyRunRewindMetadata(conversation *ConversationFile, source *Conversation
|
||||
if source.TokenDetailsMaxTokens > 0 {
|
||||
conversation.TokenDetailsMaxTokens = source.TokenDetailsMaxTokens
|
||||
}
|
||||
decision := runRewindDecision{TargetTurnSeq: turnSeq}
|
||||
if intent.ConversationState != nil {
|
||||
decision.HasClientTurnCount = true
|
||||
decision.ClientTurnCount = len(intent.ConversationState.GetTurns())
|
||||
}
|
||||
conversation.ImportedTurnIDs = rewindImportedTurnPrefix(source.ImportedTurnIDs, decision)
|
||||
}
|
||||
applyRunRewindConversationState(conversation, intent, turnSeq)
|
||||
}
|
||||
|
||||
func rewindImportedTurnPrefix(importedTurnIDs [][]byte, decision runRewindDecision) [][]byte {
|
||||
keep := decision.TargetTurnSeq - 1
|
||||
if decision.HasClientTurnCount {
|
||||
keep = int64(decision.ClientTurnCount)
|
||||
}
|
||||
if keep <= 0 || len(importedTurnIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
if keep > int64(len(importedTurnIDs)) {
|
||||
keep = int64(len(importedTurnIDs))
|
||||
}
|
||||
return cloneByteSlices(importedTurnIDs[:keep])
|
||||
}
|
||||
|
||||
func (service *Service) logRunRewindDecision(requestID string, conversationID string, eventName string, decision runRewindDecision) {
|
||||
if service == nil || !decision.Evaluated {
|
||||
return
|
||||
|
||||
@@ -50,7 +50,7 @@ func (service *Service) bootstrapRuntimeConversation(intent InboundIntent) (*Con
|
||||
}
|
||||
importedEntries := []HistoryEntry(nil)
|
||||
if len(conversation.Entries) == 0 && intent.ConversationState != nil {
|
||||
importedEntries, err = service.importConversationState(conversation, intent.ConversationState)
|
||||
importedEntries, err = service.importConversationState(conversation, intent.ConversationState, intent.PreFetchedBlobs)
|
||||
if err != nil {
|
||||
return nil, agentv1.AgentMode_AGENT_MODE_AGENT, 0, nil, err
|
||||
}
|
||||
@@ -138,6 +138,7 @@ func (service *Service) syncConversationRecord(conversationID string, conversati
|
||||
item.AutoCompactionReserveTokens = conversation.AutoCompactionReserveTokens
|
||||
item.AutoCompactionTriggeredAt = conversation.AutoCompactionTriggeredAt
|
||||
item.AutoCompactionSourceModelCallID = conversation.AutoCompactionSourceModelCallID
|
||||
item.ImportedTurnIDs = cloneByteSlices(conversation.ImportedTurnIDs)
|
||||
item.LatestRequestPrefix = cloneConversationRequestPrefix(conversation.LatestRequestPrefix)
|
||||
item.LastProviderCall = cloneConversationProviderCall(conversation.LastProviderCall)
|
||||
item.CreatedAt = conversation.CreatedAt
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"log"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
@@ -261,6 +262,8 @@ type Service struct {
|
||||
execBridge execbridge.ExecBridge
|
||||
interactionBridge interactionbridge.InteractionBridge
|
||||
appendSeq *appendSequenceTracker
|
||||
checkpointBlobMu sync.Mutex
|
||||
checkpointBlobs map[string]*checkpointBlobCacheEntry
|
||||
}
|
||||
|
||||
type agentModelMemory interface {
|
||||
@@ -300,6 +303,7 @@ func NewService(historyRoot string, resolver modeladapter.ChannelResolver) *Serv
|
||||
execBridge: execbridge.NewBridge(),
|
||||
interactionBridge: interactionbridge.NewBridge(),
|
||||
appendSeq: newAppendSequenceTracker(),
|
||||
checkpointBlobs: make(map[string]*checkpointBlobCacheEntry),
|
||||
}
|
||||
service.startHistoryMaintenance()
|
||||
store.SyncAllCursorTranscriptsBestEffort()
|
||||
@@ -328,6 +332,7 @@ func newServiceWithDependencies(store *ConversationFileStore, projector *History
|
||||
execBridge: execbridge.NewBridge(),
|
||||
interactionBridge: interactionbridge.NewBridge(),
|
||||
appendSeq: newAppendSequenceTracker(),
|
||||
checkpointBlobs: make(map[string]*checkpointBlobCacheEntry),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -557,6 +562,7 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A
|
||||
}
|
||||
intent.ConversationID = conversationID
|
||||
intent.ConversationState = runRequest.GetConversationState()
|
||||
intent.PreFetchedBlobs = runRequest.GetPreFetchedBlobs()
|
||||
intent.UserMessage = extractUserMessage(message)
|
||||
intent.RequestContext = extractRequestContext(message)
|
||||
if service.shouldIgnoreEmptyResumeRunRequest(requestID, runRequest, intent.UserMessage, intent.RequestContext) {
|
||||
@@ -604,6 +610,7 @@ func (service *Service) decodeInboundIntent(requestID string, message *agentv1.A
|
||||
intent.ConversationID = conversationID
|
||||
intent.SubagentTypeName = strings.TrimSpace(prewarmRequest.GetSubagentTypeName())
|
||||
intent.ConversationState = prewarmRequest.GetConversationState()
|
||||
intent.PreFetchedBlobs = prewarmRequest.GetPreFetchedBlobs()
|
||||
intent.Mode, intent.ModeSource, intent.HasExplicitMode, err = extractPrewarmMode(prewarmRequest)
|
||||
if err != nil {
|
||||
return InboundIntent{}, err
|
||||
@@ -819,11 +826,14 @@ func (service *Service) snapshotVisibleTurns(conversation *ConversationFile) ([]
|
||||
if service == nil || service.projector == nil || conversation == nil {
|
||||
return nil, nil
|
||||
}
|
||||
state, err := service.projector.ProjectLegacyCheckpoint(conversation)
|
||||
projection, err := service.projector.ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cloneByteSlices(state.GetTurns()), nil
|
||||
if projection == nil || projection.State == nil {
|
||||
return nil, fmt.Errorf("checkpoint projection is empty")
|
||||
}
|
||||
return cloneByteSlices(projection.State.GetTurns()), nil
|
||||
}
|
||||
|
||||
// handleCancelIntent 处理取消请求,并向客户端发送执行桥 abort。
|
||||
@@ -833,42 +843,60 @@ func (service *Service) handleCancelIntent(intent InboundIntent) error {
|
||||
return fmt.Errorf("request is not active: %s", intent.RequestID)
|
||||
}
|
||||
hasCheckpoint := checkpointConversationInitialized(stream)
|
||||
if hasCheckpoint {
|
||||
cancelReason := firstNonEmpty(intent.CancelReason, "user aborted")
|
||||
_, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
|
||||
newMetadataEntry(stream.TurnSeq, intent.RequestID, "control", map[string]any{
|
||||
"status": "canceled",
|
||||
"reason": cancelReason,
|
||||
"replay_policy": cancelReplayPolicyForReason(cancelReason),
|
||||
}),
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
stream.mu.Lock()
|
||||
pendingExecs := make([]runtimecore.PendingExec, 0, len(stream.PendingExecs))
|
||||
for _, pending := range stream.PendingExecs {
|
||||
pendingExecs = append(pendingExecs, pending)
|
||||
}
|
||||
if stream.ProviderCancel != nil {
|
||||
stream.ProviderCancel()
|
||||
stream.ProviderCancel = nil
|
||||
}
|
||||
stream.ProviderActive = false
|
||||
stream.CurrentProviderToken++
|
||||
stream.CurrentCompactionToken++
|
||||
stream.PendingProviderAction = providerActionNone
|
||||
stream.PendingCompaction = nil
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
if hasCheckpoint {
|
||||
cancelReason := firstNonEmpty(intent.CancelReason, "user aborted")
|
||||
cancelEntry := newMetadataEntry(stream.TurnSeq, intent.RequestID, "control", map[string]any{
|
||||
"status": "canceled",
|
||||
"reason": cancelReason,
|
||||
"replay_policy": cancelReplayPolicyForReason(cancelReason),
|
||||
})
|
||||
if _, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{cancelEntry}); err != nil {
|
||||
log.Printf("forwarder cancellation metadata persistence failed request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, err)
|
||||
if memoryErr := service.appendCheckpointEntries(stream, []HistoryEntry{cancelEntry}); memoryErr != nil {
|
||||
return memoryErr
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, pending := range pendingExecs {
|
||||
_ = service.broker.Publish(intent.RequestID, StreamEvent{
|
||||
Message: buildExecAbortMessage(pending),
|
||||
})
|
||||
}
|
||||
if hasCheckpoint {
|
||||
if err := service.publishCheckpoint(stream.RequestID, stream.ConversationID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
clearPendingProviderCompletion(stream)
|
||||
terminalMessage := firstNonEmpty(intent.CancelReason, "[canceled] User aborted request")
|
||||
stream.mu.Lock()
|
||||
stream.PendingProviderAction = providerActionNone
|
||||
stream.PendingExecs = make(map[string]runtimecore.PendingExec)
|
||||
stream.PendingInteractions = make(map[string]runtimecore.PendingInteraction)
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
service.setTurnPhase(stream, TurnPhaseCanceled)
|
||||
return service.broker.Cancel(intent.RequestID, firstNonEmpty(intent.CancelReason, "[canceled] User aborted request"))
|
||||
service.discardPendingCheckpoint(stream, fmt.Errorf("checkpoint superseded by cancellation"))
|
||||
if hasCheckpoint {
|
||||
if err := service.publishCheckpointWithTerminalAction(
|
||||
stream.RequestID,
|
||||
stream.ConversationID,
|
||||
checkpointCancellationAction(terminalMessage),
|
||||
); err != nil {
|
||||
return service.failTerminalCheckpointSync(stream, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return service.finishCanceledTurnAfterCheckpoint(stream, terminalMessage)
|
||||
}
|
||||
|
||||
// handleExecResult 处理客户端返回的执行桥结果,并在终态时把 tool_result 写回 history。
|
||||
@@ -2122,9 +2150,18 @@ func (service *Service) completeSuccessfulTurn(stream *ActiveStream, completion
|
||||
err,
|
||||
)
|
||||
}
|
||||
if err := service.publishCheckpoint(requestID, conversationID); err != nil {
|
||||
if err := service.publishCheckpointWithCompletion(requestID, conversationID, &completion); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
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 +2188,15 @@ 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, conversationID string, completion *pendingTurnCompletion) error {
|
||||
return service.publishCheckpointWithTerminalAction(requestID, conversationID, checkpointCompletionAction(completion))
|
||||
}
|
||||
|
||||
func (service *Service) publishCheckpointWithTerminalAction(requestID string, conversationID string, terminalAction checkpointTerminalAction) error {
|
||||
stream, ok := service.broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return fmt.Errorf("request is not active: %s", requestID)
|
||||
@@ -2160,15 +2205,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, terminalAction)
|
||||
}
|
||||
|
||||
func (service *Service) rewriteCheckpointTokenDetailsForClient(stream *ActiveStream, conversation *ConversationFile, state *agentv1.ConversationStateStructure) {
|
||||
|
||||
@@ -45,13 +45,25 @@ func (snapshot turnUsageSnapshot) requestTokensTotal() int64 {
|
||||
return snapshot.promptTokensTotal() + nonNegativeInt64(snapshot.OutputTokens)
|
||||
}
|
||||
|
||||
func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure) ([]HistoryEntry, error) {
|
||||
func (service *Service) importConversationState(item *ConversationFile, state *agentv1.ConversationStateStructure, prefetchedBlobs []*agentv1.PreFetchedBlob) ([]HistoryEntry, error) {
|
||||
if item == nil || state == nil {
|
||||
return nil, nil
|
||||
}
|
||||
blobs, err := newImportedBlobStore(prefetchedBlobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
importedIDs, err := importedTurnIDs(state.GetTurns(), blobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
item.TokenDetailsUsedTokens = state.GetTokenDetails().GetUsedTokens()
|
||||
item.ImportedTurnIDs = importedIDs
|
||||
if minimumNextTurnSeq := int64(len(item.ImportedTurnIDs)) + 1; item.NextTurnSeq < minimumNextTurnSeq {
|
||||
item.NextTurnSeq = minimumNextTurnSeq
|
||||
}
|
||||
entries := make([]HistoryEntry, 0, 2)
|
||||
if messages, err := importedConversationStateModelMessages(state); err != nil {
|
||||
if messages, err := importedConversationStateModelMessages(state, blobs); err != nil {
|
||||
return nil, err
|
||||
} else {
|
||||
for _, message := range messages {
|
||||
@@ -104,7 +116,7 @@ func (service *Service) importConversationState(item *ConversationFile, state *a
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure) ([]modeladapter.Message, error) {
|
||||
func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure, blobs importedBlobStore) ([]modeladapter.Message, error) {
|
||||
if state == nil {
|
||||
return nil, nil
|
||||
}
|
||||
@@ -113,7 +125,7 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("decode imported replay messages: %w", err)
|
||||
}
|
||||
decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns())
|
||||
decoded = restoreImportedReplayUserMessages(decoded, state.GetTurns(), blobs)
|
||||
decoded = filterLegacyPlainWriteReplay(decoded)
|
||||
decoded = filterInternalPromptContextReplay(decoded)
|
||||
messages := make([]modeladapter.Message, 0, len(decoded))
|
||||
@@ -133,35 +145,18 @@ func importedConversationStateModelMessages(state *agentv1.ConversationStateStru
|
||||
if len(rawTurn) == 0 {
|
||||
continue
|
||||
}
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(rawTurn, turn); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn: %w", err)
|
||||
turn, turnID, err := decodeImportedTurn(rawTurn, blobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
agentTurn := turn.GetAgentConversationTurn()
|
||||
if agentTurn == nil {
|
||||
continue
|
||||
if turn == nil && len(turnID) > 0 {
|
||||
return nil, fmt.Errorf("missing prefetched turn blob %x", turnID)
|
||||
}
|
||||
if rawUser := agentTurn.GetUserMessage(); len(rawUser) > 0 {
|
||||
userMessage := &agentv1.UserMessage{}
|
||||
if err := proto.Unmarshal(rawUser, userMessage); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn user_message: %w", err)
|
||||
}
|
||||
if replay, ok := promptengine.BuildUserMessageReplayMessage(userMessage); ok {
|
||||
messages = append(messages, toModelMessage(replay))
|
||||
}
|
||||
}
|
||||
for _, rawStep := range agentTurn.GetSteps() {
|
||||
if len(rawStep) == 0 {
|
||||
continue
|
||||
}
|
||||
step := &agentv1.ConversationStep{}
|
||||
if err := proto.Unmarshal(rawStep, step); err != nil {
|
||||
return nil, fmt.Errorf("decode imported turn step: %w", err)
|
||||
}
|
||||
for _, replay := range promptengine.BuildLegacyMessagesFromConversationStep(step) {
|
||||
messages = append(messages, toModelMessage(replay))
|
||||
}
|
||||
turnMessages, err := importedBlobTurnMessages(turn, blobs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
messages = append(messages, turnMessages...)
|
||||
}
|
||||
return normalizeReplayMessageSequence(messages), nil
|
||||
}
|
||||
|
||||
@@ -39,6 +39,7 @@ type ConversationFile struct {
|
||||
CurrentPlanText string `json:"current_plan_text,omitempty"`
|
||||
CurrentPlans map[string]*agentv1.PlanRegistryEntry `json:"current_plans,omitempty"`
|
||||
CurrentTodos []*agentv1.TodoItem `json:"current_todos,omitempty"`
|
||||
ImportedTurnIDs [][]byte `json:"imported_turn_ids,omitempty"`
|
||||
LatestRequestPrefix *ConversationRequestPrefix `json:"latest_request_prefix,omitempty"`
|
||||
LastProviderCall *ConversationProviderCall `json:"last_provider_call,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
@@ -163,6 +164,11 @@ type ActiveStream struct {
|
||||
ProviderUsage turnUsageSnapshot
|
||||
ProviderTerminalToolInvocation bool
|
||||
PendingCompaction *PendingCompaction
|
||||
PendingCheckpointBlobWrites map[uint32]pendingCheckpointBlobWrite
|
||||
PendingCheckpointBlobRequests map[string]uint32
|
||||
NextCheckpointBlobRequestID uint32
|
||||
NextCheckpointRevision uint64
|
||||
PendingCheckpoint *pendingCheckpointPublish
|
||||
|
||||
Backlog []StreamEvent
|
||||
Subscribers map[string]*StreamSubscriber
|
||||
@@ -219,6 +225,32 @@ type pendingTurnCompletion struct {
|
||||
Disposition pendingCompletionDisposition
|
||||
}
|
||||
|
||||
type pendingCheckpointBlobWrite struct {
|
||||
Key string
|
||||
Revision uint64
|
||||
}
|
||||
|
||||
type checkpointTerminalActionKind uint8
|
||||
|
||||
const (
|
||||
checkpointTerminalActionNone checkpointTerminalActionKind = iota
|
||||
checkpointTerminalActionComplete
|
||||
checkpointTerminalActionCancel
|
||||
)
|
||||
|
||||
type checkpointTerminalAction struct {
|
||||
kind checkpointTerminalActionKind
|
||||
completion pendingTurnCompletion
|
||||
cancelMessage string
|
||||
}
|
||||
|
||||
type pendingCheckpointPublish struct {
|
||||
Revision uint64
|
||||
State *agentv1.ConversationStateStructure
|
||||
Required map[string]struct{}
|
||||
TerminalAction checkpointTerminalAction
|
||||
}
|
||||
|
||||
type PendingCompaction struct {
|
||||
Trigger string
|
||||
ContextTokens int64
|
||||
@@ -418,6 +450,7 @@ type InboundIntent struct {
|
||||
SubagentTypeName string
|
||||
SubagentModelOverrides map[string]runtimecore.SubagentModelOverrideSelection
|
||||
ConversationState *agentv1.ConversationStateStructure
|
||||
PreFetchedBlobs []*agentv1.PreFetchedBlob
|
||||
UserMessage *agentv1.UserMessage
|
||||
RequestContext *agentv1.RequestContext
|
||||
ClientMessage *agentv1.AgentClientMessage
|
||||
|
||||
Reference in New Issue
Block a user