mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-08-18 03:57:06 +08:00
Merge branch 'main' into codex/pr270-content-addressed-read-images
# Conflicts: # frontend/src/i18n/generated/catalog.json
This commit is contained in:
@@ -215,9 +215,9 @@ func Run(resources EmbeddedResources) error {
|
||||
mainWindow = app.Window.NewWithOptions(application.WebviewWindowOptions{
|
||||
Title: appName,
|
||||
Width: 700,
|
||||
Height: 520,
|
||||
MinWidth: 640,
|
||||
MinHeight: 480,
|
||||
Height: 530,
|
||||
MinWidth: 700,
|
||||
MinHeight: 530,
|
||||
DisableResize: false,
|
||||
Frameless: goruntime.GOOS == "windows",
|
||||
URL: "/",
|
||||
|
||||
@@ -1143,10 +1143,20 @@ func isAnthropicCacheableBlock(block map[string]any) bool {
|
||||
}
|
||||
}
|
||||
|
||||
// anthropicThinkingCarrier 记录请求内最近一个有 reasoning+signature 的 assistant 轮次。
|
||||
// thinking 模式下上游要求每个 assistant 轮次都回传 thinking 块;当某轮次(如 DeepSeek
|
||||
// adaptive thinking 跳过思考的 tool-call 轮次)没有 reasoning 时,用 carrier 的
|
||||
// thinking+signature 兜底,避免上游 "thinking must be passed back" 400。
|
||||
type anthropicThinkingCarrier struct {
|
||||
reasoning string
|
||||
signature string
|
||||
}
|
||||
|
||||
func normalizeAnthropicProviderMessages(input []Message, thinkingEnabled bool, relocateImages bool) ([]string, []anthropicMessage, error) {
|
||||
systemParts := make([]string, 0, len(input))
|
||||
messages := make([]anthropicMessage, 0, len(input))
|
||||
pendingToolResults := make([]map[string]any, 0, 2)
|
||||
var thinkingCarrier *anthropicThinkingCarrier
|
||||
flushToolResults := func() {
|
||||
if len(pendingToolResults) == 0 {
|
||||
return
|
||||
@@ -1192,7 +1202,15 @@ func normalizeAnthropicProviderMessages(input []Message, thinkingEnabled bool, r
|
||||
})
|
||||
case "user", "assistant":
|
||||
flushToolResults()
|
||||
contentBlocks, err := anthropicProviderContentBlocks(message, thinkingEnabled)
|
||||
if thinkingEnabled && role == "assistant" {
|
||||
if reasoning := strings.TrimSpace(message.ReasoningContent); reasoning != "" {
|
||||
thinkingCarrier = &anthropicThinkingCarrier{
|
||||
reasoning: reasoning,
|
||||
signature: anthropicThinkingSignature(message),
|
||||
}
|
||||
}
|
||||
}
|
||||
contentBlocks, err := anthropicProviderContentBlocks(message, thinkingEnabled, thinkingCarrier)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
@@ -1301,7 +1319,7 @@ func isAnthropicImageBlock(block map[string]any) bool {
|
||||
return strings.TrimSpace(anthropicStringField(block, "type")) == "image"
|
||||
}
|
||||
|
||||
func anthropicProviderContentBlocks(message Message, thinkingEnabled bool) ([]map[string]any, error) {
|
||||
func anthropicProviderContentBlocks(message Message, thinkingEnabled bool, carrier *anthropicThinkingCarrier) ([]map[string]any, error) {
|
||||
blocks, err := anthropicContentBlocks(message)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1310,11 +1328,17 @@ func anthropicProviderContentBlocks(message Message, thinkingEnabled bool) ([]ma
|
||||
return blocks, nil
|
||||
}
|
||||
|
||||
reasoning := strings.TrimSpace(message.ReasoningContent)
|
||||
signature := anthropicThinkingSignature(message)
|
||||
if reasoning == "" && carrier != nil {
|
||||
reasoning = carrier.reasoning
|
||||
signature = carrier.signature
|
||||
}
|
||||
thinkingBlock := map[string]any{
|
||||
"type": "thinking",
|
||||
"thinking": message.ReasoningContent,
|
||||
"thinking": reasoning,
|
||||
}
|
||||
if signature := anthropicThinkingSignature(message); signature != "" {
|
||||
if signature != "" {
|
||||
thinkingBlock["signature"] = signature
|
||||
}
|
||||
return append([]map[string]any{thinkingBlock}, blocks...), nil
|
||||
@@ -1398,9 +1422,6 @@ func shouldIncludeAnthropicThinkingBlock(message Message, thinkingEnabled bool)
|
||||
if strings.TrimSpace(message.Role) != "assistant" {
|
||||
return false
|
||||
}
|
||||
if strings.TrimSpace(message.ReasoningContent) == "" {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,164 @@
|
||||
package modeladapter
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestNormalizeAnthropicProviderMessagesThinkingCarrier 验证 thinking 模式下,
|
||||
// 缺少 reasoning 的 assistant 轮次(如 DeepSeek adaptive thinking 跳过思考的
|
||||
// tool-call 轮次)会用请求内最近一个 carrier 的 thinking+signature 兜底,
|
||||
// 保证每个 assistant 轮次都有 thinking 块,避免上游 "thinking must be passed
|
||||
// back to the API" 400。
|
||||
func TestNormalizeAnthropicProviderMessagesThinkingCarrier(t *testing.T) {
|
||||
carrierToolCall := []ToolCallDescriptor{{
|
||||
ID: "call-2",
|
||||
Type: "function",
|
||||
Function: ToolCallFunctionShape{
|
||||
Name: "read",
|
||||
Arguments: `{}`,
|
||||
},
|
||||
}}
|
||||
input := []Message{
|
||||
{Role: "user", Content: "hello"},
|
||||
{Role: "assistant", Content: "let me check", ReasoningContent: "R1", ReasoningSignature: "S1"},
|
||||
{Role: "user", Content: "tool result 1"},
|
||||
{Role: "assistant", ToolCalls: carrierToolCall}, // 无 reasoning → 用 carrier
|
||||
{Role: "user", Content: "tool result 2"},
|
||||
{Role: "assistant", Content: "done", ReasoningContent: "R2", ReasoningSignature: "S2"},
|
||||
}
|
||||
|
||||
_, messages, err := normalizeAnthropicProviderMessages(input, true, false)
|
||||
if err != nil {
|
||||
t.Fatalf("normalize: %v", err)
|
||||
}
|
||||
if len(messages) != 6 {
|
||||
t.Fatalf("expected 6 messages, got %d", len(messages))
|
||||
}
|
||||
|
||||
// 第 2 条(有 reasoning)应保留自己的 thinking。
|
||||
assertAnthropicThinkingBlock(t, messages[1], "R1", "S1")
|
||||
// 第 4 条(无 reasoning 的 tool-call 轮次)应复用 carrier 的 thinking+signature。
|
||||
assertAnthropicThinkingBlock(t, messages[3], "R1", "S1")
|
||||
// 第 5 条应为 tool_result 消息(合并路径不适用时,tool-call 轮次独立成消息)。
|
||||
if role := messages[4].Role; role != "user" {
|
||||
t.Fatalf("expected messages[4] role=user, got %s", role)
|
||||
}
|
||||
// 第 6 条有自己的 thinking。
|
||||
assertAnthropicThinkingBlock(t, messages[5], "R2", "S2")
|
||||
|
||||
// tool-call 轮次应包含 tool_use 块。
|
||||
hasToolUse := false
|
||||
for _, block := range messages[3].Content {
|
||||
if strings.TrimSpace(anthropicStringField(block, "type")) == "tool_use" {
|
||||
hasToolUse = true
|
||||
}
|
||||
}
|
||||
if !hasToolUse {
|
||||
t.Fatal("expected tool_use block on the carrier-fallback assistant message")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNormalizeAnthropicProviderMessagesThinkingCarrierFirstTurn 验证请求内第一条
|
||||
// assistant 轮次就缺 reasoning 且无 carrier 时,兜底输出空 thinking 块。
|
||||
func TestNormalizeAnthropicProviderMessagesThinkingCarrierFirstTurn(t *testing.T) {
|
||||
input := []Message{
|
||||
{Role: "user", Content: "hello"},
|
||||
{Role: "assistant", Content: "ok"},
|
||||
}
|
||||
|
||||
_, messages, err := normalizeAnthropicProviderMessages(input, true, false)
|
||||
if err != nil {
|
||||
t.Fatalf("normalize: %v", err)
|
||||
}
|
||||
if len(messages) != 2 {
|
||||
t.Fatalf("expected 2 messages, got %d", len(messages))
|
||||
}
|
||||
if got := anthropicStringField(messages[1].Content[0], "type"); got != "thinking" {
|
||||
t.Fatalf("expected first block type=thinking, got %s", got)
|
||||
}
|
||||
if got := anthropicStringField(messages[1].Content[0], "thinking"); got != "" {
|
||||
t.Fatalf("expected empty fallback thinking, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNormalizeAnthropicProviderMessagesThinkingDisabled 验证 thinking 关闭时
|
||||
// 不输出任何 thinking 块(回归保护)。
|
||||
func TestNormalizeAnthropicProviderMessagesThinkingDisabled(t *testing.T) {
|
||||
input := []Message{
|
||||
{Role: "user", Content: "hello"},
|
||||
{Role: "assistant", Content: "ok", ReasoningContent: "R1", ReasoningSignature: "S1"},
|
||||
{Role: "assistant", ToolCalls: []ToolCallDescriptor{{
|
||||
ID: "call-2",
|
||||
Type: "function",
|
||||
Function: ToolCallFunctionShape{
|
||||
Name: "read",
|
||||
Arguments: `{}`,
|
||||
},
|
||||
}}},
|
||||
}
|
||||
|
||||
_, messages, err := normalizeAnthropicProviderMessages(input, false, false)
|
||||
if err != nil {
|
||||
t.Fatalf("normalize: %v", err)
|
||||
}
|
||||
for index, message := range messages {
|
||||
for _, block := range message.Content {
|
||||
if blockType := anthropicStringField(block, "type"); blockType == "thinking" {
|
||||
t.Fatalf("unexpected thinking block at messages[%d]", index)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestNormalizeAnthropicProviderMessagesThinkingMerge 验证有 reasoning 的纯
|
||||
// tool-call 轮次仍按既有逻辑合并进上一条 assistant 消息(thinking 去重,无回归)。
|
||||
func TestNormalizeAnthropicProviderMessagesThinkingMerge(t *testing.T) {
|
||||
input := []Message{
|
||||
{Role: "user", Content: "hello"},
|
||||
{Role: "assistant", Content: "let me check", ReasoningContent: "R1", ReasoningSignature: "S1"},
|
||||
{Role: "assistant", ToolCalls: []ToolCallDescriptor{{
|
||||
ID: "call-2",
|
||||
Type: "function",
|
||||
Function: ToolCallFunctionShape{
|
||||
Name: "read",
|
||||
Arguments: `{}`,
|
||||
},
|
||||
}}, ReasoningContent: "R1", ReasoningSignature: "S1"},
|
||||
}
|
||||
|
||||
_, messages, err := normalizeAnthropicProviderMessages(input, true, false)
|
||||
if err != nil {
|
||||
t.Fatalf("normalize: %v", err)
|
||||
}
|
||||
if len(messages) != 2 {
|
||||
t.Fatalf("expected 2 messages (tool-call merged), got %d", len(messages))
|
||||
}
|
||||
assertAnthropicThinkingBlock(t, messages[1], "R1", "S1")
|
||||
hasToolUse := false
|
||||
for _, block := range messages[1].Content {
|
||||
if blockType := anthropicStringField(block, "type"); blockType == "tool_use" {
|
||||
hasToolUse = true
|
||||
}
|
||||
}
|
||||
if !hasToolUse {
|
||||
t.Fatal("expected merged tool_use block on messages[1]")
|
||||
}
|
||||
}
|
||||
|
||||
func assertAnthropicThinkingBlock(t *testing.T, message anthropicMessage, wantThinking string, wantSignature string) {
|
||||
t.Helper()
|
||||
if len(message.Content) == 0 {
|
||||
t.Fatalf("expected non-empty content for %s message", message.Role)
|
||||
}
|
||||
first := message.Content[0]
|
||||
if blockType := anthropicStringField(first, "type"); blockType != "thinking" {
|
||||
t.Fatalf("expected first block type=thinking, got %s", blockType)
|
||||
}
|
||||
if got := anthropicStringField(first, "thinking"); got != wantThinking {
|
||||
t.Fatalf("expected thinking=%q, got %q", wantThinking, got)
|
||||
}
|
||||
if got := anthropicStringField(first, "signature"); got != wantSignature {
|
||||
t.Fatalf("expected signature=%q, got %q", wantSignature, got)
|
||||
}
|
||||
}
|
||||
@@ -49,7 +49,8 @@ type openAIResponsesRequestBody struct {
|
||||
}
|
||||
|
||||
type openAIResponsesReasoning struct {
|
||||
Effort string `json:"effort,omitempty"`
|
||||
Effort string `json:"effort,omitempty"`
|
||||
Summary string `json:"summary,omitempty"`
|
||||
}
|
||||
|
||||
type openAIToolAccumulator struct {
|
||||
@@ -839,7 +840,7 @@ func (adapter *OpenAIAdapter) streamChatCompletions(ctx context.Context, req Str
|
||||
}
|
||||
}
|
||||
|
||||
if choice.FinishReason != nil {
|
||||
if choice.FinishReason != nil && strings.TrimSpace(*choice.FinishReason) != "" {
|
||||
if err := flushTaggedContentTail(); err != nil {
|
||||
return fail(err)
|
||||
}
|
||||
@@ -944,7 +945,7 @@ func (adapter *OpenAIAdapter) streamResponses(ctx context.Context, req StreamReq
|
||||
requestBody.Tools = tools
|
||||
}
|
||||
if effort := strings.TrimSpace(req.ReasoningEffort); effort != "" {
|
||||
requestBody.Reasoning = &openAIResponsesReasoning{Effort: effort}
|
||||
requestBody.Reasoning = &openAIResponsesReasoning{Effort: effort, Summary: "auto"}
|
||||
requestBody.Include = []string{"reasoning.encrypted_content"}
|
||||
}
|
||||
body = requestBody
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
package modeladapter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestOpenAIResponsesRequestsReasoningSummary(t *testing.T) {
|
||||
var requestBody map[string]any
|
||||
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
if err := json.NewDecoder(request.Body).Decode(&requestBody); err != nil {
|
||||
http.Error(writer, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
writer.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = fmt.Fprint(writer, "data: {\"type\":\"response.reasoning_summary_text.delta\",\"delta\":\"check the request\"}\n\n")
|
||||
_, _ = fmt.Fprint(writer, "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-1\",\"model\":\"gpt-5.6\",\"status\":\"completed\",\"output_text\":\"done\"}}\n\n")
|
||||
_, _ = fmt.Fprint(writer, "data: [DONE]\n\n")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
adapter := &OpenAIAdapter{client: server.Client()}
|
||||
events := make([]ModelEvent, 0, 4)
|
||||
err := adapter.Stream(context.Background(), StreamRequest{
|
||||
RequestID: "request-1",
|
||||
RunID: "run-1",
|
||||
ModelCallID: "model-call-1",
|
||||
BaseURL: server.URL,
|
||||
APIKey: "test-key",
|
||||
ProviderModelID: "gpt-5.6",
|
||||
OpenAIEndpoint: "/v1/responses",
|
||||
ReasoningEffort: "high",
|
||||
Messages: []Message{{Role: "user", Content: "hello"}},
|
||||
MaxTokens: 128,
|
||||
}, func(event ModelEvent) error {
|
||||
events = append(events, event)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("stream failed: %v", err)
|
||||
}
|
||||
|
||||
reasoning, ok := requestBody["reasoning"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("reasoning request body missing: %#v", requestBody)
|
||||
}
|
||||
if got := reasoning["effort"]; got != "high" {
|
||||
t.Fatalf("reasoning.effort = %#v, want high", got)
|
||||
}
|
||||
if got := reasoning["summary"]; got != "auto" {
|
||||
t.Fatalf("reasoning.summary = %#v, want auto", got)
|
||||
}
|
||||
include, ok := requestBody["include"].([]any)
|
||||
if !ok || len(include) != 1 || include[0] != "reasoning.encrypted_content" {
|
||||
t.Fatalf("reasoning include = %#v, want encrypted content", requestBody["include"])
|
||||
}
|
||||
assertOpenAIEventKindCount(t, events, ModelEventKindThinkingDelta, 1)
|
||||
assertOpenAIEventKindCount(t, events, ModelEventKindThinkingCompleted, 1)
|
||||
assertOpenAIEventKindCount(t, events, ModelEventKindTextDelta, 1)
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesOmitsReasoningWhenEffortBlank(t *testing.T) {
|
||||
var requestBody map[string]any
|
||||
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
if err := json.NewDecoder(request.Body).Decode(&requestBody); err != nil {
|
||||
http.Error(writer, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
writer.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = fmt.Fprint(writer, "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-1\",\"model\":\"grok-composer-2.5-fast\",\"status\":\"completed\",\"output_text\":\"done\"}}\n\n")
|
||||
_, _ = fmt.Fprint(writer, "data: [DONE]\n\n")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
adapter := &OpenAIAdapter{client: server.Client()}
|
||||
err := adapter.Stream(context.Background(), StreamRequest{
|
||||
RequestID: "request-1",
|
||||
RunID: "run-1",
|
||||
ModelCallID: "model-call-1",
|
||||
BaseURL: server.URL,
|
||||
APIKey: "test-key",
|
||||
ProviderModelID: "grok-composer-2.5-fast",
|
||||
OpenAIEndpoint: "/v1/responses",
|
||||
Messages: []Message{{Role: "user", Content: "hello"}},
|
||||
MaxTokens: 128,
|
||||
}, func(ModelEvent) error { return nil })
|
||||
if err != nil {
|
||||
t.Fatalf("stream failed: %v", err)
|
||||
}
|
||||
|
||||
if _, exists := requestBody["reasoning"]; exists {
|
||||
t.Fatalf("reasoning should be omitted when effort is blank: %#v", requestBody["reasoning"])
|
||||
}
|
||||
if _, exists := requestBody["reasoning_effort"]; exists {
|
||||
t.Fatalf("reasoning_effort should be omitted when effort is blank: %#v", requestBody["reasoning_effort"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIChatCompletionsOmitsReasoningWhenEffortBlank(t *testing.T) {
|
||||
var requestBody map[string]any
|
||||
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
if err := json.NewDecoder(request.Body).Decode(&requestBody); err != nil {
|
||||
http.Error(writer, err.Error(), http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
writer.Header().Set("Content-Type", "text/event-stream")
|
||||
_, _ = fmt.Fprint(writer, "data: {\"model\":\"grok-composer-2.5-fast\",\"choices\":[{\"delta\":{\"content\":\"done\"},\"finish_reason\":\"stop\"}]}\n\n")
|
||||
_, _ = fmt.Fprint(writer, "data: [DONE]\n\n")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
adapter := &OpenAIAdapter{client: server.Client()}
|
||||
err := adapter.Stream(context.Background(), StreamRequest{
|
||||
RequestID: "request-1",
|
||||
RunID: "run-1",
|
||||
ModelCallID: "model-call-1",
|
||||
BaseURL: server.URL,
|
||||
APIKey: "test-key",
|
||||
ProviderModelID: "grok-composer-2.5-fast",
|
||||
OpenAIEndpoint: "/v1/chat/completions",
|
||||
Messages: []Message{{Role: "user", Content: "hello"}},
|
||||
MaxTokens: 128,
|
||||
}, func(ModelEvent) error { return nil })
|
||||
if err != nil {
|
||||
t.Fatalf("stream failed: %v", err)
|
||||
}
|
||||
|
||||
for _, field := range []string{"reasoning_effort", "reasoning", "include"} {
|
||||
if value, exists := requestBody[field]; exists {
|
||||
t.Fatalf("%s should be omitted when effort is blank: %#v", field, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIChatCompletionsIgnoresBlankFinishReason(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
writer.Header().Set("Content-Type", "text/event-stream")
|
||||
chunks := []string{
|
||||
`{"model":"deepseek-v4-flash","choices":[{"delta":{"reasoning_content":"first"},"finish_reason":""}]}`,
|
||||
`{"model":"deepseek-v4-flash","choices":[{"delta":{"reasoning_content":" second"},"finish_reason":""}]}`,
|
||||
`{"model":"deepseek-v4-flash","choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"Ls","arguments":""}}]},"finish_reason":""}]}`,
|
||||
`{"model":"deepseek-v4-flash","choices":[{"delta":{"tool_calls":[{"index":0,"id":"","type":"function","function":{"name":"","arguments":"{\"path\":"}}]},"finish_reason":""}]}`,
|
||||
`{"model":"deepseek-v4-flash","choices":[{"delta":{"tool_calls":[{"index":0,"id":"","type":"function","function":{"name":"","arguments":"\"/tmp\"}"}}]},"finish_reason":"tool_calls"}]}`,
|
||||
`{"model":"deepseek-v4-flash","choices":[],"usage":{"prompt_tokens":12,"completion_tokens":7}}`,
|
||||
}
|
||||
for _, chunk := range chunks {
|
||||
_, _ = fmt.Fprintf(writer, "data: %s\n\n", chunk)
|
||||
}
|
||||
_, _ = fmt.Fprint(writer, "data: [DONE]\n\n")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
adapter := &OpenAIAdapter{client: server.Client()}
|
||||
events := make([]ModelEvent, 0, 8)
|
||||
err := adapter.Stream(context.Background(), StreamRequest{
|
||||
RequestID: "request-1",
|
||||
RunID: "run-1",
|
||||
ModelCallID: "model-call-1",
|
||||
BaseURL: server.URL,
|
||||
APIKey: "test-key",
|
||||
ProviderModelID: "deepseek-v4-flash",
|
||||
OpenAIEndpoint: "/v1/chat/completions",
|
||||
Messages: []Message{{Role: "user", Content: "list files"}},
|
||||
MaxTokens: 128,
|
||||
RequestKnobs: map[string]any{},
|
||||
}, func(event ModelEvent) error {
|
||||
events = append(events, event)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("stream failed: %v", err)
|
||||
}
|
||||
|
||||
assertOpenAIEventKindCount(t, events, ModelEventKindThinkingDelta, 2)
|
||||
assertOpenAIEventKindCount(t, events, ModelEventKindThinkingCompleted, 1)
|
||||
assertOpenAIEventKindCount(t, events, ModelEventKindToolLikeCompleted, 1)
|
||||
assertOpenAIEventKindCount(t, events, ModelEventKindTurnFinished, 1)
|
||||
|
||||
toolEvent := firstOpenAIEventForTest(events, ModelEventKindToolLikeCompleted)
|
||||
if toolEvent == nil || toolEvent.ToolInvocation == nil {
|
||||
t.Fatalf("completed tool invocation missing: %#v", toolEvent)
|
||||
}
|
||||
if toolEvent.ToolInvocation.ToolName != "Ls" {
|
||||
t.Fatalf("tool name = %q, want Ls", toolEvent.ToolInvocation.ToolName)
|
||||
}
|
||||
if got := string(toolEvent.ToolInvocation.ArgsJSON); got != `{"path":"/tmp"}` {
|
||||
t.Fatalf("tool args = %q, want %q", got, `{"path":"/tmp"}`)
|
||||
}
|
||||
|
||||
finished := firstOpenAIEventForTest(events, ModelEventKindTurnFinished)
|
||||
if finished == nil || finished.FinishReason != "tool_calls" {
|
||||
t.Fatalf("finish reason = %q, want tool_calls", finished.FinishReason)
|
||||
}
|
||||
if finished.InputTokens != 12 || finished.OutputTokens != 7 {
|
||||
t.Fatalf("usage = input:%d output:%d, want input:12 output:7", finished.InputTokens, finished.OutputTokens)
|
||||
}
|
||||
}
|
||||
|
||||
func assertOpenAIEventKindCount(t *testing.T, events []ModelEvent, kind ModelEventKind, want int) {
|
||||
t.Helper()
|
||||
got := 0
|
||||
for _, event := range events {
|
||||
if event.Kind == kind {
|
||||
got++
|
||||
}
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("event kind %s count = %d, want %d; events=%#v", kind, got, want, events)
|
||||
}
|
||||
}
|
||||
|
||||
func firstOpenAIEventForTest(events []ModelEvent, kind ModelEventKind) *ModelEvent {
|
||||
for index := range events {
|
||||
if events[index].Kind == kind {
|
||||
return &events[index]
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -1,10 +1,64 @@
|
||||
package modeladapter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
legacyruntime "cursor/internal/runtime"
|
||||
)
|
||||
|
||||
type recordingModelAdapter struct {
|
||||
request StreamRequest
|
||||
}
|
||||
|
||||
func (adapter *recordingModelAdapter) Stream(_ context.Context, req StreamRequest, _ func(ModelEvent) error) error {
|
||||
adapter.request = req
|
||||
return nil
|
||||
}
|
||||
|
||||
type staticChannelResolver struct {
|
||||
channel *legacyruntime.ResolvedChannel
|
||||
}
|
||||
|
||||
func (resolver staticChannelResolver) SelectChannelForModel(context.Context, string) (*legacyruntime.ResolvedChannel, error) {
|
||||
return resolver.channel, nil
|
||||
}
|
||||
|
||||
func (staticChannelResolver) ProviderStreamIdleTimeout(context.Context) time.Duration {
|
||||
return time.Second
|
||||
}
|
||||
|
||||
func TestRouterRuntimeDisabledClearsReasoningEffort(t *testing.T) {
|
||||
openAI := &recordingModelAdapter{}
|
||||
router := &Router{
|
||||
openai: openAI,
|
||||
resolver: staticChannelResolver{channel: &legacyruntime.ResolvedChannel{
|
||||
ID: "channel-a",
|
||||
Provider: "openai",
|
||||
Model: "grok-composer-2.5-fast",
|
||||
ReasoningEffort: "medium",
|
||||
}},
|
||||
}
|
||||
requestKnobs := map[string]any{"reasoning_effort": "medium"}
|
||||
|
||||
err := router.Stream(context.Background(), StreamRequest{
|
||||
ModelID: "channel-a",
|
||||
ThinkingEffort: "disabled",
|
||||
RequestKnobs: requestKnobs,
|
||||
}, func(ModelEvent) error { return nil })
|
||||
if err != nil {
|
||||
t.Fatalf("Stream returned error: %v", err)
|
||||
}
|
||||
if got := openAI.request.ReasoningEffort; got != "" {
|
||||
t.Fatalf("ReasoningEffort = %q, want blank", got)
|
||||
}
|
||||
if _, exists := openAI.request.RequestKnobs["reasoning_effort"]; exists {
|
||||
t.Fatalf("reasoning_effort knob should be removed: %#v", openAI.request.RequestKnobs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeProviderMessagesMergesLegacyAssistantTextAndToolCallTurnsIdempotently(t *testing.T) {
|
||||
input := []Message{
|
||||
{
|
||||
|
||||
@@ -730,6 +730,9 @@ func (service *Service) handleProviderDoneEvent(stream *ActiveStream, payload *s
|
||||
service.setTurnPhase(stream, TurnPhaseFailed)
|
||||
return service.closeStreamWithProviderError(stream, conversationID, turnSeq, requestID, accumulatedText, accumulatedReasoning, accumulatedReasoningSignature, accumulatedReasoningSignatureSource, accumulatedReasoningItemID, accumulatedReasoningStatus, accumulatedReasoningSummary, usage, providerErr, !hadToolInvocation)
|
||||
}
|
||||
if err := service.flushAssistantText(stream, conversationID, turnSeq, requestID, accumulatedText, accumulatedReasoning, accumulatedReasoningSignature, accumulatedReasoningSignatureSource, accumulatedReasoningItemID, accumulatedReasoningStatus, accumulatedReasoningSummary, !hadToolInvocation); err != nil {
|
||||
return service.failStream(stream, "unknown", fmt.Errorf("flush failed provider output: %w", err))
|
||||
}
|
||||
service.setTurnPhase(stream, TurnPhaseFailed)
|
||||
return service.failStream(stream, "unknown", payload.Err)
|
||||
}
|
||||
|
||||
@@ -19,15 +19,29 @@ type pendingCheckpointBlobWrite struct {
|
||||
blob CheckpointBlob
|
||||
}
|
||||
|
||||
func clonePendingTurnCompletion(completion *pendingTurnCompletion) *pendingTurnCompletion {
|
||||
func successfulCheckpointTerminalAction(completion *pendingTurnCompletion) checkpointTerminalAction {
|
||||
if completion == nil {
|
||||
return nil
|
||||
return checkpointTerminalAction{}
|
||||
}
|
||||
return checkpointTerminalAction{
|
||||
Kind: checkpointTerminalActionComplete,
|
||||
Completion: *completion,
|
||||
}
|
||||
}
|
||||
|
||||
func failedCheckpointTerminalAction(errorCode string, errorMessage string) checkpointTerminalAction {
|
||||
return checkpointTerminalAction{
|
||||
Kind: checkpointTerminalActionFail,
|
||||
ErrorCode: strings.TrimSpace(errorCode),
|
||||
ErrorMessage: strings.TrimSpace(errorMessage),
|
||||
}
|
||||
cloned := *completion
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func (service *Service) queueCheckpointProjection(stream *ActiveStream, projection *CheckpointProjection, completion *pendingTurnCompletion) error {
|
||||
return service.queueCheckpointProjectionWithTerminal(stream, projection, successfulCheckpointTerminalAction(completion))
|
||||
}
|
||||
|
||||
func (service *Service) queueCheckpointProjectionWithTerminal(stream *ActiveStream, projection *CheckpointProjection, terminal checkpointTerminalAction) error {
|
||||
if service == nil || stream == nil || projection == nil || projection.State == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -43,8 +57,8 @@ func (service *Service) queueCheckpointProjection(stream *ActiveStream, projecti
|
||||
if stream.ConfirmedCheckpointBlobs == nil {
|
||||
stream.ConfirmedCheckpointBlobs = make(map[string]struct{})
|
||||
}
|
||||
if completion == nil && stream.PendingCheckpoint != nil {
|
||||
completion = stream.PendingCheckpoint.Completion
|
||||
if terminal.Kind == checkpointTerminalActionNone && stream.PendingCheckpoint != nil {
|
||||
terminal = stream.PendingCheckpoint.Terminal
|
||||
}
|
||||
required := make(map[string]struct{}, len(projection.Blobs))
|
||||
pendingKeys := make(map[string]struct{}, len(stream.PendingCheckpointBlobWrites))
|
||||
@@ -74,11 +88,11 @@ func (service *Service) queueCheckpointProjection(stream *ActiveStream, projecti
|
||||
toWrite = append(toWrite, pendingCheckpointBlobWrite{requestID: requestID, blob: blob})
|
||||
}
|
||||
stream.PendingCheckpoint = &pendingCheckpointPublish{
|
||||
State: state,
|
||||
Required: required,
|
||||
Completion: clonePendingTurnCompletion(completion),
|
||||
State: state,
|
||||
Required: required,
|
||||
Terminal: terminal,
|
||||
}
|
||||
if completion != nil {
|
||||
if terminal.Kind != checkpointTerminalActionNone {
|
||||
stream.Phase = TurnPhaseCheckpointing
|
||||
}
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
@@ -94,6 +108,8 @@ func (service *Service) queueCheckpointProjection(stream *ActiveStream, projecti
|
||||
if service.checkpointProjectionReady(stream) {
|
||||
return service.publishReadyCheckpoint(stream)
|
||||
}
|
||||
// Checkpoints reference these Blob IDs, so the client must confirm every
|
||||
// required Blob before the checkpoint becomes visible.
|
||||
service.scheduleStreamTimer(
|
||||
stream,
|
||||
providerTimerKey(streamTimerCheckpointBlobs, ""),
|
||||
@@ -175,21 +191,18 @@ func (service *Service) publishReadyCheckpoint(stream *ActiveStream) error {
|
||||
}
|
||||
stream.PendingCheckpoint = nil
|
||||
state := pending.State
|
||||
completion := clonePendingTurnCompletion(pending.Completion)
|
||||
terminal := pending.Terminal
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
clearStreamTimer(stream, providerTimerKey(streamTimerCheckpointBlobs, ""))
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{Message: buildCheckpointMessage(state)}); err != nil {
|
||||
if completion != nil {
|
||||
log.Printf("forwarder checkpoint publish skipped before successful terminal request_id=%s err=%v", stream.RequestID, err)
|
||||
return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion)
|
||||
if terminal.Kind != checkpointTerminalActionNone {
|
||||
log.Printf("forwarder checkpoint publish skipped before terminal request_id=%s err=%v", stream.RequestID, err)
|
||||
return service.finishCheckpointTerminalAction(stream, terminal)
|
||||
}
|
||||
return err
|
||||
}
|
||||
if completion != nil {
|
||||
return service.finishSuccessfulTurnAfterCheckpoint(stream, *completion)
|
||||
}
|
||||
return nil
|
||||
return service.finishCheckpointTerminalAction(stream, terminal)
|
||||
}
|
||||
|
||||
func (service *Service) handleCheckpointBlobTimeout(stream *ActiveStream) error {
|
||||
@@ -216,12 +229,23 @@ func (service *Service) finishAfterCheckpointSyncFailure(stream *ActiveStream, c
|
||||
if cause != nil {
|
||||
log.Printf("forwarder checkpoint blob sync skipped request_id=%s conversation_id=%s err=%v", stream.RequestID, stream.ConversationID, cause)
|
||||
}
|
||||
if pending != nil && pending.Completion != nil {
|
||||
return service.finishSuccessfulTurnAfterCheckpoint(stream, *pending.Completion)
|
||||
if pending != nil {
|
||||
return service.finishCheckpointTerminalAction(stream, pending.Terminal)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (service *Service) finishCheckpointTerminalAction(stream *ActiveStream, terminal checkpointTerminalAction) error {
|
||||
switch terminal.Kind {
|
||||
case checkpointTerminalActionComplete:
|
||||
return service.finishSuccessfulTurnAfterCheckpoint(stream, terminal.Completion)
|
||||
case checkpointTerminalActionFail:
|
||||
return service.finishFailedTurnAfterCheckpoint(stream, terminal.ErrorCode, terminal.ErrorMessage)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) discardPendingCheckpoint(stream *ActiveStream, reason string) {
|
||||
if stream == nil {
|
||||
return
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
func TestCheckpointBlobSyncPublishesCheckpointAfterAcknowledgements(t *testing.T) {
|
||||
func TestCheckpointBlobSyncWaitsForAcknowledgementsBeforePublishingNonTerminalCheckpoint(t *testing.T) {
|
||||
service, stream, projection := testCheckpointBlobProjection(t)
|
||||
if err := service.queueCheckpointProjection(stream, projection, nil); err != nil {
|
||||
t.Fatalf("queueCheckpointProjection() error = %v", err)
|
||||
@@ -23,11 +23,40 @@ func TestCheckpointBlobSyncPublishesCheckpointAfterAcknowledgements(t *testing.T
|
||||
}
|
||||
}
|
||||
|
||||
stream.mu.Lock()
|
||||
var firstRequestID uint32
|
||||
for requestID := range stream.PendingCheckpointBlobWrites {
|
||||
firstRequestID = requestID
|
||||
break
|
||||
}
|
||||
stream.mu.Unlock()
|
||||
if firstRequestID == 0 {
|
||||
t.Fatal("checkpoint projection has no pending Blob writes")
|
||||
}
|
||||
if err := service.handleCheckpointBlobResult(stream, &agentv1.KvClientMessage{
|
||||
Id: firstRequestID,
|
||||
Message: &agentv1.KvClientMessage_SetBlobResult{
|
||||
SetBlobResult: &agentv1.SetBlobResult{},
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("first Blob ACK error = %v", err)
|
||||
}
|
||||
for _, event := range readCheckpointTestEvents(t, service, stream) {
|
||||
if event.Message.GetConversationCheckpointUpdate() != nil {
|
||||
t.Fatal("checkpoint published after only a partial Blob acknowledgement")
|
||||
}
|
||||
}
|
||||
|
||||
acknowledgeCheckpointBlobs(t, service, stream)
|
||||
events = readCheckpointTestEvents(t, service, stream)
|
||||
checkpoint := events[len(events)-1].Message.GetConversationCheckpointUpdate()
|
||||
if checkpoint == nil || len(checkpoint.GetTurns()) != 1 {
|
||||
t.Fatalf("last event checkpoint = %#v, want one Blob-backed turn", checkpoint)
|
||||
checkpointCount := 0
|
||||
for _, event := range events {
|
||||
if event.Message.GetConversationCheckpointUpdate() != nil {
|
||||
checkpointCount++
|
||||
}
|
||||
}
|
||||
if checkpointCount != 1 {
|
||||
t.Fatalf("checkpoints after ACK = %d, want 1", checkpointCount)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -40,6 +69,12 @@ func TestCheckpointBlobSyncPublishesCheckpointBeforeSuccessfulTerminal(t *testin
|
||||
if err := service.queueCheckpointProjection(stream, projection, completion); err != nil {
|
||||
t.Fatalf("queueCheckpointProjection() error = %v", err)
|
||||
}
|
||||
eventsBeforeACK := readCheckpointTestEvents(t, service, stream)
|
||||
for _, event := range eventsBeforeACK {
|
||||
if event.Message.GetConversationCheckpointUpdate() != nil || event.Message.GetInteractionUpdate().GetTurnEnded() != nil || event.End {
|
||||
t.Fatalf("event before ACK = %#v, want only Blob writes", event)
|
||||
}
|
||||
}
|
||||
acknowledgeCheckpointBlobs(t, service, stream)
|
||||
|
||||
events := readCheckpointTestEvents(t, service, stream)
|
||||
@@ -84,11 +119,137 @@ func TestCheckpointBlobTimeoutDoesNotFailSuccessfulTurn(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCancellationDiscardsPendingCheckpointAndIgnoresLateAcknowledgements(t *testing.T) {
|
||||
func TestCheckpointBlobSyncPublishesCheckpointBeforeFailedTerminal(t *testing.T) {
|
||||
service, stream, _ := testCheckpointBlobProjection(t)
|
||||
if err := service.failActiveStream(
|
||||
stream,
|
||||
stream.ConversationID,
|
||||
stream.RequestID,
|
||||
"model-call-1",
|
||||
"provider_error",
|
||||
"provider failed",
|
||||
); err != nil {
|
||||
t.Fatalf("failActiveStream() error = %v", err)
|
||||
}
|
||||
|
||||
for _, event := range readCheckpointTestEvents(t, service, stream) {
|
||||
if event.Message.GetConversationCheckpointUpdate() != nil || event.End {
|
||||
t.Fatalf("event before ACK = %#v, want only Blob writes", event)
|
||||
}
|
||||
}
|
||||
stream.mu.Lock()
|
||||
phaseBeforeACK := stream.Phase
|
||||
statusBeforeACK := stream.Status
|
||||
stream.mu.Unlock()
|
||||
if phaseBeforeACK != TurnPhaseCheckpointing || isTerminalStreamStatus(statusBeforeACK) {
|
||||
t.Fatalf("before ACK phase=%s status=%s, want checkpointing and non-terminal", phaseBeforeACK, statusBeforeACK)
|
||||
}
|
||||
|
||||
acknowledgeCheckpointBlobs(t, service, stream)
|
||||
events := readCheckpointTestEvents(t, service, stream)
|
||||
checkpointIndex, endIndex := -1, -1
|
||||
for index, event := range events {
|
||||
switch {
|
||||
case event.Message.GetConversationCheckpointUpdate() != nil:
|
||||
checkpointIndex = index
|
||||
case event.End:
|
||||
endIndex = index
|
||||
if event.TerminalErrorCode != "provider_error" || event.TerminalErrorMessage != "provider failed" {
|
||||
t.Fatalf("terminal event = %#v, want provider error", event)
|
||||
}
|
||||
}
|
||||
}
|
||||
if checkpointIndex < 0 || endIndex <= checkpointIndex {
|
||||
t.Fatalf("terminal order checkpoint=%d end=%d", checkpointIndex, endIndex)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
phaseAfterACK := stream.Phase
|
||||
statusAfterACK := stream.Status
|
||||
stream.mu.Unlock()
|
||||
if phaseAfterACK != TurnPhaseFailed || statusAfterACK != StreamStatusFailed {
|
||||
t.Fatalf("after ACK phase=%s status=%s, want failed", phaseAfterACK, statusAfterACK)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckpointBlobTimeoutStillPublishesFailedTerminal(t *testing.T) {
|
||||
service, stream, _ := testCheckpointBlobProjection(t)
|
||||
if err := service.failActiveStream(
|
||||
stream,
|
||||
stream.ConversationID,
|
||||
stream.RequestID,
|
||||
"model-call-1",
|
||||
"provider_error",
|
||||
"provider failed",
|
||||
); err != nil {
|
||||
t.Fatalf("failActiveStream() error = %v", err)
|
||||
}
|
||||
if err := service.handleCheckpointBlobTimeout(stream); err != nil {
|
||||
t.Fatalf("handleCheckpointBlobTimeout() error = %v", err)
|
||||
}
|
||||
|
||||
events := readCheckpointTestEvents(t, service, stream)
|
||||
var checkpoint, failedEnd bool
|
||||
for _, event := range events {
|
||||
checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil
|
||||
failedEnd = failedEnd || event.End && event.TerminalErrorCode == "provider_error" && event.TerminalErrorMessage == "provider failed"
|
||||
}
|
||||
if checkpoint || !failedEnd {
|
||||
t.Fatalf("timeout events checkpoint=%v failed_end=%v", checkpoint, failedEnd)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManualCompactionNoopWaitsForCheckpointBeforeTerminal(t *testing.T) {
|
||||
service, stream, _ := testCheckpointBlobProjection(t)
|
||||
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
|
||||
if err != nil {
|
||||
t.Fatalf("snapshotCheckpointConversation() error = %v", err)
|
||||
}
|
||||
if _, err := service.store.SaveConversationWithEntries(stream.ConversationID, conversation, conversation.Entries); err != nil {
|
||||
t.Fatalf("SaveConversationWithEntries() error = %v", err)
|
||||
}
|
||||
if err := service.finishManualCompactionNoop(stream); err != nil {
|
||||
t.Fatalf("finishManualCompactionNoop() error = %v", err)
|
||||
}
|
||||
|
||||
for _, event := range readCheckpointTestEvents(t, service, stream) {
|
||||
if event.Message.GetInteractionUpdate().GetTurnEnded() != nil || event.End {
|
||||
t.Fatalf("terminal event before checkpoint Blob ACK = %#v", event)
|
||||
}
|
||||
}
|
||||
acknowledgeCheckpointBlobs(t, service, stream)
|
||||
|
||||
events := readCheckpointTestEvents(t, service, stream)
|
||||
checkpointIndex, turnEndedIndex, endIndex := -1, -1, -1
|
||||
for index, event := range events {
|
||||
switch {
|
||||
case event.Message.GetConversationCheckpointUpdate() != nil:
|
||||
checkpointIndex = index
|
||||
case event.Message.GetInteractionUpdate().GetTurnEnded() != nil:
|
||||
turnEndedIndex = index
|
||||
case event.End:
|
||||
endIndex = index
|
||||
}
|
||||
}
|
||||
if checkpointIndex < 0 || turnEndedIndex <= checkpointIndex || endIndex <= turnEndedIndex {
|
||||
t.Fatalf("terminal order checkpoint=%d turn_ended=%d end=%d", checkpointIndex, turnEndedIndex, endIndex)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCancellationDiscardsUnpublishedCheckpointAndIgnoresLateAcknowledgements(t *testing.T) {
|
||||
service, stream, projection := testCheckpointBlobProjection(t)
|
||||
if err := service.queueCheckpointProjection(stream, projection, nil); err != nil {
|
||||
t.Fatalf("queueCheckpointProjection() error = %v", err)
|
||||
}
|
||||
eventsBeforeCancel := readCheckpointTestEvents(t, service, stream)
|
||||
checkpointBeforeCancel := 0
|
||||
for _, event := range eventsBeforeCancel {
|
||||
if event.Message.GetConversationCheckpointUpdate() != nil {
|
||||
checkpointBeforeCancel++
|
||||
}
|
||||
}
|
||||
if checkpointBeforeCancel != 0 {
|
||||
t.Fatalf("checkpoints before cancel = %d, want 0", checkpointBeforeCancel)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
requestIDs := make([]uint32, 0, len(stream.PendingCheckpointBlobWrites))
|
||||
for requestID := range stream.PendingCheckpointBlobWrites {
|
||||
@@ -114,16 +275,19 @@ func TestCancellationDiscardsPendingCheckpointAndIgnoresLateAcknowledgements(t *
|
||||
}
|
||||
|
||||
events := readCheckpointTestEvents(t, service, stream)
|
||||
var checkpoint, canceledEnd bool
|
||||
checkpointCount := 0
|
||||
var canceledEnd bool
|
||||
for _, event := range events {
|
||||
checkpoint = checkpoint || event.Message.GetConversationCheckpointUpdate() != nil
|
||||
if event.Message.GetConversationCheckpointUpdate() != nil {
|
||||
checkpointCount++
|
||||
}
|
||||
canceledEnd = canceledEnd || event.End && event.TerminalErrorCode == "canceled"
|
||||
}
|
||||
stream.mu.Lock()
|
||||
pending := stream.PendingCheckpoint
|
||||
stream.mu.Unlock()
|
||||
if checkpoint || !canceledEnd || pending != nil {
|
||||
t.Fatalf("cancel events checkpoint=%v canceled_end=%v pending=%v", checkpoint, canceledEnd, pending != nil)
|
||||
if checkpointCount != 0 || !canceledEnd || pending != nil {
|
||||
t.Fatalf("cancel events checkpoints=%d canceled_end=%v pending=%v", checkpointCount, canceledEnd, pending != nil)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -234,7 +234,7 @@ func (service *Service) buildLegacyCompactionPlan(base *compactionPlan, conversa
|
||||
if conversation == nil || base == nil {
|
||||
return nil, nil
|
||||
}
|
||||
candidates := buildContextCompactionCandidates(checkpointProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
|
||||
candidates := buildContextCompactionCandidates(replayablePromptProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
|
||||
if len(candidates) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
@@ -260,7 +260,7 @@ func (service *Service) buildAutoCompactionPlanFromHistory(base *compactionPlan,
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
currentCandidate, hasCurrentCandidate := buildCurrentTurnCompactionCandidate(checkpointProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
|
||||
currentCandidate, hasCurrentCandidate := buildCurrentTurnCompactionCandidate(replayablePromptProjectionEntries(conversation.Entries), base.CurrentTurnSeq, base.CurrentRequestID)
|
||||
if !hasCurrentCandidate {
|
||||
return legacyPlan, nil
|
||||
}
|
||||
@@ -447,16 +447,8 @@ func (service *Service) handleCompactionEvent(stream *ActiveStream, payload *str
|
||||
if err := service.completeManualCompactionTurn(stream); err != nil {
|
||||
return service.failStream(stream, "unknown", err)
|
||||
}
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{
|
||||
Message: buildTurnEndedMessage(0, 0, 0, 0),
|
||||
}); err != nil {
|
||||
return service.failStream(stream, "unknown", err)
|
||||
}
|
||||
if err := service.broker.Complete(stream.RequestID, "", ""); err != nil {
|
||||
return service.failStream(stream, "unknown", err)
|
||||
}
|
||||
service.setTurnPhase(stream, TurnPhaseCompleted)
|
||||
return nil
|
||||
completion := manualCompactionTurnCompletion(stream)
|
||||
return service.publishCheckpointWithCompletion(stream.RequestID, stream.ConversationID, &completion)
|
||||
}
|
||||
return service.requestProviderAction(stream, providerActionResume)
|
||||
}
|
||||
@@ -500,12 +492,8 @@ func (service *Service) finishManualCompactionNoop(stream *ActiveStream) error {
|
||||
if err := service.completeManualCompactionTurn(stream); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := service.broker.Publish(stream.RequestID, StreamEvent{
|
||||
Message: buildTurnEndedMessage(0, 0, 0, 0),
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return service.broker.Complete(stream.RequestID, "", "")
|
||||
completion := manualCompactionTurnCompletion(stream)
|
||||
return service.publishCheckpointWithCompletion(stream.RequestID, stream.ConversationID, &completion)
|
||||
}
|
||||
|
||||
func (service *Service) completeManualCompactionTurn(stream *ActiveStream) error {
|
||||
@@ -530,10 +518,21 @@ func (service *Service) completeManualCompactionTurn(stream *ActiveStream) error
|
||||
if err := service.syncSummaryCarryForward(conversationID, requestID, modelCallID); err != nil {
|
||||
return err
|
||||
}
|
||||
service.setTurnPhase(stream, TurnPhaseCompleted)
|
||||
return nil
|
||||
}
|
||||
|
||||
func manualCompactionTurnCompletion(stream *ActiveStream) pendingTurnCompletion {
|
||||
if stream == nil {
|
||||
return pendingTurnCompletion{}
|
||||
}
|
||||
return pendingTurnCompletion{
|
||||
ConversationID: strings.TrimSpace(stream.ConversationID),
|
||||
RequestID: strings.TrimSpace(stream.RequestID),
|
||||
TurnSeq: stream.TurnSeq,
|
||||
ModelCallID: "turn:" + strings.TrimSpace(stream.RequestID),
|
||||
}
|
||||
}
|
||||
|
||||
func (service *Service) publishSummaryCompleted(stream *ActiveStream, hookMessage string) error {
|
||||
if service == nil || stream == nil {
|
||||
return nil
|
||||
@@ -568,6 +567,7 @@ func (service *Service) applyCompactionPlan(stream *ActiveStream, conversationID
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
originalEntryCount := len(candidateConversation.Entries)
|
||||
if err := applyCompactionToConversation(candidateConversation, plan, summaryText); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -582,9 +582,9 @@ func (service *Service) applyCompactionPlan(stream *ActiveStream, conversationID
|
||||
if validationErr := validateCompactionCandidateBudget(recompiled, plan); validationErr != nil {
|
||||
return validationErr
|
||||
}
|
||||
replacementEntries := append([]HistoryEntry(nil), candidateConversation.Entries...)
|
||||
compactionEntries := append([]HistoryEntry(nil), candidateConversation.Entries[originalEntryCount:]...)
|
||||
if service.store != nil {
|
||||
persisted, err := service.store.ReplaceEntries(conversationID, replacementEntries, func(item *ConversationFile) error {
|
||||
persisted, _, err := service.store.AppendEntriesWithUpdate(conversationID, resetEntrySequences(compactionEntries), func(item *ConversationFile) error {
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -605,10 +605,7 @@ func (service *Service) applyCompactionPlan(stream *ActiveStream, conversationID
|
||||
if item == nil {
|
||||
return nil
|
||||
}
|
||||
item.Entries = nil
|
||||
item.NextEntrySeq = 1
|
||||
item.NextTurnSeq = 1
|
||||
appendEntriesInPlace(item, resetEntrySequences(replacementEntries))
|
||||
appendEntriesInPlace(item, resetEntrySequences(compactionEntries))
|
||||
item.TokenDetailsUsedTokens = 0
|
||||
clearConversationAutoCompactionState(item)
|
||||
return nil
|
||||
@@ -643,14 +640,13 @@ func applyCompactionToConversation(conversation *ConversationFile, plan *Pending
|
||||
if conversation == nil || plan == nil {
|
||||
return nil
|
||||
}
|
||||
replacementEntries, err := buildCompactedContextEntries(conversation, plan, summaryText)
|
||||
compactionEntries, err := buildCompactedContextEntries(conversation, plan, summaryText)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
conversation.Entries = nil
|
||||
conversation.NextEntrySeq = 1
|
||||
conversation.NextTurnSeq = 1
|
||||
appendEntriesInPlace(conversation, resetEntrySequences(replacementEntries))
|
||||
// Canonical history stays append-only. The prompt projector applies the
|
||||
// latest summary marker when constructing model-visible replay.
|
||||
appendEntriesInPlace(conversation, resetEntrySequences(compactionEntries))
|
||||
conversation.TokenDetailsUsedTokens = 0
|
||||
clearConversationAutoCompactionState(conversation)
|
||||
if conversation.TokenDetailsMaxTokens == 0 {
|
||||
@@ -671,40 +667,9 @@ func buildCompactedContextEntries(conversation *ConversationFile, plan *PendingC
|
||||
if ok {
|
||||
entries = append(entries, runtimeEntry)
|
||||
}
|
||||
if conversation == nil || !plan.PreserveCurrentTurnInputs {
|
||||
return entries, nil
|
||||
}
|
||||
entries = append(entries, buildAutoCompactionPreservedCurrentTurnEntries(conversation.Entries, plan)...)
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func buildAutoCompactionPreservedCurrentTurnEntries(entries []HistoryEntry, plan *PendingCompaction) []HistoryEntry {
|
||||
if len(entries) == 0 || plan == nil || !plan.PreserveCurrentTurnInputs {
|
||||
return nil
|
||||
}
|
||||
latestToolCallID := latestCompletedToolCallIDForTurn(entries, plan.CurrentTurnSeq, plan.CurrentRequestID)
|
||||
preservedIndexes := autoCompactionPreservedEntryIndexes(entries, plan.CurrentTurnSeq, plan.CurrentRequestID, latestToolCallID)
|
||||
if len(preservedIndexes) == 0 {
|
||||
return nil
|
||||
}
|
||||
preserved := make([]HistoryEntry, 0, len(preservedIndexes))
|
||||
for index, entry := range entries {
|
||||
if _, ok := preservedIndexes[index]; !ok {
|
||||
continue
|
||||
}
|
||||
switch strings.TrimSpace(entry.Kind) {
|
||||
case "compaction_summary", "compacted_summary", "compaction_request":
|
||||
continue
|
||||
case "tool_result":
|
||||
if rewritten, ok := rewriteAutoCompactionToolResultEntry(entry, autoCompactionPreservedToolResultLimitBytes, false); ok {
|
||||
entry = rewritten
|
||||
}
|
||||
}
|
||||
preserved = append(preserved, entry)
|
||||
}
|
||||
return preserved
|
||||
}
|
||||
|
||||
func newCompactionSummaryEntry(plan *PendingCompaction, summaryText string) HistoryEntry {
|
||||
payload, _ := json.Marshal(compactionSummaryEntryPayload{
|
||||
Summary: strings.TrimSpace(summaryText),
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protojson"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
func TestApplyCompactionToConversationPreservesCanonicalHistory(t *testing.T) {
|
||||
conversation := compactionAppendOnlyConversation(t)
|
||||
originalEntries := append([]HistoryEntry(nil), conversation.Entries...)
|
||||
plan := &PendingCompaction{
|
||||
Trigger: "manual",
|
||||
CurrentTurnSeq: 2,
|
||||
CurrentRequestID: "request-2",
|
||||
}
|
||||
|
||||
if err := applyCompactionToConversation(conversation, plan, "earlier context summary"); err != nil {
|
||||
t.Fatalf("applyCompactionToConversation() error = %v", err)
|
||||
}
|
||||
if len(conversation.Entries) <= len(originalEntries) {
|
||||
t.Fatalf("entries after compaction = %d, want the %d original entries plus a summary marker", len(conversation.Entries), len(originalEntries))
|
||||
}
|
||||
if !reflect.DeepEqual(conversation.Entries[:len(originalEntries)], originalEntries) {
|
||||
t.Fatal("compaction changed the canonical history prefix")
|
||||
}
|
||||
|
||||
projector := NewHistoryProjector()
|
||||
projection, err := projector.ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if len(projection.State.GetTurns()) != 2 {
|
||||
t.Fatalf("checkpoint turns after compaction = %d, want 2 visible turns", len(projection.State.GetTurns()))
|
||||
}
|
||||
replay, err := projector.ProjectPromptReplay(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectPromptReplay() error = %v", err)
|
||||
}
|
||||
if len(replay) != 1 || replay[0].Role != "user" || !strings.Contains(replay[0].Content, "earlier context summary") {
|
||||
t.Fatalf("prompt replay after compaction = %#v, want only the compacted summary", replay)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompactedPromptProjectionPlacesSummaryBeforePreservedCurrentTurn(t *testing.T) {
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
RootConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 1,
|
||||
NextEntrySeq: 1,
|
||||
}
|
||||
appendEntriesInPlace(conversation, []HistoryEntry{
|
||||
compactionTestUserEntry(t, 1, "request-1", "current question", "message-1"),
|
||||
newToolCallEntry(1, "request-1", "call-1", "Read", "", "", checkpointTestReadToolCall(t, nil)),
|
||||
newToolResultEntry(1, "request-1", "call-1", "Read", `{"path":"/tmp/example.txt"}`, "file contents", "", checkpointTestReadToolCall(t, nil)),
|
||||
})
|
||||
plan := &PendingCompaction{
|
||||
Trigger: "auto",
|
||||
CurrentTurnSeq: 1,
|
||||
CurrentRequestID: "request-1",
|
||||
PreserveCurrentTurnInputs: true,
|
||||
}
|
||||
if err := applyCompactionToConversation(conversation, plan, "current progress summary"); err != nil {
|
||||
t.Fatalf("applyCompactionToConversation() error = %v", err)
|
||||
}
|
||||
|
||||
projected := compactedPromptProjectionEntries(conversation.Entries)
|
||||
promptKinds := make([]string, 0, len(projected))
|
||||
for _, entry := range projected {
|
||||
if isPromptReplayEntryKind(entry.Kind) {
|
||||
promptKinds = append(promptKinds, entry.Kind)
|
||||
}
|
||||
}
|
||||
want := []string{"compacted_summary", "user_message", "tool_call", "tool_result"}
|
||||
if !reflect.DeepEqual(promptKinds, want) {
|
||||
t.Fatalf("compacted prompt entry order = %#v, want %#v", promptKinds, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompactionPlanningDoesNotRecompactArchivedHistory(t *testing.T) {
|
||||
conversation := compactionAppendOnlyConversation(t)
|
||||
if err := applyCompactionToConversation(conversation, &PendingCompaction{
|
||||
Trigger: "manual",
|
||||
CurrentTurnSeq: 2,
|
||||
CurrentRequestID: "request-2",
|
||||
}, "archived history summary"); err != nil {
|
||||
t.Fatalf("applyCompactionToConversation() error = %v", err)
|
||||
}
|
||||
appendEntriesInPlace(conversation, []HistoryEntry{
|
||||
compactionTestUserEntry(t, 3, "request-3", "new question", "message-3"),
|
||||
})
|
||||
|
||||
plan, err := (&Service{}).buildLegacyCompactionPlan(&compactionPlan{
|
||||
CurrentTurnSeq: 3,
|
||||
CurrentRequestID: "request-3",
|
||||
}, conversation, false, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("buildLegacyCompactionPlan() error = %v", err)
|
||||
}
|
||||
if plan != nil {
|
||||
t.Fatalf("buildLegacyCompactionPlan() = %#v, want no already summarized candidates", plan)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyCompactionPlanPersistsHistoryAppendOnly(t *testing.T) {
|
||||
store := NewConversationFileStore(t.TempDir())
|
||||
conversation := compactionAppendOnlyConversation(t)
|
||||
if _, _, err := store.AppendEntries(conversation.ConversationID, resetEntrySequences(conversation.Entries)); err != nil {
|
||||
t.Fatalf("AppendEntries() error = %v", err)
|
||||
}
|
||||
persisted, err := store.LoadConversation(conversation.ConversationID)
|
||||
if err != nil {
|
||||
t.Fatalf("initial LoadConversation() error = %v", err)
|
||||
}
|
||||
originalEntries := append([]HistoryEntry(nil), persisted.Entries...)
|
||||
projector := NewHistoryProjector()
|
||||
service := &Service{
|
||||
store: store,
|
||||
projector: projector,
|
||||
compiler: compactionProjectionCompiler{projector: projector},
|
||||
}
|
||||
stream := &ActiveStream{
|
||||
RequestID: "request-2",
|
||||
ConversationID: conversation.ConversationID,
|
||||
TurnSeq: 2,
|
||||
Mode: agentv1.AgentMode_AGENT_MODE_AGENT,
|
||||
CheckpointConversation: persisted,
|
||||
}
|
||||
plan := &PendingCompaction{
|
||||
Trigger: "manual",
|
||||
CurrentTurnSeq: 2,
|
||||
CurrentRequestID: "request-2",
|
||||
ContextWindowSize: 1_000_000,
|
||||
}
|
||||
if err := service.applyCompactionPlan(stream, conversation.ConversationID, plan, "persisted summary"); err != nil {
|
||||
t.Fatalf("applyCompactionPlan() error = %v", err)
|
||||
}
|
||||
|
||||
loaded, err := store.LoadConversation(conversation.ConversationID)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConversation() error = %v", err)
|
||||
}
|
||||
if len(loaded.Entries) <= len(originalEntries) {
|
||||
t.Fatalf("persisted entries after compaction = %d, want more than %d", len(loaded.Entries), len(originalEntries))
|
||||
}
|
||||
for index := range originalEntries {
|
||||
if !reflect.DeepEqual(loaded.Entries[index], originalEntries[index]) {
|
||||
t.Fatalf("persisted history entry %d changed after compaction:\ngot %#v\nwant %#v", index, loaded.Entries[index], originalEntries[index])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type compactionProjectionCompiler struct {
|
||||
projector *HistoryProjector
|
||||
}
|
||||
|
||||
func (compiler compactionProjectionCompiler) Compile(conversation *ConversationFile, _ agentv1.AgentMode, _ string, _ string) (CompiledConversation, error) {
|
||||
messages, err := compiler.projector.ProjectPromptReplay(conversation)
|
||||
return CompiledConversation{Messages: messages}, err
|
||||
}
|
||||
|
||||
func (compactionProjectionCompiler) DerivePromptContexts(*ConversationFile, agentv1.AgentMode, string) ([]PromptContextMessage, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func compactionAppendOnlyConversation(t *testing.T) *ConversationFile {
|
||||
t.Helper()
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
RootConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 1,
|
||||
NextEntrySeq: 1,
|
||||
TokenDetailsUsedTokens: 42_000,
|
||||
TokenDetailsMaxTokens: 50_000,
|
||||
}
|
||||
appendEntriesInPlace(conversation, []HistoryEntry{
|
||||
compactionTestUserEntry(t, 1, "request-1", "first question", "message-1"),
|
||||
newAssistantTextEntry(1, "request-1", "first answer", "", ""),
|
||||
compactionTestUserEntry(t, 2, "request-2", "second question", "message-2"),
|
||||
newAssistantTextEntry(2, "request-2", "second answer", "", ""),
|
||||
})
|
||||
return conversation
|
||||
}
|
||||
|
||||
func compactionTestUserEntry(t *testing.T, turnSeq int64, requestID string, text string, messageID string) HistoryEntry {
|
||||
t.Helper()
|
||||
payload, err := protojson.Marshal(&agentv1.UserMessage{Text: text, MessageId: messageID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal user message: %v", err)
|
||||
}
|
||||
return HistoryEntry{
|
||||
TurnSeq: turnSeq,
|
||||
RequestID: requestID,
|
||||
Role: "user",
|
||||
Kind: "user_message",
|
||||
Payload: payload,
|
||||
}
|
||||
}
|
||||
|
||||
var _ PromptCompiler = compactionProjectionCompiler{}
|
||||
@@ -201,6 +201,41 @@ func buildShellOutputDeltaMessage(delta *agentv1.ShellOutputDeltaUpdate) *agentv
|
||||
}
|
||||
}
|
||||
|
||||
// buildShellToolCallDeltaMessage maps client shell output to the delta consumed by Cursor's terminal bubble.
|
||||
func buildShellToolCallDeltaMessage(callID string, modelCallID string, output *agentv1.ShellOutputDeltaUpdate) *agentv1.AgentServerMessage {
|
||||
if output == nil {
|
||||
return nil
|
||||
}
|
||||
var delta *agentv1.ShellToolCallDelta
|
||||
switch event := output.GetEvent().(type) {
|
||||
case *agentv1.ShellOutputDeltaUpdate_Stdout:
|
||||
content := event.Stdout.GetData()
|
||||
if content == "" {
|
||||
return nil
|
||||
}
|
||||
delta = &agentv1.ShellToolCallDelta{
|
||||
Delta: &agentv1.ShellToolCallDelta_Stdout{
|
||||
Stdout: &agentv1.ShellToolCallStdoutDelta{Content: content},
|
||||
},
|
||||
}
|
||||
case *agentv1.ShellOutputDeltaUpdate_Stderr:
|
||||
content := event.Stderr.GetData()
|
||||
if content == "" {
|
||||
return nil
|
||||
}
|
||||
delta = &agentv1.ShellToolCallDelta{
|
||||
Delta: &agentv1.ShellToolCallDelta_Stderr{
|
||||
Stderr: &agentv1.ShellToolCallStderrDelta{Content: content},
|
||||
},
|
||||
}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
return buildToolCallDeltaMessage(callID, modelCallID, &agentv1.ToolCallDelta{
|
||||
Delta: &agentv1.ToolCallDelta_ShellToolCallDelta{ShellToolCallDelta: delta},
|
||||
})
|
||||
}
|
||||
|
||||
// buildTurnEndedMessage 构造 turn 结束消息,并携带标准化后的 token 统计。
|
||||
func buildTurnEndedMessage(inputTokens int64, outputTokens int64, cacheReadTokens int64, cacheWriteTokens int64) *agentv1.AgentServerMessage {
|
||||
inputTokensValue := inputTokens
|
||||
|
||||
@@ -121,10 +121,15 @@ func (store *ConversationFileStore) LoadConversation(conversationID string) (*Co
|
||||
|
||||
// AppendEntries 把已经发生的语义事件追加到 context.json,并同步 state.json。
|
||||
func (store *ConversationFileStore) AppendEntries(conversationID string, entries []HistoryEntry) (*ConversationFile, []HistoryEntry, error) {
|
||||
return store.AppendEntriesWithUpdate(conversationID, entries, nil)
|
||||
}
|
||||
|
||||
// AppendEntriesWithUpdate 原子追加 context entries,并在同一把会话锁内更新 state metadata。
|
||||
func (store *ConversationFileStore) AppendEntriesWithUpdate(conversationID string, entries []HistoryEntry, update func(*ConversationFile) error) (*ConversationFile, []HistoryEntry, error) {
|
||||
if store == nil {
|
||||
return nil, nil, fmt.Errorf("conversation file store is nil")
|
||||
}
|
||||
if len(entries) == 0 {
|
||||
if len(entries) == 0 && update == nil {
|
||||
conversation, err := store.LoadConversation(conversationID)
|
||||
return conversation, nil, err
|
||||
}
|
||||
@@ -162,6 +167,11 @@ func (store *ConversationFileStore) AppendEntries(conversationID string, entries
|
||||
conversation.Mode = alias
|
||||
}
|
||||
assigned := appendEntriesInPlace(conversation, entries)
|
||||
if update != nil {
|
||||
if err := update(conversation); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
}
|
||||
deriveConversationLoopState(conversation)
|
||||
if err := store.writeConversationLocked(normalizedConversationID, conversation); err != nil {
|
||||
return nil, nil, err
|
||||
@@ -690,9 +700,21 @@ func appendEntriesInPlace(conversation *ConversationFile, entries []HistoryEntry
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
assigned := make([]HistoryEntry, 0, len(entries))
|
||||
existingIdempotencyKeys := make(map[string]struct{})
|
||||
for _, existing := range conversation.Entries {
|
||||
if key := strings.TrimSpace(existing.IdempotencyKey); key != "" {
|
||||
existingIdempotencyKeys[key] = struct{}{}
|
||||
}
|
||||
}
|
||||
maxTurnSeq := conversation.NextTurnSeq - 1
|
||||
for _, entry := range entries {
|
||||
next := entry
|
||||
if key := strings.TrimSpace(next.IdempotencyKey); key != "" {
|
||||
if _, exists := existingIdempotencyKeys[key]; exists {
|
||||
continue
|
||||
}
|
||||
existingIdempotencyKeys[key] = struct{}{}
|
||||
}
|
||||
if next.CreatedAt.IsZero() {
|
||||
next.CreatedAt = now
|
||||
}
|
||||
@@ -750,6 +772,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 +905,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
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
)
|
||||
|
||||
func TestImportedConversationStateRestoresBlobOnlyForkAndCheckpointPrefix(t *testing.T) {
|
||||
parent := compactionAppendOnlyConversation(t)
|
||||
parent.Entries = parent.Entries[:2]
|
||||
parent.NextEntrySeq = 3
|
||||
parent.NextTurnSeq = 2
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(parent)
|
||||
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, err := newRuntimeConversation("fork-conversation", agentv1.AgentMode_AGENT_MODE_AGENT)
|
||||
if err != nil {
|
||||
t.Fatalf("newRuntimeConversation() error = %v", err)
|
||||
}
|
||||
entries, err := (&Service{}).importConversationState(conversation, state, prefetched)
|
||||
if err != nil {
|
||||
t.Fatalf("importConversationState() error = %v", err)
|
||||
}
|
||||
if len(conversation.ImportedTurnIDs) != 1 || conversation.NextTurnSeq != 2 {
|
||||
t.Fatalf("imported prefix turns=%d next_turn_seq=%d, want 1 and 2", len(conversation.ImportedTurnIDs), conversation.NextTurnSeq)
|
||||
}
|
||||
if len(entries) != 2 {
|
||||
t.Fatalf("imported model entries = %d, want parent user and assistant", len(entries))
|
||||
}
|
||||
appendEntriesInPlace(conversation, append(entries,
|
||||
compactionTestUserEntry(t, 2, "request-2", "fork question", "message-2"),
|
||||
))
|
||||
forkProjection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("fork ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if len(forkProjection.State.GetTurns()) != 2 {
|
||||
t.Fatalf("fork checkpoint turns = %d, want imported parent plus local fork turn", len(forkProjection.State.GetTurns()))
|
||||
}
|
||||
if string(forkProjection.State.GetTurns()[0]) != string(projection.State.GetTurns()[0]) {
|
||||
t.Fatal("fork checkpoint did not preserve the imported parent turn ID as its prefix")
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportedConversationStateRejectsUnresolvedBlobTurn(t *testing.T) {
|
||||
turnID := sha256.Sum256([]byte("missing imported turn"))
|
||||
conversation, err := newRuntimeConversation("fork-conversation", agentv1.AgentMode_AGENT_MODE_AGENT)
|
||||
if err != nil {
|
||||
t.Fatalf("newRuntimeConversation() error = %v", err)
|
||||
}
|
||||
if _, err := (&Service{}).importConversationState(conversation, &agentv1.ConversationStateStructure{
|
||||
Turns: [][]byte{turnID[:]},
|
||||
}, nil); err == nil {
|
||||
t.Fatal("importConversationState() accepted an unresolved Blob turn")
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportedTurnIDsPersistThroughConversationStore(t *testing.T) {
|
||||
store := NewConversationFileStore(t.TempDir())
|
||||
turnID := sha256.Sum256([]byte("parent turn"))
|
||||
conversation, err := newRuntimeConversation("fork-conversation", agentv1.AgentMode_AGENT_MODE_AGENT)
|
||||
if err != nil {
|
||||
t.Fatalf("newRuntimeConversation() error = %v", err)
|
||||
}
|
||||
conversation.ImportedTurnIDs = [][]byte{turnID[:]}
|
||||
persisted, err := store.SaveConversationWithEntries(conversation.ConversationID, conversation, []HistoryEntry{
|
||||
compactionTestUserEntry(t, 2, "request-2", "fork question", "message-2"),
|
||||
})
|
||||
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, 3)
|
||||
for index := range ids {
|
||||
digest := sha256.Sum256([]byte{byte(index + 1)})
|
||||
ids[index] = digest[:]
|
||||
}
|
||||
trimmed := rewindImportedTurnPrefix(ids, runRewindDecision{
|
||||
TargetTurnSeq: 4,
|
||||
HasClientTurnCount: true,
|
||||
ClientTurnCount: 1,
|
||||
})
|
||||
if len(trimmed) != 1 || string(trimmed[0]) != string(ids[0]) {
|
||||
t.Fatalf("rewindImportedTurnPrefix() = %x, want first imported turn only", trimmed)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,235 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAppendEntriesDeduplicatesIdempotencyKey(t *testing.T) {
|
||||
store := NewConversationFileStore(t.TempDir())
|
||||
entry := HistoryEntry{
|
||||
TurnSeq: 1,
|
||||
RequestID: "request-1",
|
||||
IdempotencyKey: "provider-interrupted-output:test",
|
||||
Role: "assistant",
|
||||
Kind: "assistant_text",
|
||||
Payload: json.RawMessage(`{"text":"partial"}`),
|
||||
}
|
||||
|
||||
if _, assigned, err := store.AppendEntries("conversation-1", []HistoryEntry{entry}); err != nil {
|
||||
t.Fatalf("first AppendEntries() error = %v", err)
|
||||
} else if len(assigned) != 1 {
|
||||
t.Fatalf("first AppendEntries() assigned = %d, want 1", len(assigned))
|
||||
}
|
||||
if _, assigned, err := store.AppendEntries("conversation-1", []HistoryEntry{entry}); err != nil {
|
||||
t.Fatalf("duplicate AppendEntries() error = %v", err)
|
||||
} else if len(assigned) != 0 {
|
||||
t.Fatalf("duplicate AppendEntries() assigned = %d, want 0", len(assigned))
|
||||
}
|
||||
|
||||
conversation, err := store.LoadConversation("conversation-1")
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConversation() error = %v", err)
|
||||
}
|
||||
if len(conversation.Entries) != 1 {
|
||||
t.Fatalf("persisted entries = %d, want 1", len(conversation.Entries))
|
||||
}
|
||||
}
|
||||
|
||||
func TestCancelPersistsInterruptedProviderOutputIdempotently(t *testing.T) {
|
||||
service, stream, _ := testCheckpointBlobProjection(t)
|
||||
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
|
||||
if err != nil {
|
||||
t.Fatalf("snapshotCheckpointConversation() error = %v", err)
|
||||
}
|
||||
if _, err := service.store.SaveConversationWithEntries(stream.ConversationID, conversation, conversation.Entries); err != nil {
|
||||
t.Fatalf("SaveConversationWithEntries() error = %v", err)
|
||||
}
|
||||
|
||||
stream.mu.Lock()
|
||||
stream.CurrentModelCallID = "model-call-1"
|
||||
stream.ProviderAccumulatedText = "partial answer"
|
||||
stream.ProviderAccumulatedReasoning = "partial reasoning"
|
||||
stream.mu.Unlock()
|
||||
|
||||
cancel := InboundIntent{
|
||||
Kind: "cancel",
|
||||
RequestID: stream.RequestID,
|
||||
CancelReason: "[canceled] Superseded by newer request",
|
||||
}
|
||||
if err := service.handleCancelIntent(cancel); err != nil {
|
||||
t.Fatalf("first handleCancelIntent() error = %v", err)
|
||||
}
|
||||
stream.mu.Lock()
|
||||
stream.ProviderAccumulatedText = "late duplicate fragment"
|
||||
stream.mu.Unlock()
|
||||
if err := service.handleCancelIntent(cancel); err != nil {
|
||||
t.Fatalf("duplicate handleCancelIntent() error = %v", err)
|
||||
}
|
||||
|
||||
persisted, err := service.store.LoadConversation(stream.ConversationID)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConversation() error = %v", err)
|
||||
}
|
||||
assistantEntries := 0
|
||||
cancelEntries := 0
|
||||
for _, entry := range persisted.Entries {
|
||||
if entry.Kind == "metadata" {
|
||||
var payload metadataPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
t.Fatalf("decode metadata entry: %v", err)
|
||||
}
|
||||
if payload.Type == "control" && readStringValue(payload.Value["status"]) == "canceled" {
|
||||
cancelEntries++
|
||||
}
|
||||
}
|
||||
if entry.Kind != "assistant_text" {
|
||||
continue
|
||||
}
|
||||
var payload assistantTextPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
t.Fatalf("decode assistant entry: %v", err)
|
||||
}
|
||||
if payload.Text == "partial answer" {
|
||||
assistantEntries++
|
||||
}
|
||||
}
|
||||
if assistantEntries != 1 {
|
||||
t.Fatalf("persisted interrupted assistant entries = %d, want 1", assistantEntries)
|
||||
}
|
||||
if cancelEntries != 1 {
|
||||
t.Fatalf("persisted cancel metadata entries = %d, want 1", cancelEntries)
|
||||
}
|
||||
|
||||
replay, err := service.projector.ProjectPromptReplay(persisted)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectPromptReplay() error = %v", err)
|
||||
}
|
||||
found := false
|
||||
for _, message := range replay {
|
||||
if message.Role == "assistant" && strings.TrimSpace(message.Content) == "partial answer" && strings.TrimSpace(message.ReasoningContent) == "partial reasoning" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("replay = %#v, want interrupted assistant output", replay)
|
||||
}
|
||||
checkpoint, err := service.projector.ProjectCheckpointProjection(persisted)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if checkpoint == nil || checkpoint.State == nil || len(checkpoint.State.GetTurns()) != 1 {
|
||||
t.Fatalf("checkpoint state = %#v, want interrupted turn", checkpoint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenericProviderFailurePersistsAccumulatedOutput(t *testing.T) {
|
||||
service, stream, _ := testCheckpointBlobProjection(t)
|
||||
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
|
||||
if err != nil {
|
||||
t.Fatalf("snapshotCheckpointConversation() error = %v", err)
|
||||
}
|
||||
if _, err := service.store.SaveConversationWithEntries(stream.ConversationID, conversation, conversation.Entries); err != nil {
|
||||
t.Fatalf("SaveConversationWithEntries() error = %v", err)
|
||||
}
|
||||
|
||||
stream.mu.Lock()
|
||||
stream.CurrentModelCallID = "model-call-1"
|
||||
stream.ProviderActive = true
|
||||
stream.ProviderAccumulatedText = "partial answer before transport failure"
|
||||
stream.Status = StreamStatusStreaming
|
||||
stream.Phase = TurnPhaseProviderRunning
|
||||
stream.mu.Unlock()
|
||||
if err := service.handleProviderDoneEvent(stream, &streamProviderEvent{
|
||||
Done: true,
|
||||
Err: errors.New("transport failed"),
|
||||
}); err != nil {
|
||||
t.Fatalf("handleProviderDoneEvent() error = %v", err)
|
||||
}
|
||||
|
||||
persisted, err := service.store.LoadConversation(stream.ConversationID)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConversation() error = %v", err)
|
||||
}
|
||||
foundPartialOutput := false
|
||||
for _, entry := range persisted.Entries {
|
||||
if entry.Kind != "assistant_text" {
|
||||
continue
|
||||
}
|
||||
var payload assistantTextPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
t.Fatalf("decode assistant entry: %v", err)
|
||||
}
|
||||
if payload.Text == "partial answer before transport failure" {
|
||||
foundPartialOutput = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !foundPartialOutput {
|
||||
t.Fatal("generic provider failure discarded accumulated assistant output")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCancelPreservesPersistedTurnActivityWithoutLiveAccumulator(t *testing.T) {
|
||||
service, stream, _ := testCheckpointBlobProjection(t)
|
||||
conversation, _, _, err := service.snapshotCheckpointConversation(stream)
|
||||
if err != nil {
|
||||
t.Fatalf("snapshotCheckpointConversation() error = %v", err)
|
||||
}
|
||||
if _, err := service.store.SaveConversationWithEntries(stream.ConversationID, conversation, conversation.Entries); err != nil {
|
||||
t.Fatalf("SaveConversationWithEntries() error = %v", err)
|
||||
}
|
||||
|
||||
if err := service.handleCancelIntent(InboundIntent{
|
||||
Kind: "cancel",
|
||||
RequestID: stream.RequestID,
|
||||
CancelReason: "new_message_submitted",
|
||||
}); err != nil {
|
||||
t.Fatalf("handleCancelIntent() error = %v", err)
|
||||
}
|
||||
|
||||
persisted, err := service.store.LoadConversation(stream.ConversationID)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadConversation() error = %v", err)
|
||||
}
|
||||
replay, err := service.projector.ProjectPromptReplay(persisted)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectPromptReplay() error = %v", err)
|
||||
}
|
||||
for _, message := range replay {
|
||||
if message.Role == "assistant" && strings.TrimSpace(message.Content) == "hi" {
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatalf("replay = %#v, want persisted assistant activity", replay)
|
||||
}
|
||||
|
||||
func TestProjectPromptReplayPreservesLegacyCanceledTurnActivity(t *testing.T) {
|
||||
cancelEntry := newMetadataEntry(1, "request-1", "control", map[string]any{
|
||||
"status": "canceled",
|
||||
"reason": "new_message_submitted",
|
||||
"replay_policy": cancelReplayPolicyKeepStableInput,
|
||||
})
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
NextTurnSeq: 2,
|
||||
Entries: []HistoryEntry{
|
||||
newAssistantTextEntry(1, "request-1", "persisted activity", "", ""),
|
||||
cancelEntry,
|
||||
},
|
||||
}
|
||||
|
||||
replay, err := NewHistoryProjector().ProjectPromptReplay(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectPromptReplay() error = %v", err)
|
||||
}
|
||||
for _, message := range replay {
|
||||
if message.Role == "assistant" && strings.TrimSpace(message.Content) == "persisted activity" {
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatalf("replay = %#v, want legacy canceled activity", replay)
|
||||
}
|
||||
@@ -249,12 +249,12 @@ func (projector *HistoryProjector) ProjectPromptReplay(conversation *Conversatio
|
||||
replayMessages[index].OpenAIResponsesReasoningSummary = append(json.RawMessage(nil), payload.ReasoningSummary...)
|
||||
applyPromptProviderMetadataToFirstToolCall(&replayMessages[index], payload.ProviderItemID, payload.ProviderCallID, payload.ProviderStatus)
|
||||
}
|
||||
for index := range replayMessages {
|
||||
if strings.TrimSpace(replayMessages[index].Role) == "tool" {
|
||||
toolName := firstNonEmpty(strings.TrimSpace(replayMessages[index].Name), strings.TrimSpace(payload.ToolName))
|
||||
replayMessages[index].Content = limitProjectedToolResultReplay(toolName, replayMessages[index].Content, payload.ResultText, true, historicalToolResult)
|
||||
for _, replay := range replayMessages {
|
||||
if strings.TrimSpace(replay.Role) == "tool" {
|
||||
toolName := firstNonEmpty(strings.TrimSpace(replay.Name), strings.TrimSpace(payload.ToolName))
|
||||
replay.Content = limitProjectedToolResultReplay(toolName, replay.Content, payload.ResultText, true, historicalToolResult)
|
||||
}
|
||||
messages = append(messages, toModelMessage(replayMessages[index]))
|
||||
messages = append(messages, toModelMessage(replay))
|
||||
}
|
||||
continue
|
||||
}
|
||||
@@ -326,20 +326,24 @@ func compactedPromptProjectionEntries(entries []HistoryEntry) []HistoryEntry {
|
||||
latestToolCallID := latestCompletedToolCallIDForTurn(entries, compactionPayload.CurrentTurnSeq, compactionPayload.CurrentRequestID)
|
||||
preservedIndexes = autoCompactionPreservedEntryIndexes(entries, compactionPayload.CurrentTurnSeq, compactionPayload.CurrentRequestID, latestToolCallID)
|
||||
}
|
||||
filtered := make([]HistoryEntry, 0, len(entries)-compactionIndex)
|
||||
for index, entry := range entries {
|
||||
if index < compactionIndex && isPromptReplayEntryKind(entry.Kind) {
|
||||
if _, ok := preservedIndexes[index]; !ok {
|
||||
continue
|
||||
}
|
||||
filtered := make([]HistoryEntry, 0, len(entries)-compactionIndex+len(preservedIndexes))
|
||||
for index := 0; index < compactionIndex; index++ {
|
||||
if !isPromptReplayEntryKind(entries[index].Kind) {
|
||||
filtered = append(filtered, entries[index])
|
||||
}
|
||||
if index < compactionIndex {
|
||||
if rewritten, ok := compactedProjectionPreservedEntry(entry); ok {
|
||||
entry = rewritten
|
||||
}
|
||||
}
|
||||
filtered = append(filtered, entries[compactionIndex])
|
||||
for index := 0; index < compactionIndex; index++ {
|
||||
if _, ok := preservedIndexes[index]; !ok || isCompactionSummaryKind(entries[index].Kind) {
|
||||
continue
|
||||
}
|
||||
entry := entries[index]
|
||||
if rewritten, ok := compactedProjectionPreservedEntry(entry); ok {
|
||||
entry = rewritten
|
||||
}
|
||||
filtered = append(filtered, entry)
|
||||
}
|
||||
filtered = append(filtered, entries[compactionIndex+1:]...)
|
||||
return filtered
|
||||
}
|
||||
|
||||
@@ -355,6 +359,7 @@ const (
|
||||
cancelReplayPolicyDropTurn = "drop_turn"
|
||||
cancelReplayPolicyDropUnstarted = "drop_unstarted_turn"
|
||||
cancelReplayPolicyKeepStableInput = "keep_stable_input"
|
||||
cancelReplayPolicyKeepInterrupted = "keep_interrupted_output"
|
||||
)
|
||||
|
||||
func sanitizeCanceledReplayEntries(entries []HistoryEntry) []HistoryEntry {
|
||||
@@ -370,13 +375,19 @@ func sanitizeCanceledReplayEntries(entries []HistoryEntry) []HistoryEntry {
|
||||
for _, entry := range entries {
|
||||
if entry.TurnSeq > 0 {
|
||||
if policy, canceled := canceledTurns[entry.TurnSeq]; canceled {
|
||||
if policy == cancelReplayPolicyDropUnstarted {
|
||||
if policy == cancelReplayPolicyKeepInterrupted {
|
||||
filtered = append(filtered, entry)
|
||||
continue
|
||||
}
|
||||
if policy != cancelReplayPolicyDropTurn {
|
||||
if _, active := activeCanceledTurns[entry.TurnSeq]; active {
|
||||
policy = cancelReplayPolicyKeepStableInput
|
||||
} else {
|
||||
policy = cancelReplayPolicyDropTurn
|
||||
filtered = append(filtered, entry)
|
||||
continue
|
||||
}
|
||||
}
|
||||
if policy == cancelReplayPolicyDropUnstarted {
|
||||
policy = cancelReplayPolicyDropTurn
|
||||
}
|
||||
if policy == cancelReplayPolicyDropTurn || !isStableCanceledTurnInputEntry(entry) {
|
||||
continue
|
||||
}
|
||||
@@ -450,6 +461,8 @@ func normalizeCancelReplayPolicy(policy string, reason string) string {
|
||||
return cancelReplayPolicyDropUnstarted
|
||||
case cancelReplayPolicyKeepStableInput:
|
||||
return cancelReplayPolicyKeepStableInput
|
||||
case cancelReplayPolicyKeepInterrupted:
|
||||
return cancelReplayPolicyKeepInterrupted
|
||||
default:
|
||||
return cancelReplayPolicyForReason(reason)
|
||||
}
|
||||
@@ -566,7 +579,7 @@ func (projector *HistoryProjector) ProjectCheckpointProjection(conversation *Con
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
state.Turns = turnIDs
|
||||
state.Turns = append(cloneByteSlices(conversation.ImportedTurnIDs), turnIDs...)
|
||||
replayMessages, err := projector.ProjectPromptReplay(conversation)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -612,13 +625,29 @@ func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpoin
|
||||
}
|
||||
grouped[entry.TurnSeq] = append(grouped[entry.TurnSeq], entry)
|
||||
}
|
||||
|
||||
turnIDs := make([][]byte, 0, len(order))
|
||||
logicalTurns := make([][]HistoryEntry, 0, len(order))
|
||||
for _, turnSeq := range order {
|
||||
entries := grouped[turnSeq]
|
||||
if checkpointTurnHasUserMessage(entries) {
|
||||
logicalTurns = append(logicalTurns, append([]HistoryEntry(nil), entries...))
|
||||
continue
|
||||
}
|
||||
if len(logicalTurns) == 0 {
|
||||
continue
|
||||
}
|
||||
last := len(logicalTurns) - 1
|
||||
logicalTurns[last] = append(logicalTurns[last], entries...)
|
||||
}
|
||||
|
||||
turnIDs := make([][]byte, 0, len(logicalTurns))
|
||||
for _, entries := range logicalTurns {
|
||||
completedToolCalls, err := collectCheckpointCompletedToolCalls(entries)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var userMessageID []byte
|
||||
var turnRequestID string
|
||||
stepIDs := make([][]byte, 0, len(entries))
|
||||
steps := make([]*agentv1.ConversationStep, 0, len(entries))
|
||||
seenToolCalls := make(map[string]struct{})
|
||||
openToolCalls := make(map[string]struct{})
|
||||
for _, entry := range entries {
|
||||
@@ -645,59 +674,48 @@ func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpoin
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(payload.ReasoningContent) != "" {
|
||||
stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{
|
||||
steps = append(steps, &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{
|
||||
steps = append(steps, &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{
|
||||
steps = append(steps, &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{}
|
||||
toolCallID := strings.TrimSpace(payload.ToolCallID)
|
||||
if err := protojson.Unmarshal(payload.ToolCall, toolCall); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !shouldPersistToolResultName(firstNonEmpty(strings.TrimSpace(payload.ToolName), inferToolName(toolCall))) {
|
||||
continue
|
||||
if completedPayload := completedToolCalls[toolCallID]; len(completedPayload) > 0 {
|
||||
completedToolCall := &agentv1.ToolCall{}
|
||||
if err := protojson.Unmarshal(completedPayload, completedToolCall); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
proto.Merge(toolCall, completedToolCall)
|
||||
}
|
||||
stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{
|
||||
steps = append(steps, &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 != "" {
|
||||
if toolCallID != "" {
|
||||
seenToolCalls[toolCallID] = struct{}{}
|
||||
openToolCalls[toolCallID] = struct{}{}
|
||||
}
|
||||
@@ -706,22 +724,19 @@ func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpoin
|
||||
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
|
||||
}
|
||||
toolCallID := strings.TrimSpace(payload.ToolCallID)
|
||||
if toolCallID != "" {
|
||||
delete(openToolCalls, toolCallID)
|
||||
}
|
||||
if _, ok := seenToolCalls[toolCallID]; ok {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(payload.ReasoningContent) != "" {
|
||||
stepID, err := addCheckpointStepBlob(blobs, &agentv1.ConversationStep{
|
||||
steps = append(steps, &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
|
||||
@@ -730,21 +745,22 @@ func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpoin
|
||||
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{
|
||||
steps = append(steps, &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 {
|
||||
if len(userMessageID) == 0 {
|
||||
continue
|
||||
}
|
||||
stepIDs := make([][]byte, 0, len(steps))
|
||||
for _, step := range steps {
|
||||
stepID, err := addCheckpointStepBlob(blobs, step)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stepIDs = append(stepIDs, stepID)
|
||||
}
|
||||
agentTurn := &agentv1.AgentConversationTurnStructure{
|
||||
UserMessage: userMessageID,
|
||||
Steps: stepIDs,
|
||||
@@ -765,6 +781,32 @@ func projectCheckpointTurnBlobs(conversation *ConversationFile, blobs *checkpoin
|
||||
return turnIDs, nil
|
||||
}
|
||||
|
||||
func checkpointTurnHasUserMessage(entries []HistoryEntry) bool {
|
||||
for _, entry := range entries {
|
||||
if strings.TrimSpace(entry.Kind) == "user_message" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func collectCheckpointCompletedToolCalls(entries []HistoryEntry) (map[string]json.RawMessage, error) {
|
||||
completed := make(map[string]json.RawMessage)
|
||||
for _, entry := range entries {
|
||||
if strings.TrimSpace(entry.Kind) != "tool_result" {
|
||||
continue
|
||||
}
|
||||
var payload toolResultEntryPayload
|
||||
if err := json.Unmarshal(entry.Payload, &payload); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if toolCallID := strings.TrimSpace(payload.ToolCallID); toolCallID != "" && len(payload.ToolCall) > 0 {
|
||||
completed[toolCallID] = payload.ToolCall
|
||||
}
|
||||
}
|
||||
return completed, nil
|
||||
}
|
||||
|
||||
func addCheckpointStepBlob(blobs *checkpointBlobGraph, step *agentv1.ConversationStep) ([]byte, error) {
|
||||
payload, err := proto.Marshal(step)
|
||||
if err != nil {
|
||||
@@ -1209,7 +1251,7 @@ func trimReplayDanglingAssistantToolCalls(messages []modeladapter.Message) []mod
|
||||
return trimmed
|
||||
}
|
||||
|
||||
func shouldPersistToolResultName(toolName string) bool {
|
||||
func shouldPersistCheckpointReplayToolResultName(toolName string) bool {
|
||||
switch strings.TrimSpace(toolName) {
|
||||
case "PatchEdit", "PatchEditLines", "PatchEditSpan", "Edit", "Write", "GenerateImage":
|
||||
return true
|
||||
@@ -1218,60 +1260,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
|
||||
@@ -1282,7 +1270,7 @@ func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []pro
|
||||
if strings.TrimSpace(message.Role) == "assistant" && len(message.ToolCalls) > 0 {
|
||||
nextToolCalls := make([]promptengine.ToolCallDescriptor, 0, len(message.ToolCalls))
|
||||
for _, toolCall := range message.ToolCalls {
|
||||
if !shouldPersistToolResultName(toolCall.Function.Name) {
|
||||
if !shouldPersistCheckpointReplayToolResultName(toolCall.Function.Name) {
|
||||
skippedToolCallIDs[strings.TrimSpace(toolCall.ID)] = struct{}{}
|
||||
continue
|
||||
}
|
||||
@@ -1299,7 +1287,7 @@ func filterCheckpointPersistentToolReplay(messages []promptengine.Message) []pro
|
||||
if _, ok := skippedToolCallIDs[strings.TrimSpace(message.ToolCallID)]; ok {
|
||||
continue
|
||||
}
|
||||
if !shouldPersistToolResultName(message.Name) {
|
||||
if !shouldPersistCheckpointReplayToolResultName(message.Name) {
|
||||
continue
|
||||
}
|
||||
}
|
||||
@@ -1308,7 +1296,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
|
||||
}
|
||||
@@ -1317,16 +1305,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)
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -9,6 +11,7 @@ import (
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
promptengine "cursor/internal/backend/agent/prompt"
|
||||
)
|
||||
|
||||
func TestProjectCheckpointProjectionBuildsResolvableForkState(t *testing.T) {
|
||||
@@ -84,6 +87,94 @@ func TestProjectCheckpointProjectionBuildsResolvableForkState(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectCheckpointProjectionMergesResumeActivityIntoPreviousUserTurn(t *testing.T) {
|
||||
userPayload, err := protojson.Marshal(&agentv1.UserMessage{
|
||||
Text: "original question",
|
||||
MessageId: "message-1",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal user message: %v", err)
|
||||
}
|
||||
firstAnswer := newAssistantTextEntry(1, "request-1", "before resume", "", "")
|
||||
firstAnswer.Seq = 2
|
||||
resumedAnswer := newAssistantTextEntry(2, "request-resume", "after resume", "", "")
|
||||
resumedAnswer.Seq = 3
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
RootConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 3,
|
||||
NextEntrySeq: 4,
|
||||
Entries: []HistoryEntry{
|
||||
{Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: userPayload},
|
||||
firstAnswer,
|
||||
resumedAnswer,
|
||||
},
|
||||
}
|
||||
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if len(projection.State.GetTurns()) != 1 {
|
||||
t.Fatalf("checkpoint turns = %d, want one logical user turn", len(projection.State.GetTurns()))
|
||||
}
|
||||
|
||||
blobs := make(map[string][]byte, len(projection.Blobs))
|
||||
for _, blob := range projection.Blobs {
|
||||
blobs[string(blob.ID)] = blob.Data
|
||||
}
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(blobs[string(projection.State.GetTurns()[0])], turn); err != nil {
|
||||
t.Fatalf("decode checkpoint turn: %v", err)
|
||||
}
|
||||
agentTurn := turn.GetAgentConversationTurn()
|
||||
if agentTurn == nil {
|
||||
t.Fatal("checkpoint turn does not contain an agent turn")
|
||||
}
|
||||
if len(agentTurn.GetUserMessage()) == 0 {
|
||||
t.Fatal("checkpoint turn lost the original user message Blob id")
|
||||
}
|
||||
if _, ok := blobs[string(agentTurn.GetUserMessage())]; !ok {
|
||||
t.Fatal("checkpoint turn references a missing user message Blob")
|
||||
}
|
||||
if agentTurn.GetRequestId() != "request-1" {
|
||||
t.Fatalf("checkpoint request id = %q, want original request", agentTurn.GetRequestId())
|
||||
}
|
||||
|
||||
steps := checkpointProjectionSteps(t, projection)
|
||||
if len(steps) != 2 || steps[0].GetAssistantMessage().GetText() != "before resume" || steps[1].GetAssistantMessage().GetText() != "after resume" {
|
||||
t.Fatalf("checkpoint steps did not preserve resumed activity: %#v", steps)
|
||||
}
|
||||
|
||||
replay, err := promptengine.DecodeReplayMessages(projection.State.GetRootPromptMessagesJson())
|
||||
if err != nil {
|
||||
t.Fatalf("decode root prompt replay: %v", err)
|
||||
}
|
||||
if len(replay) != 3 || replay[0].Role != "user" || replay[1].Content != "before resume" || replay[2].Content != "after resume" {
|
||||
t.Fatalf("checkpoint turn merge changed model replay: %#v", replay)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectCheckpointProjectionOmitsActivityWithoutAnyUserMessage(t *testing.T) {
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 2,
|
||||
Entries: []HistoryEntry{
|
||||
newAssistantTextEntry(1, "request-resume", "orphaned resume output", "", ""),
|
||||
},
|
||||
}
|
||||
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if len(projection.State.GetTurns()) != 0 || len(projection.Blobs) != 0 {
|
||||
t.Fatalf("checkpoint emitted a turn without a user message: %#v", projection)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectCheckpointProjectionKeepsForkPointIsolatedFromLaterHistory(t *testing.T) {
|
||||
firstUser, err := protojson.Marshal(&agentv1.UserMessage{Text: "first question", MessageId: "message-1"})
|
||||
if err != nil {
|
||||
@@ -138,3 +229,242 @@ func TestProjectCheckpointProjectionKeepsForkPointIsolatedFromLaterHistory(t *te
|
||||
t.Fatalf("fork snapshots are not isolated: midpoint=%#v latest=%#v", midpointMessages, latestMessages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectCheckpointProjectionMergesToolCallWithCompletedResult(t *testing.T) {
|
||||
userPayload, err := protojson.Marshal(&agentv1.UserMessage{Text: "inspect file", MessageId: "message-1"})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal user message: %v", err)
|
||||
}
|
||||
startedAt := uint64(100)
|
||||
toolCallID := "call-1"
|
||||
startedToolCall := checkpointTestToolCallPayload(t, &agentv1.ToolCall{
|
||||
ToolCallId: &toolCallID,
|
||||
StartedAtMs: &startedAt,
|
||||
Tool: &agentv1.ToolCall_ReadToolCall{
|
||||
ReadToolCall: &agentv1.ReadToolCall{
|
||||
Args: &agentv1.ReadToolArgs{Path: "/tmp/example.txt"},
|
||||
},
|
||||
},
|
||||
})
|
||||
completedAt := uint64(200)
|
||||
completedToolCall := checkpointTestToolCallPayload(t, &agentv1.ToolCall{
|
||||
CompletedAtMs: &completedAt,
|
||||
Tool: &agentv1.ToolCall_ReadToolCall{
|
||||
ReadToolCall: &agentv1.ReadToolCall{
|
||||
Result: &agentv1.ReadToolResult{
|
||||
Result: &agentv1.ReadToolResult_Success{
|
||||
Success: &agentv1.ReadToolSuccess{
|
||||
Path: "/tmp/example.txt",
|
||||
TotalLines: 1,
|
||||
Output: &agentv1.ReadToolSuccess_Content{Content: "file contents"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 2,
|
||||
Entries: []HistoryEntry{
|
||||
{Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: userPayload},
|
||||
newAssistantTextEntry(1, "request-1", "before", "", ""),
|
||||
newToolCallEntry(1, "request-1", "call-1", "Read", "", "", startedToolCall),
|
||||
newToolResultEntry(1, "request-1", "call-1", "Read", `{"path":"/tmp/example.txt"}`, "file contents", "", completedToolCall),
|
||||
newAssistantTextEntry(1, "request-1", "after", "", ""),
|
||||
},
|
||||
}
|
||||
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
if len(projection.Blobs) != 5 {
|
||||
t.Fatalf("checkpoint blobs = %d, want user, three final steps, and turn", len(projection.Blobs))
|
||||
}
|
||||
steps := checkpointProjectionSteps(t, projection)
|
||||
if len(steps) != 3 {
|
||||
t.Fatalf("checkpoint steps = %d, want assistant, completed Read, assistant", len(steps))
|
||||
}
|
||||
if steps[0].GetAssistantMessage().GetText() != "before" || steps[2].GetAssistantMessage().GetText() != "after" {
|
||||
t.Fatalf("checkpoint step ordering changed: %#v", steps)
|
||||
}
|
||||
mergedToolCall := steps[1].GetToolCall()
|
||||
readCall := mergedToolCall.GetReadToolCall()
|
||||
if readCall == nil || readCall.GetResult().GetSuccess().GetContent() != "file contents" {
|
||||
t.Fatalf("checkpoint Read step does not contain completed result: %#v", steps[1].GetToolCall())
|
||||
}
|
||||
if readCall.GetArgs().GetPath() != "/tmp/example.txt" || mergedToolCall.GetToolCallId() != toolCallID || mergedToolCall.GetStartedAtMs() != startedAt || mergedToolCall.GetCompletedAtMs() != completedAt {
|
||||
t.Fatalf("checkpoint Read step lost started-call fields: %#v", mergedToolCall)
|
||||
}
|
||||
|
||||
replay, err := promptengine.DecodeReplayMessages(projection.State.GetRootPromptMessagesJson())
|
||||
if err != nil {
|
||||
t.Fatalf("decode root prompt replay: %v", err)
|
||||
}
|
||||
for _, message := range replay {
|
||||
if message.Name == "Read" || len(message.ToolCalls) > 0 {
|
||||
t.Fatalf("UI-only Read result leaked into root prompt replay: %#v", replay)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectCheckpointProjectionIsIdempotentAndDoesNotMutateHistory(t *testing.T) {
|
||||
userPayload, err := protojson.Marshal(&agentv1.UserMessage{Text: "inspect file", MessageId: "message-1"})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal user message: %v", err)
|
||||
}
|
||||
completedToolCall := checkpointTestReadToolCall(t, &agentv1.ReadToolResult{
|
||||
Result: &agentv1.ReadToolResult_Success{
|
||||
Success: &agentv1.ReadToolSuccess{
|
||||
Path: "/tmp/example.txt",
|
||||
TotalLines: 1,
|
||||
Output: &agentv1.ReadToolSuccess_Content{Content: "file contents"},
|
||||
},
|
||||
},
|
||||
})
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 2,
|
||||
Entries: []HistoryEntry{
|
||||
{Seq: 1, TurnSeq: 1, RequestID: "request-1", Role: "user", Kind: "user_message", Payload: userPayload},
|
||||
newToolCallEntry(1, "request-1", "call-1", "Read", "", "", checkpointTestReadToolCall(t, nil)),
|
||||
newToolResultEntry(1, "request-1", "call-1", "Read", `{"path":"/tmp/example.txt"}`, "file contents", "", completedToolCall),
|
||||
},
|
||||
}
|
||||
before, err := json.Marshal(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal conversation before projection: %v", err)
|
||||
}
|
||||
|
||||
projector := NewHistoryProjector()
|
||||
first, err := projector.ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("first projection: %v", err)
|
||||
}
|
||||
second, err := projector.ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("second projection: %v", err)
|
||||
}
|
||||
|
||||
if !proto.Equal(first.State, second.State) {
|
||||
t.Fatalf("repeated projection changed checkpoint state: first=%#v second=%#v", first.State, second.State)
|
||||
}
|
||||
if len(first.Blobs) != len(second.Blobs) {
|
||||
t.Fatalf("repeated projection changed blob count: first=%d second=%d", len(first.Blobs), len(second.Blobs))
|
||||
}
|
||||
for index := range first.Blobs {
|
||||
if !bytes.Equal(first.Blobs[index].ID, second.Blobs[index].ID) || !bytes.Equal(first.Blobs[index].Data, second.Blobs[index].Data) {
|
||||
t.Fatalf("repeated projection changed blob %d", index)
|
||||
}
|
||||
}
|
||||
after, err := json.Marshal(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal conversation after projection: %v", err)
|
||||
}
|
||||
if !bytes.Equal(before, after) {
|
||||
t.Fatalf("checkpoint projection mutated semantic history:\nbefore=%s\nafter=%s", before, after)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectCheckpointProjectionKeepsStartedToolCallWhenResultPayloadIsMissing(t *testing.T) {
|
||||
startedToolCall := checkpointTestReadToolCall(t, nil)
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 2,
|
||||
Entries: []HistoryEntry{
|
||||
testCheckpointUserEntry(t),
|
||||
newToolCallEntry(1, "request-1", "call-1", "Read", "", "", startedToolCall),
|
||||
newToolResultEntry(1, "request-1", "call-1", "Read", `{"path":"/tmp/example.txt"}`, "read failed", "", nil),
|
||||
},
|
||||
}
|
||||
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
steps := checkpointProjectionSteps(t, projection)
|
||||
if len(steps) != 1 {
|
||||
t.Fatalf("checkpoint steps = %d, want the original Read step", len(steps))
|
||||
}
|
||||
readCall := steps[0].GetToolCall().GetReadToolCall()
|
||||
if readCall == nil || readCall.GetArgs().GetPath() != "/tmp/example.txt" || readCall.GetResult() != nil {
|
||||
t.Fatalf("checkpoint did not preserve the original Read call: %#v", steps[0].GetToolCall())
|
||||
}
|
||||
}
|
||||
|
||||
func TestProjectCheckpointProjectionAppendsLegacyResultWithoutToolCallEntry(t *testing.T) {
|
||||
completedToolCall := checkpointTestReadToolCall(t, &agentv1.ReadToolResult{
|
||||
Result: &agentv1.ReadToolResult_Error{Error: &agentv1.ReadToolError{ErrorMessage: "not readable"}},
|
||||
})
|
||||
conversation := &ConversationFile{
|
||||
ConversationID: "conversation-1",
|
||||
Mode: "agent",
|
||||
NextTurnSeq: 2,
|
||||
Entries: []HistoryEntry{
|
||||
testCheckpointUserEntry(t),
|
||||
newToolResultEntry(1, "request-1", "call-1", "Read", `{"path":"/tmp/example.txt"}`, "not readable", "", completedToolCall),
|
||||
},
|
||||
}
|
||||
|
||||
projection, err := NewHistoryProjector().ProjectCheckpointProjection(conversation)
|
||||
if err != nil {
|
||||
t.Fatalf("ProjectCheckpointProjection() error = %v", err)
|
||||
}
|
||||
steps := checkpointProjectionSteps(t, projection)
|
||||
if len(steps) != 1 || steps[0].GetToolCall().GetReadToolCall().GetResult().GetError().GetErrorMessage() != "not readable" {
|
||||
t.Fatalf("legacy result-only Read step was not preserved: %#v", steps)
|
||||
}
|
||||
}
|
||||
|
||||
func checkpointTestReadToolCall(t *testing.T, result *agentv1.ReadToolResult) []byte {
|
||||
t.Helper()
|
||||
return checkpointTestToolCallPayload(t, &agentv1.ToolCall{
|
||||
Tool: &agentv1.ToolCall_ReadToolCall{
|
||||
ReadToolCall: &agentv1.ReadToolCall{
|
||||
Args: &agentv1.ReadToolArgs{Path: "/tmp/example.txt"},
|
||||
Result: result,
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func checkpointTestToolCallPayload(t *testing.T, toolCall *agentv1.ToolCall) []byte {
|
||||
t.Helper()
|
||||
payload, err := protojson.Marshal(toolCall)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal Read tool call: %v", err)
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
func checkpointProjectionSteps(t *testing.T, projection *CheckpointProjection) []*agentv1.ConversationStep {
|
||||
t.Helper()
|
||||
if projection == nil || projection.State == nil || len(projection.State.GetTurns()) != 1 {
|
||||
t.Fatalf("checkpoint turns = %#v, want exactly one turn", projection)
|
||||
}
|
||||
blobs := make(map[string][]byte, len(projection.Blobs))
|
||||
for _, blob := range projection.Blobs {
|
||||
blobs[string(blob.ID)] = blob.Data
|
||||
}
|
||||
turn := &agentv1.ConversationTurnStructure{}
|
||||
if err := proto.Unmarshal(blobs[string(projection.State.GetTurns()[0])], turn); err != nil {
|
||||
t.Fatalf("decode checkpoint turn: %v", err)
|
||||
}
|
||||
agentTurn := turn.GetAgentConversationTurn()
|
||||
if agentTurn == nil {
|
||||
t.Fatal("checkpoint turn does not contain an agent turn")
|
||||
}
|
||||
steps := make([]*agentv1.ConversationStep, 0, len(agentTurn.GetSteps()))
|
||||
for _, stepID := range agentTurn.GetSteps() {
|
||||
step := &agentv1.ConversationStep{}
|
||||
if err := proto.Unmarshal(blobs[string(stepID)], step); err != nil {
|
||||
t.Fatalf("decode checkpoint step: %v", err)
|
||||
}
|
||||
steps = append(steps, step)
|
||||
}
|
||||
return steps
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -3,6 +3,8 @@ package forwarder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
@@ -561,6 +563,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) {
|
||||
@@ -608,6 +611,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
|
||||
@@ -838,14 +842,22 @@ func (service *Service) handleCancelIntent(intent InboundIntent) error {
|
||||
}
|
||||
hasCheckpoint := checkpointConversationInitialized(stream)
|
||||
if hasCheckpoint {
|
||||
preservedInterruptedOutput, err := service.persistInterruptedProviderOutput(stream)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
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),
|
||||
}),
|
||||
replayPolicy := cancelReplayPolicyForReason(cancelReason)
|
||||
if preservedInterruptedOutput || checkpointTurnHasReplayActivity(stream) {
|
||||
replayPolicy = cancelReplayPolicyKeepInterrupted
|
||||
}
|
||||
cancelEntry := newMetadataEntry(stream.TurnSeq, intent.RequestID, "control", map[string]any{
|
||||
"status": "canceled",
|
||||
"reason": cancelReason,
|
||||
"replay_policy": replayPolicy,
|
||||
})
|
||||
cancelEntry.IdempotencyKey = cancelMetadataIdempotencyKey(stream.TurnSeq, intent.RequestID)
|
||||
_, err = service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{cancelEntry})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -873,6 +885,89 @@ func (service *Service) handleCancelIntent(intent InboundIntent) error {
|
||||
return service.broker.Cancel(intent.RequestID, firstNonEmpty(intent.CancelReason, "[canceled] User aborted request"))
|
||||
}
|
||||
|
||||
func checkpointTurnHasReplayActivity(stream *ActiveStream) bool {
|
||||
if stream == nil {
|
||||
return false
|
||||
}
|
||||
stream.mu.Lock()
|
||||
defer stream.mu.Unlock()
|
||||
if stream.CheckpointConversation == nil {
|
||||
return false
|
||||
}
|
||||
for _, entry := range stream.CheckpointConversation.Entries {
|
||||
if entry.TurnSeq == stream.TurnSeq && isCanceledTurnActivityEntry(entry) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// persistInterruptedProviderOutput commits the current provider pass before cancellation.
|
||||
// The entry key is stable for this provider pass, so repeated cancellation handling is a no-op.
|
||||
func (service *Service) persistInterruptedProviderOutput(stream *ActiveStream) (bool, error) {
|
||||
if stream == nil {
|
||||
return false, nil
|
||||
}
|
||||
stream.mu.Lock()
|
||||
turnSeq := stream.TurnSeq
|
||||
requestID := strings.TrimSpace(stream.RequestID)
|
||||
modelCallID := strings.TrimSpace(stream.CurrentModelCallID)
|
||||
providerPass := stream.ProviderPassCount
|
||||
text := stream.ProviderAccumulatedText
|
||||
reasoning := stream.ProviderAccumulatedReasoning
|
||||
reasoningSignature := stream.ProviderAccumulatedReasoningSignature
|
||||
reasoningSignatureSource := stream.ProviderAccumulatedReasoningSignatureSource
|
||||
reasoningItemID := stream.ProviderAccumulatedReasoningItemID
|
||||
reasoningStatus := stream.ProviderAccumulatedReasoningStatus
|
||||
reasoningSummary := append([]byte(nil), stream.ProviderAccumulatedReasoningSummary...)
|
||||
stream.mu.Unlock()
|
||||
if strings.TrimSpace(text) == "" && !hasReplayableReasoningPayload(reasoning, reasoningSignature, reasoningSignatureSource) {
|
||||
return false, nil
|
||||
}
|
||||
key := interruptedProviderOutputIdempotencyKey(turnSeq, requestID, modelCallID, providerPass)
|
||||
_, err := service.appendConversationEntries(stream, stream.ConversationID, []HistoryEntry{
|
||||
{
|
||||
TurnSeq: turnSeq,
|
||||
RequestID: requestID,
|
||||
IdempotencyKey: key,
|
||||
Role: "assistant",
|
||||
Kind: "assistant_text",
|
||||
Payload: newAssistantTextPayload(
|
||||
text,
|
||||
reasoning,
|
||||
reasoningSignature,
|
||||
reasoningSignatureSource,
|
||||
reasoningItemID,
|
||||
reasoningStatus,
|
||||
reasoningSummary,
|
||||
),
|
||||
},
|
||||
})
|
||||
return true, err
|
||||
}
|
||||
|
||||
func interruptedProviderOutputIdempotencyKey(turnSeq int64, requestID string, modelCallID string, providerPass int) string {
|
||||
payload := strings.Join([]string{
|
||||
"provider_interrupted_output",
|
||||
fmt.Sprintf("%d", turnSeq),
|
||||
strings.TrimSpace(requestID),
|
||||
strings.TrimSpace(modelCallID),
|
||||
fmt.Sprintf("%d", providerPass),
|
||||
}, "\x00")
|
||||
digest := sha256.Sum256([]byte(payload))
|
||||
return "provider-interrupted-output:" + hex.EncodeToString(digest[:])
|
||||
}
|
||||
|
||||
func cancelMetadataIdempotencyKey(turnSeq int64, requestID string) string {
|
||||
payload := strings.Join([]string{
|
||||
"cancel",
|
||||
fmt.Sprintf("%d", turnSeq),
|
||||
strings.TrimSpace(requestID),
|
||||
}, "\x00")
|
||||
digest := sha256.Sum256([]byte(payload))
|
||||
return "cancel:" + hex.EncodeToString(digest[:])
|
||||
}
|
||||
|
||||
// handleExecResult 处理客户端返回的执行桥结果,并在终态时把 tool_result 写回 history。
|
||||
func (service *Service) handleExecResult(intent InboundIntent) error {
|
||||
stream, ok := service.broker.Get(intent.RequestID)
|
||||
@@ -914,6 +1009,11 @@ func (service *Service) handleExecResult(intent InboundIntent) error {
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
if message := buildShellToolCallDeltaMessage(pending.ToolCallID, pending.ModelCallID, result.ShellOutputDelta); message != nil {
|
||||
if err := service.broker.Publish(intent.RequestID, StreamEvent{Message: message}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
if !result.IsTerminal {
|
||||
return nil
|
||||
@@ -1665,6 +1765,13 @@ func (service *Service) handleToolInvocation(stream *ActiveStream, invocation ru
|
||||
stream.ToolInvocationCount++
|
||||
stream.UpdatedAt = time.Now().UTC()
|
||||
stream.mu.Unlock()
|
||||
if !isKnownToolName(trimmedToolName) {
|
||||
displayToolName := trimmedToolName
|
||||
if displayToolName == "" {
|
||||
displayToolName = "<empty>"
|
||||
}
|
||||
return service.completePreDispatchToolError(stream, invocation, nil, false, false, fmt.Errorf("Model hallucination: attempted to invoke a nonexistent tool: %s", displayToolName))
|
||||
}
|
||||
if !isToolAllowedInMode(mode, subagentTypeName, trimmedToolName) {
|
||||
return service.completePreDispatchToolError(stream, invocation, nil, false, false, fmt.Errorf("tool invocation is not enabled in mode %s: %s", mode.String(), invocation.ToolName))
|
||||
}
|
||||
@@ -2163,6 +2270,15 @@ func (service *Service) finishSuccessfulTurnAfterCheckpoint(stream *ActiveStream
|
||||
return nil
|
||||
}
|
||||
|
||||
func (service *Service) finishFailedTurnAfterCheckpoint(stream *ActiveStream, terminalCode string, terminalMessage string) error {
|
||||
if stream == nil {
|
||||
return nil
|
||||
}
|
||||
err := service.broker.Fail(stream.RequestID, terminalCode, terminalMessage)
|
||||
service.setTurnPhase(stream, TurnPhaseFailed)
|
||||
return err
|
||||
}
|
||||
|
||||
func (service *Service) failStreamIfNonTerminal(stream *ActiveStream, terminalCode string, cause error) error {
|
||||
if stream == nil || cause == nil {
|
||||
return nil
|
||||
@@ -2182,6 +2298,10 @@ func (service *Service) publishCheckpoint(requestID string, conversationID strin
|
||||
}
|
||||
|
||||
func (service *Service) publishCheckpointWithCompletion(requestID string, _ string, completion *pendingTurnCompletion) error {
|
||||
return service.publishCheckpointWithTerminalAction(requestID, successfulCheckpointTerminalAction(completion))
|
||||
}
|
||||
|
||||
func (service *Service) publishCheckpointWithTerminalAction(requestID string, terminal checkpointTerminalAction) error {
|
||||
stream, ok := service.broker.Get(requestID)
|
||||
if !ok || stream == nil {
|
||||
return fmt.Errorf("request is not active: %s", requestID)
|
||||
@@ -2199,7 +2319,7 @@ func (service *Service) publishCheckpointWithCompletion(requestID string, _ stri
|
||||
}
|
||||
projection.State.PendingToolCalls = buildPendingToolCalls(pendingExecs, pendingInteractions)
|
||||
service.rewriteCheckpointTokenDetailsForClient(stream, conversation, projection.State)
|
||||
return service.queueCheckpointProjection(stream, projection, completion)
|
||||
return service.queueCheckpointProjectionWithTerminal(stream, projection, terminal)
|
||||
}
|
||||
|
||||
func (service *Service) rewriteCheckpointTokenDetailsForClient(stream *ActiveStream, conversation *ConversationFile, state *agentv1.ConversationStateStructure) {
|
||||
@@ -2339,18 +2459,20 @@ func (service *Service) failActiveStream(stream *ActiveStream, conversationID st
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
service.setTurnPhase(stream, TurnPhaseFailed)
|
||||
var firstErr error
|
||||
if err := service.syncSummaryCarryForward(conversationID, requestID, modelCallID); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
if err := service.syncSummaryCarryForward(conversationID, requestID, modelCallID); err != nil {
|
||||
log.Printf(
|
||||
"forwarder summary sync before failed terminal skipped request_id=%s model_call_id=%s err=%v",
|
||||
strings.TrimSpace(requestID),
|
||||
strings.TrimSpace(modelCallID),
|
||||
err,
|
||||
)
|
||||
}
|
||||
if err := service.publishCheckpoint(requestID, conversationID); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
terminal := failedCheckpointTerminalAction(terminalCode, terminalMessage)
|
||||
if err := service.publishCheckpointWithTerminalAction(requestID, terminal); err != nil {
|
||||
log.Printf("forwarder checkpoint queue before failed terminal skipped request_id=%s err=%v", strings.TrimSpace(requestID), err)
|
||||
return service.finishFailedTurnAfterCheckpoint(stream, terminalCode, terminalMessage)
|
||||
}
|
||||
if err := service.broker.Fail(requestID, terminalCode, terminalMessage); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
return firstErr
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildRunEntries 构造一次 run intent 需要写入 history 的首批 entry。
|
||||
@@ -2453,6 +2575,16 @@ func newAssistantTextEntry(turnSeq int64, requestID string, text string, reasoni
|
||||
}
|
||||
|
||||
func newAssistantTextEntryWithProviderMetadata(turnSeq int64, requestID string, text string, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, reasoningItemID string, reasoningStatus string, reasoningSummary json.RawMessage) HistoryEntry {
|
||||
return HistoryEntry{
|
||||
TurnSeq: turnSeq,
|
||||
RequestID: strings.TrimSpace(requestID),
|
||||
Role: "assistant",
|
||||
Kind: "assistant_text",
|
||||
Payload: newAssistantTextPayload(text, reasoningContent, reasoningSignature, reasoningSignatureSource, reasoningItemID, reasoningStatus, reasoningSummary),
|
||||
}
|
||||
}
|
||||
|
||||
func newAssistantTextPayload(text string, reasoningContent string, reasoningSignature string, reasoningSignatureSource string, reasoningItemID string, reasoningStatus string, reasoningSummary json.RawMessage) json.RawMessage {
|
||||
payload, _ := json.Marshal(assistantTextPayload{
|
||||
Text: text,
|
||||
ReasoningContent: reasoningContent,
|
||||
@@ -2462,13 +2594,7 @@ func newAssistantTextEntryWithProviderMetadata(turnSeq int64, requestID string,
|
||||
ReasoningStatus: strings.TrimSpace(reasoningStatus),
|
||||
ReasoningSummary: append(json.RawMessage(nil), reasoningSummary...),
|
||||
})
|
||||
return HistoryEntry{
|
||||
TurnSeq: turnSeq,
|
||||
RequestID: strings.TrimSpace(requestID),
|
||||
Role: "assistant",
|
||||
Kind: "assistant_text",
|
||||
Payload: payload,
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
// newToolCallEntry 构造 tool_call entry。
|
||||
|
||||
@@ -0,0 +1,150 @@
|
||||
package forwarder
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"cursor/gen/agentv1"
|
||||
execbridge "cursor/internal/backend/agent/bridge/exec"
|
||||
runtimecore "cursor/internal/backend/agent/core"
|
||||
)
|
||||
|
||||
func TestHandleExecResultPublishesShellToolCallDelta(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
shellStream func() *agentv1.ShellStream
|
||||
wantStdout string
|
||||
wantStderr string
|
||||
}{
|
||||
{
|
||||
name: "stdout",
|
||||
shellStream: func() *agentv1.ShellStream {
|
||||
return &agentv1.ShellStream{Event: &agentv1.ShellStream_Stdout{
|
||||
Stdout: &agentv1.ShellStreamStdout{Data: "stdout chunk\n"},
|
||||
}}
|
||||
},
|
||||
wantStdout: "stdout chunk\n",
|
||||
},
|
||||
{
|
||||
name: "stderr",
|
||||
shellStream: func() *agentv1.ShellStream {
|
||||
return &agentv1.ShellStream{Event: &agentv1.ShellStream_Stderr{
|
||||
Stderr: &agentv1.ShellStreamStderr{Data: "stderr chunk\n"},
|
||||
}}
|
||||
},
|
||||
wantStderr: "stderr chunk\n",
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
broker := NewStreamBroker()
|
||||
service := &Service{
|
||||
broker: broker,
|
||||
execBridge: execbridge.NewBridge(),
|
||||
}
|
||||
stream, err := broker.OpenStream(
|
||||
"request-1", "conversation-1", 1, "default", "default",
|
||||
agentv1.AgentMode_AGENT_MODE_AGENT, "run command",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("OpenStream() error = %v", err)
|
||||
}
|
||||
pending := runtimecore.PendingExec{
|
||||
MessageID: 42,
|
||||
ExecID: "exec-shell-1",
|
||||
ModelCallID: "model-call-1",
|
||||
ToolCallID: "tool-call-1",
|
||||
ExecKind: "shell",
|
||||
}
|
||||
stream.mu.Lock()
|
||||
stream.PendingExecs[pending.ExecID] = pending
|
||||
stream.mu.Unlock()
|
||||
|
||||
if err := service.handleExecResult(InboundIntent{
|
||||
Kind: "exec_result",
|
||||
RequestID: "request-1",
|
||||
ExecClientMessage: &agentv1.ExecClientMessage{
|
||||
Id: pending.MessageID,
|
||||
ExecId: pending.ExecID,
|
||||
Message: &agentv1.ExecClientMessage_ShellStream{
|
||||
ShellStream: test.shellStream(),
|
||||
},
|
||||
},
|
||||
}); err != nil {
|
||||
t.Fatalf("handleExecResult() error = %v", err)
|
||||
}
|
||||
|
||||
events, err := broker.ReadFromCursor("request-1", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFromCursor() error = %v", err)
|
||||
}
|
||||
if len(events) != 2 {
|
||||
t.Fatalf("published events = %d, want compatibility and tool-call deltas", len(events))
|
||||
}
|
||||
|
||||
var compatibilityCount, toolCallDeltaCount int
|
||||
for _, event := range events {
|
||||
update := event.Message.GetInteractionUpdate()
|
||||
if update.GetShellOutputDelta() != nil {
|
||||
compatibilityCount++
|
||||
}
|
||||
deltaUpdate := update.GetToolCallDelta()
|
||||
if deltaUpdate == nil {
|
||||
continue
|
||||
}
|
||||
toolCallDeltaCount++
|
||||
if deltaUpdate.GetCallId() != pending.ToolCallID || deltaUpdate.GetModelCallId() != pending.ModelCallID {
|
||||
t.Fatalf("tool-call delta ids = call %q model %q", deltaUpdate.GetCallId(), deltaUpdate.GetModelCallId())
|
||||
}
|
||||
shellDelta := deltaUpdate.GetToolCallDelta().GetShellToolCallDelta()
|
||||
if shellDelta == nil || shellDelta.GetStdout().GetContent() != test.wantStdout || shellDelta.GetStderr().GetContent() != test.wantStderr {
|
||||
t.Fatalf("shell tool-call delta = %#v", shellDelta)
|
||||
}
|
||||
}
|
||||
if compatibilityCount != 1 || toolCallDeltaCount != 1 {
|
||||
t.Fatalf("published compatibility=%d tool_call_delta=%d, want one each", compatibilityCount, toolCallDeltaCount)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildShellToolCallDeltaMessageIgnoresNonOutputEvents(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
output *agentv1.ShellOutputDeltaUpdate
|
||||
}{
|
||||
{name: "nil"},
|
||||
{
|
||||
name: "start",
|
||||
output: &agentv1.ShellOutputDeltaUpdate{Event: &agentv1.ShellOutputDeltaUpdate_Start{
|
||||
Start: &agentv1.ShellStreamStart{},
|
||||
}},
|
||||
},
|
||||
{
|
||||
name: "exit",
|
||||
output: &agentv1.ShellOutputDeltaUpdate{Event: &agentv1.ShellOutputDeltaUpdate_Exit{
|
||||
Exit: &agentv1.ShellStreamExit{},
|
||||
}},
|
||||
},
|
||||
{
|
||||
name: "empty stdout",
|
||||
output: &agentv1.ShellOutputDeltaUpdate{Event: &agentv1.ShellOutputDeltaUpdate_Stdout{
|
||||
Stdout: &agentv1.ShellStreamStdout{},
|
||||
}},
|
||||
},
|
||||
{
|
||||
name: "empty stderr",
|
||||
output: &agentv1.ShellOutputDeltaUpdate{Event: &agentv1.ShellOutputDeltaUpdate_Stderr{
|
||||
Stderr: &agentv1.ShellStreamStderr{},
|
||||
}},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if message := buildShellToolCallDeltaMessage("tool-call-1", "model-call-1", test.output); message != nil {
|
||||
t.Fatalf("buildShellToolCallDeltaMessage() = %#v, want nil", message)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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 := importedConversationStateModelMessagesWithBlobs(state, blobs); err != nil {
|
||||
return nil, err
|
||||
} else {
|
||||
for _, message := range messages {
|
||||
@@ -105,6 +117,10 @@ func (service *Service) importConversationState(item *ConversationFile, state *a
|
||||
}
|
||||
|
||||
func importedConversationStateModelMessages(state *agentv1.ConversationStateStructure) ([]modeladapter.Message, error) {
|
||||
return importedConversationStateModelMessagesWithBlobs(state, nil)
|
||||
}
|
||||
|
||||
func importedConversationStateModelMessagesWithBlobs(state *agentv1.ConversationStateStructure, blobs importedBlobStore) ([]modeladapter.Message, error) {
|
||||
if state == nil {
|
||||
return nil, nil
|
||||
}
|
||||
@@ -113,7 +129,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 +149,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
|
||||
}
|
||||
|
||||
@@ -181,6 +181,25 @@ func supportedToolNamesForMode(mode agentv1.AgentMode) map[string]struct{} {
|
||||
}
|
||||
}
|
||||
|
||||
func isKnownToolName(toolName string) bool {
|
||||
trimmedToolName := strings.TrimSpace(toolName)
|
||||
if trimmedToolName == "" {
|
||||
return false
|
||||
}
|
||||
for _, supported := range []map[string]struct{}{
|
||||
agentModeToolNames,
|
||||
askModeToolNames,
|
||||
planModeToolNames,
|
||||
debugModeToolNames,
|
||||
multitaskModeToolNames,
|
||||
} {
|
||||
if _, ok := supported[trimmedToolName]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func isToolAllowedInMode(mode agentv1.AgentMode, subagentTypeName string, toolName string) bool {
|
||||
trimmedToolName := strings.TrimSpace(toolName)
|
||||
if trimmedToolName == "" {
|
||||
|
||||
@@ -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"`
|
||||
@@ -79,6 +80,7 @@ type HistoryEntry struct {
|
||||
Seq int64 `json:"seq"`
|
||||
TurnSeq int64 `json:"turn_seq"`
|
||||
RequestID string `json:"request_id,omitempty"`
|
||||
IdempotencyKey string `json:"idempotency_key,omitempty"`
|
||||
Role string `json:"role"`
|
||||
Kind string `json:"kind"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
@@ -223,10 +225,25 @@ type pendingTurnCompletion struct {
|
||||
Disposition pendingCompletionDisposition
|
||||
}
|
||||
|
||||
type checkpointTerminalActionKind uint8
|
||||
|
||||
const (
|
||||
checkpointTerminalActionNone checkpointTerminalActionKind = iota
|
||||
checkpointTerminalActionComplete
|
||||
checkpointTerminalActionFail
|
||||
)
|
||||
|
||||
type checkpointTerminalAction struct {
|
||||
Kind checkpointTerminalActionKind
|
||||
Completion pendingTurnCompletion
|
||||
ErrorCode string
|
||||
ErrorMessage string
|
||||
}
|
||||
|
||||
type pendingCheckpointPublish struct {
|
||||
State *agentv1.ConversationStateStructure
|
||||
Required map[string]struct{}
|
||||
Completion *pendingTurnCompletion
|
||||
State *agentv1.ConversationStateStructure
|
||||
Required map[string]struct{}
|
||||
Terminal checkpointTerminalAction
|
||||
}
|
||||
|
||||
type PendingCompaction struct {
|
||||
@@ -428,6 +445,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
|
||||
|
||||
@@ -16,6 +16,7 @@ func (store *Store) LegacyRuntimeSnapshot(ctx context.Context) (legacyruntime.Ru
|
||||
for _, item := range cfg.ModelAdapters {
|
||||
adapters = append(adapters, legacyruntime.ModelAdapterConfig{
|
||||
ID: item.ID,
|
||||
Sort: item.Sort,
|
||||
DisplayName: item.DisplayName,
|
||||
Type: item.Type,
|
||||
BaseURL: item.BaseURL,
|
||||
|
||||
@@ -147,6 +147,7 @@ func (manager *Manager) LegacyRuntimeSnapshot(_ context.Context) (legacyruntime.
|
||||
for _, item := range cfg.ModelAdapters {
|
||||
adapters = append(adapters, legacyruntime.ModelAdapterConfig{
|
||||
ID: item.ID,
|
||||
Sort: item.Sort,
|
||||
DisplayName: item.DisplayName,
|
||||
Type: item.Type,
|
||||
BaseURL: item.BaseURL,
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
@@ -21,6 +22,7 @@ const (
|
||||
|
||||
type ModelAdapterConfig struct {
|
||||
ID string `json:"id,omitempty" yaml:"-"`
|
||||
Sort int `json:"sort" yaml:"sort"`
|
||||
DisplayName string `json:"displayName" yaml:"displayName"`
|
||||
Type string `json:"type" yaml:"type"`
|
||||
BaseURL string `json:"baseURL" yaml:"baseURL"`
|
||||
@@ -104,6 +106,7 @@ func NormalizeModelAdapterConfigs(input []ModelAdapterConfig) ([]ModelAdapterCon
|
||||
}
|
||||
nextType := normalizeModelAdapterType(item.Type)
|
||||
next := ModelAdapterConfig{
|
||||
Sort: item.Sort,
|
||||
DisplayName: strings.TrimSpace(item.DisplayName),
|
||||
Type: nextType,
|
||||
BaseURL: baseURL,
|
||||
@@ -138,8 +141,8 @@ func NormalizeModelAdapterConfigs(input []ModelAdapterConfig) ([]ModelAdapterCon
|
||||
return nil, errors.New("模型适配器 tooltipData 不能为空")
|
||||
case next.ModelID == "":
|
||||
return nil, errors.New("模型适配器 modelID 不能为空")
|
||||
case next.Type == "openai" && next.ReasoningEffort == "":
|
||||
return nil, errors.New("模型适配器 reasoningEffort 仅支持 low、medium、high、xhigh、max")
|
||||
case next.Type == "openai" && !isSupportedReasoningEffort(next.ReasoningEffort):
|
||||
return nil, errors.New("模型适配器 reasoningEffort 仅支持空值、low、medium、high、xhigh、max")
|
||||
case next.Type == "openai" && next.OpenAIEndpoint == "":
|
||||
return nil, errors.New("模型适配器 openAIEndpoint 仅支持 /v1/responses、/v1/chat/completions 或 /custom(自定义路径)")
|
||||
case next.Type == "openai" && next.OpenAIExtraParamsEnabled:
|
||||
@@ -164,9 +167,30 @@ func NormalizeModelAdapterConfigs(input []ModelAdapterConfig) ([]ModelAdapterCon
|
||||
seenChannelIDs[next.ID] = struct{}{}
|
||||
normalized = append(normalized, next)
|
||||
}
|
||||
normalizeModelAdapterSorts(normalized)
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeModelAdapterSorts(adapters []ModelAdapterConfig) {
|
||||
sort.SliceStable(adapters, func(leftIndex, rightIndex int) bool {
|
||||
left := adapters[leftIndex].Sort
|
||||
right := adapters[rightIndex].Sort
|
||||
switch {
|
||||
case left <= 0 && right <= 0:
|
||||
return false
|
||||
case left <= 0:
|
||||
return false
|
||||
case right <= 0:
|
||||
return true
|
||||
default:
|
||||
return left < right
|
||||
}
|
||||
})
|
||||
for index := range adapters {
|
||||
adapters[index].Sort = index + 1
|
||||
}
|
||||
}
|
||||
|
||||
func validateJSONMap(value string, fieldName string) error {
|
||||
text := strings.TrimSpace(value)
|
||||
if text == "" {
|
||||
@@ -200,13 +224,15 @@ func validateHeadersJSON(value string) error {
|
||||
}
|
||||
|
||||
func normalizeReasoningEffort(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "medium":
|
||||
return "medium"
|
||||
case "low", "high", "xhigh", "max":
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
}
|
||||
|
||||
func isSupportedReasoningEffort(value string) bool {
|
||||
switch value {
|
||||
case "", "low", "medium", "high", "xhigh", "max":
|
||||
return true
|
||||
default:
|
||||
return ""
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
package config
|
||||
|
||||
import "testing"
|
||||
|
||||
func testModelAdapter(displayName string, sortValue int) ModelAdapterConfig {
|
||||
return ModelAdapterConfig{
|
||||
Sort: sortValue,
|
||||
DisplayName: displayName,
|
||||
Type: "openai",
|
||||
BaseURL: "https://api.example.com/v1",
|
||||
APIKey: "test-key",
|
||||
TooltipData: displayName,
|
||||
ModelID: displayName,
|
||||
ReasoningEffort: "medium",
|
||||
OpenAIEndpoint: "/v1/responses",
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeModelAdapterConfigsPreservesLegacyArrayOrder(t *testing.T) {
|
||||
adapters, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{
|
||||
testModelAdapter("first", 0),
|
||||
testModelAdapter("second", 0),
|
||||
testModelAdapter("third", 0),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NormalizeModelAdapterConfigs returned error: %v", err)
|
||||
}
|
||||
|
||||
for index, expectedName := range []string{"first", "second", "third"} {
|
||||
if adapters[index].DisplayName != expectedName {
|
||||
t.Fatalf("adapter %d = %q, want %q", index, adapters[index].DisplayName, expectedName)
|
||||
}
|
||||
if adapters[index].Sort != index+1 {
|
||||
t.Fatalf("adapter %d sort = %d, want %d", index, adapters[index].Sort, index+1)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeModelAdapterConfigsUsesStableExplicitSort(t *testing.T) {
|
||||
adapters, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{
|
||||
testModelAdapter("legacy", 0),
|
||||
testModelAdapter("third", 30),
|
||||
testModelAdapter("first", 10),
|
||||
testModelAdapter("second-a", 20),
|
||||
testModelAdapter("second-b", 20),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NormalizeModelAdapterConfigs returned error: %v", err)
|
||||
}
|
||||
|
||||
expectedNames := []string{"first", "second-a", "second-b", "third", "legacy"}
|
||||
for index, expectedName := range expectedNames {
|
||||
if adapters[index].DisplayName != expectedName {
|
||||
t.Fatalf("adapter %d = %q, want %q", index, adapters[index].DisplayName, expectedName)
|
||||
}
|
||||
if adapters[index].Sort != index+1 {
|
||||
t.Fatalf("adapter %d sort = %d, want %d", index, adapters[index].Sort, index+1)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeModelAdapterConfigsAllowsBlankReasoningEffort(t *testing.T) {
|
||||
adapter := testModelAdapter("non-reasoning-model", 1)
|
||||
adapter.ReasoningEffort = ""
|
||||
|
||||
adapters, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{adapter})
|
||||
if err != nil {
|
||||
t.Fatalf("NormalizeModelAdapterConfigs returned error: %v", err)
|
||||
}
|
||||
if got := adapters[0].ReasoningEffort; got != "" {
|
||||
t.Fatalf("ReasoningEffort = %q, want blank", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeModelAdapterConfigsRejectsUnknownReasoningEffort(t *testing.T) {
|
||||
adapter := testModelAdapter("invalid-reasoning-effort", 1)
|
||||
adapter.ReasoningEffort = "unsupported"
|
||||
|
||||
if _, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{adapter}); err == nil {
|
||||
t.Fatal("NormalizeModelAdapterConfigs should reject an unknown reasoning effort")
|
||||
}
|
||||
}
|
||||
@@ -145,7 +145,7 @@ var bootstrapStatsigTemplate = statsigBootstrapTemplate{
|
||||
bootstrapStatsigGlassCustomThemeSupport: buildEnabledStatsigGate(bootstrapStatsigGlassCustomThemeSupport),
|
||||
bootstrapStatsigGlassAutomationsUI: buildEnabledStatsigGate(bootstrapStatsigGlassAutomationsUI),
|
||||
bootstrapStatsigTerminalUI2: buildEnabledStatsigGate(bootstrapStatsigTerminalUI2),
|
||||
bootstrapStatsigDisableTerminalOutputUIStreaming: buildEnabledStatsigGate(bootstrapStatsigDisableTerminalOutputUIStreaming),
|
||||
bootstrapStatsigDisableTerminalOutputUIStreaming: buildDisabledStatsigGate(bootstrapStatsigDisableTerminalOutputUIStreaming),
|
||||
bootstrapStatsigBrowserCanvas: buildEnabledStatsigGate(bootstrapStatsigBrowserCanvas),
|
||||
bootstrapStatsigEnableMultitaskMode: buildEnabledStatsigGate(bootstrapStatsigEnableMultitaskMode),
|
||||
bootstrapStatsigDecomposeAlwaysLocalExtHostGate: buildDisabledStatsigGate(bootstrapStatsigDecomposeAlwaysLocalExtHostGate),
|
||||
@@ -728,8 +728,10 @@ func buildCLIModelDetails(adapters []legacyruntime.ModelAdapterConfig) []map[str
|
||||
continue
|
||||
}
|
||||
models = append(models, map[string]any{
|
||||
"modelId": channelID,
|
||||
"displayModelId": channelID,
|
||||
"modelId": channelID,
|
||||
"displayModelId": channelID,
|
||||
"displayName": strings.TrimSpace(adapter.DisplayName),
|
||||
"displayNameShort": strings.TrimSpace(adapter.DisplayName),
|
||||
"apiKeyCredentials": map[string]any{
|
||||
"apiKey": strings.TrimSpace(adapter.APIKey),
|
||||
"baseUrl": strings.TrimSpace(adapter.BaseURL),
|
||||
@@ -833,7 +835,7 @@ func defaultThinkingEffortForAdapter(adapter legacyruntime.ModelAdapterConfig) s
|
||||
if strings.EqualFold(strings.TrimSpace(adapter.Type), "anthropic") {
|
||||
return normalizeAvailableModelThinkingEffort(adapter.AnthropicThinkingEffort, true, "xhigh")
|
||||
}
|
||||
return normalizeAvailableModelThinkingEffort(adapter.ReasoningEffort, true, "medium")
|
||||
return normalizeAvailableModelThinkingEffort(adapter.ReasoningEffort, true, "disabled")
|
||||
}
|
||||
|
||||
func normalizeAvailableModelThinkingEffort(raw string, allowMax bool, fallback string) string {
|
||||
|
||||
@@ -11,25 +11,59 @@ import (
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
func TestBuildCLIModelDetailsPreservesChannelCredentials(t *testing.T) {
|
||||
func TestBuildCLIModelDetailsPreservesChannelMetadata(t *testing.T) {
|
||||
adapters := []legacyruntime.ModelAdapterConfig{
|
||||
{ID: " channel-a ", ModelID: "model-a", APIKey: "provider-secret-a", BaseURL: "https://provider-a.example/v1"},
|
||||
{ID: "channel-b", ModelID: "model-a"},
|
||||
{ID: " channel-a ", DisplayName: " Model A ", ModelID: "model-a", APIKey: "provider-secret-a", BaseURL: "https://provider-a.example/v1"},
|
||||
{ID: "channel-b", DisplayName: "Model B", ModelID: "model-a"},
|
||||
{ID: "", ModelID: "model-c"},
|
||||
}
|
||||
|
||||
got := buildCLIModelDetails(adapters)
|
||||
want := []map[string]any{
|
||||
{"modelId": "channel-a", "displayModelId": "channel-a", "apiKeyCredentials": map[string]any{"apiKey": "provider-secret-a", "baseUrl": "https://provider-a.example/v1"}},
|
||||
{"modelId": "channel-b", "displayModelId": "channel-b", "apiKeyCredentials": map[string]any{"apiKey": "", "baseUrl": ""}},
|
||||
{"modelId": "channel-a", "displayModelId": "channel-a", "displayName": "Model A", "displayNameShort": "Model A", "apiKeyCredentials": map[string]any{"apiKey": "provider-secret-a", "baseUrl": "https://provider-a.example/v1"}},
|
||||
{"modelId": "channel-b", "displayModelId": "channel-b", "displayName": "Model B", "displayNameShort": "Model B", "apiKeyCredentials": map[string]any{"apiKey": "", "baseUrl": ""}},
|
||||
}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("build CLI model details: got %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultThinkingEffortForOpenAIAdapterUsesDisabledWhenUnset(t *testing.T) {
|
||||
adapter := legacyruntime.ModelAdapterConfig{Type: "openai", ReasoningEffort: ""}
|
||||
|
||||
if got := defaultThinkingEffortForAdapter(adapter); got != "disabled" {
|
||||
t.Fatalf("default thinking effort = %q, want disabled", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildAvailableModelEntriesUsesDisabledVariantWhenReasoningEffortUnset(t *testing.T) {
|
||||
entries := buildAvailableModelEntries([]legacyruntime.ModelAdapterConfig{{
|
||||
ID: "channel-a",
|
||||
DisplayName: "Model A",
|
||||
ModelID: "model-a",
|
||||
Type: "openai",
|
||||
}})
|
||||
if len(entries) != 1 {
|
||||
t.Fatalf("entry count = %d, want 1", len(entries))
|
||||
}
|
||||
|
||||
variants, ok := entries[0]["variants"].([]map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("variants type = %T, want []map[string]any", entries[0]["variants"])
|
||||
}
|
||||
if len(variants) == 0 {
|
||||
t.Fatal("variants should not be empty")
|
||||
}
|
||||
if got := variants[0]["variantStringRepresentation"]; got != "channel-a:disabled" {
|
||||
t.Fatalf("first variant representation = %#v, want channel-a:disabled", got)
|
||||
}
|
||||
if got := variants[0]["isDefaultNonMaxConfig"]; got != true {
|
||||
t.Fatalf("disabled variant default flag = %#v, want true", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEncodeCLIModelsUsesAgentModelDetailsWireFormat(t *testing.T) {
|
||||
payload := map[string]any{"models": buildCLIModelDetails([]legacyruntime.ModelAdapterConfig{{ID: "channel-a", APIKey: "provider-secret", BaseURL: "https://provider.example/v1"}})}
|
||||
payload := map[string]any{"models": buildCLIModelDetails([]legacyruntime.ModelAdapterConfig{{ID: "channel-a", DisplayName: "Model A", APIKey: "provider-secret", BaseURL: "https://provider.example/v1"}})}
|
||||
encoded, err := encodeMockProto("aiserver.v1.GetUsableModelsResponse", payload)
|
||||
if err != nil {
|
||||
t.Fatalf("encode CLI models: %v", err)
|
||||
@@ -46,6 +80,9 @@ func TestEncodeCLIModelsUsesAgentModelDetailsWireFormat(t *testing.T) {
|
||||
if model.GetModelId() != "channel-a" || model.GetDisplayModelId() != "channel-a" {
|
||||
t.Fatalf("decoded channel IDs: model=%q display=%q", model.GetModelId(), model.GetDisplayModelId())
|
||||
}
|
||||
if model.GetDisplayName() != "Model A" || model.GetDisplayNameShort() != "Model A" {
|
||||
t.Fatalf("decoded display names: name=%q short=%q", model.GetDisplayName(), model.GetDisplayNameShort())
|
||||
}
|
||||
if credentials := model.GetApiKeyCredentials(); credentials == nil || credentials.GetApiKey() != "provider-secret" || credentials.GetBaseUrl() != "https://provider.example/v1" {
|
||||
t.Fatalf("decoded relay credentials: %#v", credentials)
|
||||
}
|
||||
@@ -73,3 +110,26 @@ func TestBuildBootstrapStatsigConfigJSONDisablesAlwaysLocalDecompositionGate(t *
|
||||
t.Fatalf("unexpected rule_id: %q", ruleID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBootstrapStatsigConfigJSONEnablesTerminalOutputUIStreaming(t *testing.T) {
|
||||
payload, err := buildBootstrapStatsigConfigJSON(12345, "test-auth-id")
|
||||
if err != nil {
|
||||
t.Fatalf("build bootstrap statsig config: %v", err)
|
||||
}
|
||||
|
||||
var decoded statsigBootstrapTemplate
|
||||
if err := json.Unmarshal(payload, &decoded); err != nil {
|
||||
t.Fatalf("decode bootstrap statsig config: %v", err)
|
||||
}
|
||||
|
||||
gate, ok := decoded.FeatureGates[bootstrapStatsigDisableTerminalOutputUIStreaming]
|
||||
if !ok {
|
||||
t.Fatalf("missing feature gate %q", bootstrapStatsigDisableTerminalOutputUIStreaming)
|
||||
}
|
||||
if value, _ := gate["value"].(bool); value {
|
||||
t.Fatalf("expected %q to be disabled", bootstrapStatsigDisableTerminalOutputUIStreaming)
|
||||
}
|
||||
if ruleID, _ := gate["rule_id"].(string); ruleID != "local_disabled" {
|
||||
t.Fatalf("unexpected rule_id: %q", ruleID)
|
||||
}
|
||||
}
|
||||
|
||||
+2
-101
@@ -15,19 +15,11 @@ import (
|
||||
"github.com/wailsapp/wails/v3/pkg/events"
|
||||
)
|
||||
|
||||
// modelEditorContext 保存当前模型编辑器窗口的初始化上下文。
|
||||
type modelEditorContext struct {
|
||||
Index int `json:"index"`
|
||||
AdapterJSON string `json:"adapterJSON"`
|
||||
}
|
||||
|
||||
// WindowService 定义了当前模块中的 WindowService 类型。
|
||||
type WindowService struct {
|
||||
app *application.App
|
||||
updater *updater.Manager
|
||||
modelConfigWindow *application.WebviewWindow
|
||||
modelEditorWindow *application.WebviewWindow
|
||||
editorCtx *modelEditorContext
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
@@ -102,8 +94,8 @@ func (s *WindowService) OpenModelConfigWindow() {
|
||||
Title: "模型配置",
|
||||
Width: 980,
|
||||
Height: 700,
|
||||
MinWidth: 820,
|
||||
MinHeight: 560,
|
||||
MinWidth: 980,
|
||||
MinHeight: 700,
|
||||
DisableResize: false,
|
||||
Frameless: goruntime.GOOS == "windows",
|
||||
URL: "/#/model-config",
|
||||
@@ -144,97 +136,6 @@ func (s *WindowService) OpenModelConfigWindow() {
|
||||
s.modelConfigWindow = win
|
||||
}
|
||||
|
||||
// OpenModelEditorWindow 打开模型编辑器独立窗口。
|
||||
// index < 0 表示新增,>= 0 表示编辑对应索引的适配器。
|
||||
// adapterJSON 为编辑器初始数据的 JSON 字符串。
|
||||
func (s *WindowService) OpenModelEditorWindow(index int, adapterJSON string) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if s.app == nil {
|
||||
return
|
||||
}
|
||||
|
||||
s.editorCtx = &modelEditorContext{
|
||||
Index: index,
|
||||
AdapterJSON: adapterJSON,
|
||||
}
|
||||
|
||||
if s.modelEditorWindow != nil {
|
||||
s.modelEditorWindow.Show()
|
||||
s.modelEditorWindow.Focus()
|
||||
return
|
||||
}
|
||||
|
||||
title := "新增模型配置"
|
||||
if index >= 0 {
|
||||
title = "编辑模型配置"
|
||||
}
|
||||
|
||||
win := s.app.Window.NewWithOptions(application.WebviewWindowOptions{
|
||||
Title: title,
|
||||
Width: 840,
|
||||
Height: 680,
|
||||
MinWidth: 740,
|
||||
MinHeight: 600,
|
||||
DisableResize: false,
|
||||
Frameless: goruntime.GOOS == "windows",
|
||||
URL: fmt.Sprintf("/#/model-editor?index=%d", index),
|
||||
Hidden: false,
|
||||
HideOnEscape: false,
|
||||
MinimiseButtonState: application.ButtonEnabled,
|
||||
MaximiseButtonState: application.ButtonEnabled,
|
||||
CloseButtonState: application.ButtonEnabled,
|
||||
BackgroundColour: application.RGBA{Red: 25, Green: 25, Blue: 25, Alpha: 255},
|
||||
Mac: application.MacWindow{
|
||||
Backdrop: application.MacBackdropLiquidGlass,
|
||||
DisableShadow: false,
|
||||
TitleBar: application.MacTitleBar{
|
||||
AppearsTransparent: true,
|
||||
Hide: false,
|
||||
HideTitle: true,
|
||||
FullSizeContent: true,
|
||||
UseToolbar: false,
|
||||
HideToolbarSeparator: true,
|
||||
},
|
||||
WebviewPreferences: application.MacWebviewPreferences{
|
||||
FullscreenEnabled: u.False,
|
||||
TextInteractionEnabled: u.True,
|
||||
AllowsBackForwardNavigationGestures: u.False,
|
||||
},
|
||||
},
|
||||
Windows: application.WindowsWindow{
|
||||
HiddenOnTaskbar: false,
|
||||
},
|
||||
})
|
||||
|
||||
win.RegisterHook(events.Common.WindowClosing, func(e *application.WindowEvent) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.modelEditorWindow = nil
|
||||
s.editorCtx = nil
|
||||
})
|
||||
|
||||
s.modelEditorWindow = win
|
||||
}
|
||||
|
||||
// GetModelEditorContext 返回当前编辑器窗口的初始化上下文。
|
||||
func (s *WindowService) GetModelEditorContext() map[string]any {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
if s.editorCtx == nil {
|
||||
return map[string]any{
|
||||
"index": -1,
|
||||
"adapterJSON": "{}",
|
||||
}
|
||||
}
|
||||
return map[string]any{
|
||||
"index": s.editorCtx.Index,
|
||||
"adapterJSON": s.editorCtx.AdapterJSON,
|
||||
}
|
||||
}
|
||||
|
||||
// OpenHistoryWindow 用于处理与 OpenHistoryWindow 相关的逻辑。
|
||||
func (s *WindowService) OpenHistoryWindow() {
|
||||
_ = os.MkdirAll(client.ResolveLogsRootPath(), 0o755)
|
||||
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
"fmt"
|
||||
"hash/fnv"
|
||||
"io"
|
||||
"math"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sort"
|
||||
@@ -438,7 +437,6 @@ func (s *ProxyService) TestModelAdapter(adapter serverconfig.ModelAdapterConfig)
|
||||
AdapterID: normalized.ID,
|
||||
RequestHash: requestHash,
|
||||
Status: string(ModelAdapterTestStatusRunning),
|
||||
SummaryText: "测试中...",
|
||||
TestedAt: time.Now().UTC().Format(time.RFC3339Nano),
|
||||
}
|
||||
s.storeAndEmitModelAdapterTestResult(running)
|
||||
@@ -516,7 +514,6 @@ func (s *ProxyService) runModelAdapterTest(adapter serverconfig.ModelAdapterConf
|
||||
TestedAt: time.Now().UTC().Format(time.RFC3339Nano),
|
||||
RawResponse: strings.TrimSpace(metrics.rawResponse),
|
||||
}
|
||||
result.SummaryText = buildModelAdapterTestSummaryText(result)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
@@ -747,13 +744,6 @@ func buildErroredModelAdapterTestResult(adapterID string, requestHash string, er
|
||||
}
|
||||
}
|
||||
|
||||
func buildModelAdapterTestSummaryText(result ModelAdapterTestResult) string {
|
||||
if strings.TrimSpace(result.Status) != string(ModelAdapterTestStatusSuccess) {
|
||||
return firstNonEmptyTrimmed(result.SummaryText, "测试失败")
|
||||
}
|
||||
return fmt.Sprintf("%d t/s | 首字 %s", int(math.Round(maxFloat64(result.TokensPerSecond, 0))), formatModelAdapterTestDuration(result.FirstTextTokenMS))
|
||||
}
|
||||
|
||||
func buildModelAdapterHTTPStatusError(prefix string, resp *http.Response) error {
|
||||
if resp == nil {
|
||||
return fmt.Errorf("%s response is nil", strings.TrimSpace(prefix))
|
||||
@@ -849,17 +839,6 @@ func buildModelAdapterTestErrorSummary(err error) string {
|
||||
}
|
||||
}
|
||||
|
||||
func formatModelAdapterTestDuration(durationMS int64) string {
|
||||
if durationMS < 1000 {
|
||||
if durationMS < 0 {
|
||||
durationMS = 0
|
||||
}
|
||||
return fmt.Sprintf("%d ms", durationMS)
|
||||
}
|
||||
seconds := float64(durationMS) / 1000
|
||||
return fmt.Sprintf("%.1f s", seconds)
|
||||
}
|
||||
|
||||
func estimateBenchmarkTextTokens(text string) int64 {
|
||||
trimmed := strings.TrimSpace(text)
|
||||
if trimmed == "" {
|
||||
@@ -971,6 +950,8 @@ func normalizeModelAdapterTestType(value string) string {
|
||||
|
||||
func normalizeModelAdapterTestReasoning(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "":
|
||||
return ""
|
||||
case "low", "medium", "high", "xhigh", "max":
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
default:
|
||||
@@ -1056,13 +1037,6 @@ func normalizeModelAdapterTestInt(value int) int {
|
||||
return value
|
||||
}
|
||||
|
||||
func maxFloat64(value float64, fallback float64) float64 {
|
||||
if value < fallback {
|
||||
return fallback
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func firstNonEmptyTrimmed(values ...string) string {
|
||||
for _, value := range values {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
serverconfig "cursor/internal/backend/server/config"
|
||||
)
|
||||
|
||||
func TestNormalizeModelAdapterTestProviderReasoningPreservesBlank(t *testing.T) {
|
||||
adapter := serverconfig.ModelAdapterConfig{Type: "openai", ReasoningEffort: ""}
|
||||
|
||||
if got := normalizeModelAdapterTestProviderReasoning(adapter); got != "" {
|
||||
t.Fatalf("reasoning effort = %q, want blank", got)
|
||||
}
|
||||
}
|
||||
@@ -30,6 +30,7 @@ const (
|
||||
var cursorStateDisabledStatsigGates = []string{
|
||||
"decompose_always_local_ext_host",
|
||||
"cursor_extensions_isolation_v2",
|
||||
"disable_terminal_output_ui_streaming",
|
||||
}
|
||||
|
||||
// InjectCursorUserInfo synchronizes the Cursor user-level auth cache used by the
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
package cursor
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSyncCursorAuthStateDBDisablesCachedTerminalOutputUIStreamingIdempotently(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "state.vscdb")
|
||||
db, err := sql.Open("sqlite", path)
|
||||
if err != nil {
|
||||
t.Fatalf("open temporary state db: %v", err)
|
||||
}
|
||||
if _, err := db.Exec("CREATE TABLE ItemTable (key TEXT UNIQUE ON CONFLICT REPLACE, value BLOB)"); err != nil {
|
||||
db.Close()
|
||||
t.Fatalf("create ItemTable: %v", err)
|
||||
}
|
||||
bootstrap := map[string]any{
|
||||
"feature_gates": map[string]any{
|
||||
"disable_terminal_output_ui_streaming": map[string]any{
|
||||
"value": true,
|
||||
"rule_id": "local_enabled",
|
||||
"groupName": "local_enabled",
|
||||
},
|
||||
"unrelated_gate": map[string]any{"value": true},
|
||||
},
|
||||
"hash_used": "none",
|
||||
}
|
||||
raw, err := json.Marshal(bootstrap)
|
||||
if err != nil {
|
||||
db.Close()
|
||||
t.Fatalf("encode bootstrap: %v", err)
|
||||
}
|
||||
if _, err := db.Exec("INSERT INTO ItemTable(key, value) VALUES(?, ?)", cursorStateStatsigBootstrapKey, raw); err != nil {
|
||||
db.Close()
|
||||
t.Fatalf("insert bootstrap: %v", err)
|
||||
}
|
||||
if err := db.Close(); err != nil {
|
||||
t.Fatalf("close setup db: %v", err)
|
||||
}
|
||||
|
||||
values := map[string]string{"cursorAuth/cachedEmail": "local@example.com"}
|
||||
if err := syncCursorAuthStateDB(path, values); err != nil {
|
||||
t.Fatalf("first state sync: %v", err)
|
||||
}
|
||||
first := readCursorStatsigBootstrapForTest(t, path)
|
||||
assertCursorStatsigGateValueForTest(t, first, "disable_terminal_output_ui_streaming", false)
|
||||
assertCursorStatsigGateValueForTest(t, first, "unrelated_gate", true)
|
||||
|
||||
if err := syncCursorAuthStateDB(path, values); err != nil {
|
||||
t.Fatalf("second state sync: %v", err)
|
||||
}
|
||||
second := readCursorStatsigBootstrapForTest(t, path)
|
||||
if string(second) != string(first) {
|
||||
t.Fatalf("repeated sync changed bootstrap:\nfirst: %s\nsecond: %s", first, second)
|
||||
}
|
||||
}
|
||||
|
||||
func readCursorStatsigBootstrapForTest(t *testing.T, path string) []byte {
|
||||
t.Helper()
|
||||
db, err := sql.Open("sqlite", path)
|
||||
if err != nil {
|
||||
t.Fatalf("open state db: %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
var raw []byte
|
||||
if err := db.QueryRowContext(context.Background(), "SELECT value FROM ItemTable WHERE key = ?", cursorStateStatsigBootstrapKey).Scan(&raw); err != nil {
|
||||
t.Fatalf("read bootstrap: %v", err)
|
||||
}
|
||||
return raw
|
||||
}
|
||||
|
||||
func assertCursorStatsigGateValueForTest(t *testing.T, raw []byte, name string, want bool) {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
FeatureGates map[string]struct {
|
||||
Value bool `json:"value"`
|
||||
} `json:"feature_gates"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &payload); err != nil {
|
||||
t.Fatalf("decode bootstrap: %v", err)
|
||||
}
|
||||
gate, ok := payload.FeatureGates[name]
|
||||
if !ok {
|
||||
t.Fatalf("missing gate %q", name)
|
||||
}
|
||||
if gate.Value != want {
|
||||
t.Fatalf("gate %q value=%t, want %t", name, gate.Value, want)
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -36,6 +37,8 @@ const (
|
||||
// ModelAdapterConfig 定义了当前模块中的 ModelAdapterConfig 类型。
|
||||
type ModelAdapterConfig struct {
|
||||
ID string `json:"id,omitempty"`
|
||||
// Sort 表示模型渠道的展示顺序。
|
||||
Sort int `json:"sort"`
|
||||
// DisplayName 表示当前声明中的 DisplayName。
|
||||
DisplayName string `json:"displayName"`
|
||||
// Type 表示当前声明中的 Type。
|
||||
@@ -103,6 +106,7 @@ func NormalizeModelAdapterConfigs(input []ModelAdapterConfig) ([]ModelAdapterCon
|
||||
return nil, err
|
||||
}
|
||||
next := ModelAdapterConfig{
|
||||
Sort: item.Sort,
|
||||
DisplayName: strings.TrimSpace(item.DisplayName),
|
||||
Type: normalizeModelAdapterType(item.Type),
|
||||
BaseURL: baseURL,
|
||||
@@ -137,8 +141,8 @@ func NormalizeModelAdapterConfigs(input []ModelAdapterConfig) ([]ModelAdapterCon
|
||||
return nil, errors.New("模型适配器 tooltipData 不能为空")
|
||||
case next.ModelID == "":
|
||||
return nil, errors.New("模型适配器 modelID 不能为空")
|
||||
case next.Type == "openai" && next.ReasoningEffort == "":
|
||||
return nil, errors.New("模型适配器 reasoningEffort 仅支持 low、medium、high、xhigh、max")
|
||||
case next.Type == "openai" && !isSupportedReasoningEffort(next.ReasoningEffort):
|
||||
return nil, errors.New("模型适配器 reasoningEffort 仅支持空值、low、medium、high、xhigh、max")
|
||||
case next.Type == "openai" && next.OpenAIEndpoint == "":
|
||||
return nil, errors.New("模型适配器 openAIEndpoint 仅支持 /v1/responses 或 /v1/chat/completions")
|
||||
case next.Type == "openai" && next.OpenAIExtraParamsEnabled:
|
||||
@@ -163,9 +167,30 @@ func NormalizeModelAdapterConfigs(input []ModelAdapterConfig) ([]ModelAdapterCon
|
||||
seenChannelIDs[next.ID] = struct{}{}
|
||||
normalized = append(normalized, next)
|
||||
}
|
||||
normalizeModelAdapterSorts(normalized)
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeModelAdapterSorts(adapters []ModelAdapterConfig) {
|
||||
sort.SliceStable(adapters, func(leftIndex, rightIndex int) bool {
|
||||
left := adapters[leftIndex].Sort
|
||||
right := adapters[rightIndex].Sort
|
||||
switch {
|
||||
case left <= 0 && right <= 0:
|
||||
return false
|
||||
case left <= 0:
|
||||
return false
|
||||
case right <= 0:
|
||||
return true
|
||||
default:
|
||||
return left < right
|
||||
}
|
||||
})
|
||||
for index := range adapters {
|
||||
adapters[index].Sort = index + 1
|
||||
}
|
||||
}
|
||||
|
||||
func validateJSONMap(value string, fieldName string) error {
|
||||
text := strings.TrimSpace(value)
|
||||
if text == "" {
|
||||
@@ -199,13 +224,15 @@ func validateHeadersJSON(value string) error {
|
||||
}
|
||||
|
||||
func normalizeReasoningEffort(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "", "medium":
|
||||
return "medium"
|
||||
case "low", "high", "xhigh", "max":
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
}
|
||||
|
||||
func isSupportedReasoningEffort(value string) bool {
|
||||
switch value {
|
||||
case "", "low", "medium", "high", "xhigh", "max":
|
||||
return true
|
||||
default:
|
||||
return ""
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
package runtime
|
||||
|
||||
import "testing"
|
||||
|
||||
func testRuntimeModelAdapter(reasoningEffort string) ModelAdapterConfig {
|
||||
return ModelAdapterConfig{
|
||||
DisplayName: "non-reasoning-model",
|
||||
Type: "openai",
|
||||
BaseURL: "https://api.example.com/v1",
|
||||
APIKey: "test-key",
|
||||
TooltipData: "non-reasoning-model",
|
||||
ModelID: "non-reasoning-model",
|
||||
ReasoningEffort: reasoningEffort,
|
||||
OpenAIEndpoint: "/v1/responses",
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeModelAdapterConfigsAllowsBlankReasoningEffort(t *testing.T) {
|
||||
adapters, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{testRuntimeModelAdapter("")})
|
||||
if err != nil {
|
||||
t.Fatalf("NormalizeModelAdapterConfigs returned error: %v", err)
|
||||
}
|
||||
if got := adapters[0].ReasoningEffort; got != "" {
|
||||
t.Fatalf("ReasoningEffort = %q, want blank", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeModelAdapterConfigsRejectsUnknownReasoningEffort(t *testing.T) {
|
||||
if _, err := NormalizeModelAdapterConfigs([]ModelAdapterConfig{testRuntimeModelAdapter("unsupported")}); err == nil {
|
||||
t.Fatal("NormalizeModelAdapterConfigs should reject an unknown reasoning effort")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user