From 9eb24bb4d423802aebdddb9962894383d5d80bae Mon Sep 17 00:00:00 2001 From: leookun Date: Fri, 7 Aug 2026 22:39:58 +0800 Subject: [PATCH] feat: add shell tool call delta message handling Implemented the buildShellToolCallDeltaMessage function to map client shell output to the delta format for Cursor's terminal bubble. This includes handling both stdout and stderr events. Updated the service to publish these messages when processing execution results. Additionally, added a test for enabling terminal output UI streaming in the bootstrap statsig configuration. --- internal/backend/forwarder/events.go | 35 ++++ internal/backend/forwarder/service.go | 5 + .../forwarder/shell_stream_delta_test.go | 150 ++++++++++++++++++ internal/backend/server/upstream/mocks.go | 2 +- .../backend/server/upstream/mocks_test.go | 23 +++ internal/cursor/state_db.go | 1 + internal/cursor/state_db_test.go | 94 +++++++++++ 7 files changed, 309 insertions(+), 1 deletion(-) create mode 100644 internal/backend/forwarder/shell_stream_delta_test.go create mode 100644 internal/cursor/state_db_test.go diff --git a/internal/backend/forwarder/events.go b/internal/backend/forwarder/events.go index ef69776..f080a49 100644 --- a/internal/backend/forwarder/events.go +++ b/internal/backend/forwarder/events.go @@ -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 diff --git a/internal/backend/forwarder/service.go b/internal/backend/forwarder/service.go index 066acfe..3f2e223 100644 --- a/internal/backend/forwarder/service.go +++ b/internal/backend/forwarder/service.go @@ -1003,6 +1003,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 diff --git a/internal/backend/forwarder/shell_stream_delta_test.go b/internal/backend/forwarder/shell_stream_delta_test.go new file mode 100644 index 0000000..0d96102 --- /dev/null +++ b/internal/backend/forwarder/shell_stream_delta_test.go @@ -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) + } + }) + } +} diff --git a/internal/backend/server/upstream/mocks.go b/internal/backend/server/upstream/mocks.go index cc61898..7501a8e 100644 --- a/internal/backend/server/upstream/mocks.go +++ b/internal/backend/server/upstream/mocks.go @@ -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), diff --git a/internal/backend/server/upstream/mocks_test.go b/internal/backend/server/upstream/mocks_test.go index 48c4fa6..5d3fc6a 100644 --- a/internal/backend/server/upstream/mocks_test.go +++ b/internal/backend/server/upstream/mocks_test.go @@ -73,3 +73,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) + } +} diff --git a/internal/cursor/state_db.go b/internal/cursor/state_db.go index e7bc0c6..940c0f1 100644 --- a/internal/cursor/state_db.go +++ b/internal/cursor/state_db.go @@ -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 diff --git a/internal/cursor/state_db_test.go b/internal/cursor/state_db_test.go new file mode 100644 index 0000000..040b917 --- /dev/null +++ b/internal/cursor/state_db_test.go @@ -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) + } +}