From a779320d873ddf57e0d677d5ecaf9c0e75b20640 Mon Sep 17 00:00:00 2001 From: Ethan <39577870+ethanndickson@users.noreply.github.com> Date: Tue, 4 Aug 2026 16:24:14 +1000 Subject: [PATCH] feat(scaletest): llm-mock tool calls and paced streaming (#26850) `coder exp scaletest llm-mock` used to return only canned text, so Coder Agents scaletests pointed at it never exercised the path that matters most: the agentic tool-call loop where the model asks for a tool, the workspace runs it, the result is fed back, and the model is re-prompted before it finally answers. This teaches the mock to reproduce that loop deterministically so scaletests actually drive real tool execution and hold streams open the way a live model would. It adds `--tool-calls-per-turn` and `--tool-call-command` so the OpenAI Chat Completions endpoint emits a controllable number of `execute` tool calls per turn (only when the request advertises an `execute` tool, otherwise it falls back to text, so it stays a safe drop-in). It also adds paced streaming with `--min-stream-duration`/`--max-stream-duration` (randomized per response) and `--response-payload-size`, so runs can simulate slow or long-lived responses instead of flushing everything at once. Closes CODAGT-307 Closes GRU-48 --- cli/exp_scaletest_llmmock.go | 45 +- scaletest/llmmock/openai_stream.go | 162 ++++++ .../llmmock/openai_stream_internal_test.go | 178 ++++++ scaletest/llmmock/openai_toolcalls.go | 64 +++ .../llmmock/openai_toolcalls_internal_test.go | 214 +++++++ scaletest/llmmock/server.go | 542 ++++++++++-------- 6 files changed, 971 insertions(+), 234 deletions(-) create mode 100644 scaletest/llmmock/openai_stream.go create mode 100644 scaletest/llmmock/openai_stream_internal_test.go create mode 100644 scaletest/llmmock/openai_toolcalls.go create mode 100644 scaletest/llmmock/openai_toolcalls_internal_test.go diff --git a/cli/exp_scaletest_llmmock.go b/cli/exp_scaletest_llmmock.go index fa61b8e378..ff8fd75aff 100644 --- a/cli/exp_scaletest_llmmock.go +++ b/cli/exp_scaletest_llmmock.go @@ -19,7 +19,11 @@ func (*RootCmd) scaletestLLMMock() *serpent.Command { var ( address string artificialLatency time.Duration + minStreamDuration time.Duration + maxStreamDuration time.Duration responsePayloadSize int64 + toolCallsPerTurn int64 + toolCallCommand string pprofEnable bool pprofAddress string @@ -34,6 +38,13 @@ func (*RootCmd) scaletestLLMMock() *serpent.Command { ctx, stop := signal.NotifyContext(inv.Context(), StopSignals...) defer stop() + if (minStreamDuration > 0) != (maxStreamDuration > 0) { + return xerrors.New("--min-stream-duration and --max-stream-duration must be set together") + } + if minStreamDuration > maxStreamDuration { + return xerrors.New("--min-stream-duration must not exceed --max-stream-duration") + } + logger := slog.Make(sloghuman.Sink(inv.Stderr)).Leveled(slog.LevelInfo) if pprofEnable { @@ -46,9 +57,11 @@ func (*RootCmd) scaletestLLMMock() *serpent.Command { Address: address, Logger: logger, ArtificialLatency: artificialLatency, + MinStreamDuration: minStreamDuration, + MaxStreamDuration: maxStreamDuration, ResponsePayloadSize: int(responsePayloadSize), - PprofEnable: pprofEnable, - PprofAddress: pprofAddress, + ToolCallsPerTurn: int(toolCallsPerTurn), + ToolCallCommand: toolCallCommand, TraceEnable: traceEnable, } srv := new(llmmock.Server) @@ -87,6 +100,20 @@ func (*RootCmd) scaletestLLMMock() *serpent.Command { Description: "Artificial latency to add to each response (e.g., 100ms, 1s). Simulates slow upstream processing.", Value: serpent.DurationOf(&artificialLatency), }, + { + Flag: "min-stream-duration", + Env: "CODER_SCALETEST_LLM_MOCK_MIN_STREAM_DURATION", + Default: "0s", + Description: "Minimum duration to stream a text response over (e.g., 5s, 10s). Set with max-stream-duration to pace response chunks.", + Value: serpent.DurationOf(&minStreamDuration), + }, + { + Flag: "max-stream-duration", + Env: "CODER_SCALETEST_LLM_MOCK_MAX_STREAM_DURATION", + Default: "0s", + Description: "Maximum duration to stream a text response over (e.g., 10s, 30s). Set with min-stream-duration to pace response chunks.", + Value: serpent.DurationOf(&maxStreamDuration), + }, { Flag: "response-payload-size", Env: "CODER_SCALETEST_LLM_MOCK_RESPONSE_PAYLOAD_SIZE", @@ -94,6 +121,20 @@ func (*RootCmd) scaletestLLMMock() *serpent.Command { Description: "Size in bytes of the response payload. If 0, uses default context-aware responses.", Value: serpent.Int64Of(&responsePayloadSize), }, + { + Flag: "tool-calls-per-turn", + Env: "CODER_SCALETEST_LLM_MOCK_TOOL_CALLS_PER_TURN", + Default: "0", + Description: "Number of execute tool calls to emit per user turn. Set to 0 for text-only responses. OpenAI Chat Completions only.", + Value: serpent.Int64Of(&toolCallsPerTurn), + }, + { + Flag: "tool-call-command", + Env: "CODER_SCALETEST_LLM_MOCK_TOOL_CALL_COMMAND", + Default: "echo scaletest", + Description: "Shell command sent in each mock execute tool call when tool calls are enabled. OpenAI Chat Completions only.", + Value: serpent.StringOf(&toolCallCommand), + }, { Flag: "pprof-enable", Env: "CODER_SCALETEST_LLM_MOCK_PPROF_ENABLE", diff --git a/scaletest/llmmock/openai_stream.go b/scaletest/llmmock/openai_stream.go new file mode 100644 index 0000000000..840d127f6b --- /dev/null +++ b/scaletest/llmmock/openai_stream.go @@ -0,0 +1,162 @@ +package llmmock + +import ( + "context" + "encoding/json" + "io" + "net/http" + + "cdr.dev/slog/v3" +) + +type openAIToolCallDelta struct { + Index int `json:"index"` + ID string `json:"id,omitempty"` + Type string `json:"type,omitempty"` + Function openAIToolCallFunction `json:"function"` +} + +type openAIStreamDelta struct { + Role string `json:"role,omitempty"` + Content *string `json:"content,omitempty"` + ToolCalls []openAIToolCallDelta `json:"tool_calls,omitempty"` +} + +type openAIStreamChoice struct { + Index int `json:"index"` + Delta openAIStreamDelta `json:"delta"` + FinishReason *string `json:"finish_reason"` +} + +type openAIStreamChunk struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []openAIStreamChoice `json:"choices"` +} + +type openAIStreamWriter struct { + logger slog.Logger + w http.ResponseWriter + flusher http.Flusher + response openAIResponse +} + +func (s *Server) newOpenAIStreamWriter(ctx context.Context, w http.ResponseWriter, resp openAIResponse) (openAIStreamWriter, bool) { + flusher, ok := w.(http.Flusher) + if !ok { + s.logger.Error(ctx, "responseWriter does not support flushing", + slog.F("response_id", resp.ID), + ) + return openAIStreamWriter{}, false + } + + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.WriteHeader(http.StatusOK) + + return openAIStreamWriter{ + logger: s.logger, + w: w, + flusher: flusher, + response: resp, + }, true +} + +func (sw openAIStreamWriter) writeDelta(ctx context.Context, delta openAIStreamDelta, finishReason *string) bool { + chunk := openAIStreamChunk{ + ID: sw.response.ID, + Object: "chat.completion.chunk", + Created: sw.response.Created, + Model: sw.response.Model, + Choices: []openAIStreamChoice{{ + Index: 0, + Delta: delta, + FinishReason: finishReason, + }}, + } + data, _ := json.Marshal(chunk) + return sw.writeData(ctx, data) +} + +func (sw openAIStreamWriter) writeDone(ctx context.Context) bool { + return sw.writeData(ctx, []byte("[DONE]")) +} + +func (sw openAIStreamWriter) writeData(ctx context.Context, data []byte) bool { + if _, err := io.WriteString(sw.w, "data: "); err != nil { + return sw.logWriteError(ctx, err) + } + if _, err := sw.w.Write(data); err != nil { + return sw.logWriteError(ctx, err) + } + if _, err := io.WriteString(sw.w, "\n\n"); err != nil { + return sw.logWriteError(ctx, err) + } + sw.flusher.Flush() + return true +} + +func (sw openAIStreamWriter) logWriteError(ctx context.Context, err error) bool { + sw.logger.Error(ctx, "failed to write OpenAI stream chunk", + slog.F("response_id", sw.response.ID), + slog.Error(err), + slog.F("error_type", "write_error"), + slog.F("likely_cause", "network_error"), + ) + return false +} + +func (s *Server) sendOpenAIStream(ctx context.Context, w http.ResponseWriter, resp openAIResponse) { + writer, ok := s.newOpenAIStreamWriter(ctx, w, resp) + if !ok { + return + } + + choice := resp.Choices[0] + if len(choice.Message.ToolCalls) > 0 { + if !writer.writeDelta(ctx, openAIStreamDelta{ + Role: "assistant", + ToolCalls: openAIStreamToolCallDeltas(choice.Message.ToolCalls), + }, nil) { + return + } + } else if !s.writeOpenAITextStream(ctx, writer, choice.Message.Content) { + return + } + + if !writer.writeDelta(ctx, openAIStreamDelta{}, &choice.FinishReason) { + return + } + _ = writer.writeDone(ctx) +} + +func (s *Server) writeOpenAITextStream(ctx context.Context, writer openAIStreamWriter, content string) bool { + first := true + for chunk := range s.streamContentChunks(ctx, s.randomStreamDuration(), content) { + delta := openAIStreamDelta{Content: &chunk} + if first { + delta.Role = "assistant" + first = false + } + if !writer.writeDelta(ctx, delta, nil) { + return false + } + } + return ctx.Err() == nil +} + +func openAIStreamToolCallDeltas(toolCalls []openAIToolCall) []openAIToolCallDelta { + deltas := make([]openAIToolCallDelta, 0, len(toolCalls)) + for i, toolCall := range toolCalls { + deltas = append(deltas, openAIToolCallDelta{ + Index: i, + ID: toolCall.ID, + Type: toolCall.Type, + Function: toolCall.Function, + }) + } + return deltas +} diff --git a/scaletest/llmmock/openai_stream_internal_test.go b/scaletest/llmmock/openai_stream_internal_test.go new file mode 100644 index 0000000000..c4c0919138 --- /dev/null +++ b/scaletest/llmmock/openai_stream_internal_test.go @@ -0,0 +1,178 @@ +package llmmock + +import ( + "context" + "net/http/httptest" + "slices" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/testutil" +) + +func TestStreamWordChunks(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + content string + want []string + }{ + { + name: "words", + content: "hello world test", + want: []string{"hello ", "world ", "test"}, + }, + { + name: "trailing space", + content: "hello world ", + want: []string{"hello ", "world "}, + }, + { + name: "single word", + content: "single", + want: []string{"single"}, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + require.Equal(t, tc.want, streamWordChunks(tc.content)) + }) + } +} + +func TestStreamContentChunksWordSplit(t *testing.T) { + t.Parallel() + + srv := &Server{} + ctx := testutil.Context(t, testutil.WaitShort) + got := slices.Collect(srv.streamContentChunks(ctx, time.Nanosecond, "hello world test")) + require.Equal(t, []string{"hello ", "world ", "test"}, got) +} + +func TestStreamContentChunksFixedWindow(t *testing.T) { + t.Parallel() + + srv := &Server{responsePayloadSize: 1} + ctx := testutil.Context(t, testutil.WaitShort) + content := strings.Repeat("x", 2050) + + got := slices.Collect(srv.streamContentChunks(ctx, time.Nanosecond, content)) + require.Len(t, got, 3) + require.Equal(t, streamFixedWindowSize, len(got[0])) + require.Equal(t, streamFixedWindowSize, len(got[1])) + require.Equal(t, 2, len(got[2])) + require.Equal(t, content, strings.Join(got, "")) +} + +func TestStreamContentChunksShortCircuits(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + totalDuration time.Duration + content string + want []string + }{ + { + name: "non-positive duration", + totalDuration: 0, + content: "anything", + want: []string{"anything"}, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + srv := &Server{} + ctx := testutil.Context(t, testutil.WaitShort) + got := slices.Collect(srv.streamContentChunks(ctx, tc.totalDuration, tc.content)) + require.Equal(t, tc.want, got) + }) + } +} + +func TestStreamPacedChunksStopsOnCancel(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort)) + defer cancel() + + var got []string + for chunk := range streamPacedChunks(ctx, time.Hour, 3, slices.Values([]string{"a", "b", "c"})) { + got = append(got, chunk) + cancel() + } + require.Equal(t, []string{"a"}, got) +} + +func TestRandomStreamDuration(t *testing.T) { + t.Parallel() + + require.Zero(t, (&Server{}).randomStreamDuration()) + require.Equal(t, 5*time.Second, (&Server{ + minStreamDuration: 5 * time.Second, + maxStreamDuration: 5 * time.Second, + }).randomStreamDuration()) + + duration := (&Server{ + minStreamDuration: time.Second, + maxStreamDuration: 2 * time.Second, + }).randomStreamDuration() + require.GreaterOrEqual(t, duration, time.Second) + require.Less(t, duration, 2*time.Second) +} + +func TestSendOpenAIStreamTextContent(t *testing.T) { + t.Parallel() + + // Nanosecond pacing keeps the test instant while forcing the multi-chunk + // path so only the first delta should carry the assistant role. + srv := &Server{ + minStreamDuration: time.Nanosecond, + maxStreamDuration: time.Nanosecond, + } + writer := httptest.NewRecorder() + resp := openAIResponse{ + ID: "chatcmpl-text", + Object: "chat.completion", + Created: 7, + Model: "scaletest-model", + Choices: []openAIResponseChoice{{ + Message: openAIMessage{Role: "assistant", Content: "hello world test"}, + FinishReason: openAIStopFinishReason, + }}, + } + + ctx := testutil.Context(t, testutil.WaitShort) + srv.sendOpenAIStream(ctx, writer, resp) + events := sseDataEvents(t, writer.Body.String()) + require.Len(t, events, 5) // 3 word chunks + finish + [DONE] + + first := decodeStreamChunk(t, events[0]) + require.Len(t, first.Choices, 1) + require.Nil(t, first.Choices[0].FinishReason) + require.Equal(t, "assistant", first.Choices[0].Delta.Role) + require.NotNil(t, first.Choices[0].Delta.Content) + require.Equal(t, "hello ", *first.Choices[0].Delta.Content) + + second := decodeStreamChunk(t, events[1]) + require.Len(t, second.Choices, 1) + require.Nil(t, second.Choices[0].FinishReason) + require.Empty(t, second.Choices[0].Delta.Role) + require.NotNil(t, second.Choices[0].Delta.Content) + require.Equal(t, "world ", *second.Choices[0].Delta.Content) + + finish := decodeStreamChunk(t, events[3]) + require.Len(t, finish.Choices, 1) + require.NotNil(t, finish.Choices[0].FinishReason) + require.Equal(t, openAIStopFinishReason, *finish.Choices[0].FinishReason) + require.Empty(t, finish.Choices[0].Delta.Role) + require.Nil(t, finish.Choices[0].Delta.Content) + + require.Equal(t, "[DONE]", events[4]) +} diff --git a/scaletest/llmmock/openai_toolcalls.go b/scaletest/llmmock/openai_toolcalls.go new file mode 100644 index 0000000000..071ddfc777 --- /dev/null +++ b/scaletest/llmmock/openai_toolcalls.go @@ -0,0 +1,64 @@ +package llmmock + +import ( + "encoding/json" + "fmt" + "slices" + + "github.com/google/uuid" +) + +const executeToolName = "execute" + +func (s *Server) buildOpenAIChoice(req llmRequest) openAIResponseChoice { + executeToolIncluded := slices.ContainsFunc(req.Tools, func(tool openAITool) bool { + return tool.Function.Name == executeToolName + }) + + if s.needsOpenAIToolCall(req) && executeToolIncluded { + return openAIResponseChoice{ + Message: openAIMessage{ + Role: "assistant", + ToolCalls: []openAIToolCall{executeToolCall(s.toolCallCommand)}, + }, + FinishReason: openAIToolCallFinishReason, + } + } + + return openAIResponseChoice{ + Message: openAIMessage{ + Role: "assistant", + Content: s.responseText(openAIDefaultResponseText), + }, + FinishReason: openAIStopFinishReason, + } +} + +func (s *Server) needsOpenAIToolCall(req llmRequest) bool { + if s.toolCallsPerTurn <= 0 { + return false + } + + completedToolCalls := 0 + for _, msg := range slices.Backward(req.Messages) { + switch msg.Role { + case "tool": + completedToolCalls++ + case "user": + return completedToolCalls < s.toolCallsPerTurn + } + } + return false +} + +func executeToolCall(command string) openAIToolCall { + payload, _ := json.Marshal(map[string]string{"command": command}) + return openAIToolCall{ + ID: fmt.Sprintf("call_%s", uuid.New().String()[:8]), + Type: "function", + Function: openAIToolCallFunction{ + Name: executeToolName, + Arguments: string(payload), + }, + } +} diff --git a/scaletest/llmmock/openai_toolcalls_internal_test.go b/scaletest/llmmock/openai_toolcalls_internal_test.go new file mode 100644 index 0000000000..187aa936a1 --- /dev/null +++ b/scaletest/llmmock/openai_toolcalls_internal_test.go @@ -0,0 +1,214 @@ +package llmmock + +import ( + "encoding/json" + "net/http/httptest" + "strings" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/testutil" +) + +const testToolCallCommand = "echo scaletest" + +func TestBuildOpenAIChoice(t *testing.T) { + t.Parallel() + + t.Run("TextOnlyDoesNotRequireTools", func(t *testing.T) { + t.Parallel() + + srv := &Server{toolCallsPerTurn: 0} + choice := srv.buildOpenAIChoice(openAIExecuteRequest(0)) + require.Equal(t, openAIStopFinishReason, choice.FinishReason) + require.Empty(t, choice.Message.ToolCalls) + require.Equal(t, openAIDefaultResponseText, choice.Message.Content) + }) + + t.Run("NoUserMessageReturnsText", func(t *testing.T) { + t.Parallel() + + srv := &Server{toolCallsPerTurn: 1, toolCallCommand: testToolCallCommand} + choice := srv.buildOpenAIChoice(llmRequest{ + Model: "scaletest-model", + Messages: []llmRequestMessage{{Role: "system"}}, + }) + require.Equal(t, openAIStopFinishReason, choice.FinishReason) + require.Empty(t, choice.Message.ToolCalls) + require.Equal(t, openAIDefaultResponseText, choice.Message.Content) + }) + + t.Run("EmitsToolCallNoneCompleted", func(t *testing.T) { + t.Parallel() + + srv := &Server{toolCallsPerTurn: 2, toolCallCommand: testToolCallCommand} + choice := srv.buildOpenAIChoice(openAIExecuteRequest(0)) + require.Equal(t, openAIToolCallFinishReason, choice.FinishReason) + require.Len(t, choice.Message.ToolCalls, 1) + toolCall := choice.Message.ToolCalls[0] + requireExecuteToolCall(t, toolCall, testToolCallCommand) + + toolCallJSON, err := json.Marshal(toolCall) + require.NoError(t, err) + var toolCallFields map[string]json.RawMessage + require.NoError(t, json.Unmarshal(toolCallJSON, &toolCallFields)) + require.NotContains(t, toolCallFields, "index") + }) + + t.Run("EmitsToolCallOneCompleted", func(t *testing.T) { + t.Parallel() + + srv := &Server{toolCallsPerTurn: 2, toolCallCommand: testToolCallCommand} + choice := srv.buildOpenAIChoice(openAIExecuteRequest(1)) + require.Equal(t, openAIToolCallFinishReason, choice.FinishReason) + require.Len(t, choice.Message.ToolCalls, 1) + requireExecuteToolCall(t, choice.Message.ToolCalls[0], testToolCallCommand) + }) + + t.Run("StopsAfterAllToolCallsCompleted", func(t *testing.T) { + t.Parallel() + + srv := &Server{toolCallsPerTurn: 2, toolCallCommand: testToolCallCommand} + choice := srv.buildOpenAIChoice(openAIExecuteRequest(2)) + require.Equal(t, openAIStopFinishReason, choice.FinishReason) + require.Empty(t, choice.Message.ToolCalls) + require.Equal(t, openAIDefaultResponseText, choice.Message.Content) + }) + + t.Run("FallsBackToTextWhenExecuteToolMissing", func(t *testing.T) { + t.Parallel() + + srv := &Server{toolCallsPerTurn: 1, toolCallCommand: testToolCallCommand} + req := openAIExecuteRequest(0) + req.Tools = nil + choice := srv.buildOpenAIChoice(req) + require.Equal(t, openAIStopFinishReason, choice.FinishReason) + require.Empty(t, choice.Message.ToolCalls) + require.Equal(t, openAIDefaultResponseText, choice.Message.Content) + }) +} + +func TestSendOpenAIStreamIncludesToolCalls(t *testing.T) { + t.Parallel() + + toolCall := executeToolCall(testToolCallCommand) + + srv := &Server{} + writer := httptest.NewRecorder() + resp := openAIResponse{ + ID: "chatcmpl-test", + Object: "chat.completion", + Created: 7, + Model: "scaletest-model", + Choices: []openAIResponseChoice{{ + Message: openAIMessage{Role: "assistant", ToolCalls: []openAIToolCall{toolCall}}, + FinishReason: openAIToolCallFinishReason, + }}, + } + + ctx := testutil.Context(t, testutil.WaitShort) + srv.sendOpenAIStream(ctx, writer, resp) + events := sseDataEvents(t, writer.Body.String()) + require.Len(t, events, 3) + + first := decodeStreamChunk(t, events[0]) + require.Len(t, first.Choices, 1) + require.Nil(t, first.Choices[0].FinishReason) + require.Equal(t, "assistant", first.Choices[0].Delta.Role) + require.Len(t, first.Choices[0].Delta.ToolCalls, 1) + requireExecuteStreamToolCall(t, events[0], first.Choices[0].Delta.ToolCalls[0], testToolCallCommand) + + second := decodeStreamChunk(t, events[1]) + require.Len(t, second.Choices, 1) + require.NotNil(t, second.Choices[0].FinishReason) + require.Equal(t, openAIToolCallFinishReason, *second.Choices[0].FinishReason) + require.Empty(t, second.Choices[0].Delta.ToolCalls) + + require.Equal(t, "[DONE]", events[2]) +} + +func decodeStreamChunk(t *testing.T, data string) openAIStreamChunk { + t.Helper() + + var chunk openAIStreamChunk + require.NoError(t, json.Unmarshal([]byte(data), &chunk)) + return chunk +} + +func sseDataEvents(t *testing.T, body string) []string { + t.Helper() + + var events []string + for _, event := range strings.Split(body, "\n\n") { + if event == "" { + continue + } + + var dataLines []string + for _, line := range strings.Split(event, "\n") { + data, ok := strings.CutPrefix(line, "data: ") + if ok { + dataLines = append(dataLines, data) + } + } + if len(dataLines) > 0 { + events = append(events, strings.Join(dataLines, "\n")) + } + } + return events +} + +// requireExecuteStreamToolCall asserts the streamed tool-call delta and that +// its JSON payload includes the index field, which the production type marks +// as non-omitempty but is otherwise indistinguishable from a zero default in +// the decoded struct. +func requireExecuteStreamToolCall(t *testing.T, rawEvent string, toolCall openAIToolCallDelta, command string) { + t.Helper() + + require.Zero(t, toolCall.Index) + + var rawChunk struct { + Choices []struct { + Delta struct { + ToolCalls []map[string]json.RawMessage `json:"tool_calls"` + } `json:"delta"` + } `json:"choices"` + } + require.NoError(t, json.Unmarshal([]byte(rawEvent), &rawChunk)) + require.Len(t, rawChunk.Choices, 1) + require.Len(t, rawChunk.Choices[0].Delta.ToolCalls, 1) + require.Contains(t, rawChunk.Choices[0].Delta.ToolCalls[0], "index") + + requireExecuteToolCall(t, openAIToolCall{ + ID: toolCall.ID, + Type: toolCall.Type, + Function: toolCall.Function, + }, command) +} + +func requireExecuteToolCall(t *testing.T, toolCall openAIToolCall, command string) { + t.Helper() + + require.True(t, strings.HasPrefix(toolCall.ID, "call_"), "tool call ID %q", toolCall.ID) + require.Equal(t, "function", toolCall.Type) + require.Equal(t, executeToolName, toolCall.Function.Name) + expectedPayload, err := json.Marshal(map[string]string{"command": command}) + require.NoError(t, err) + require.JSONEq(t, string(expectedPayload), toolCall.Function.Arguments) +} + +func openAIExecuteRequest(completedToolCalls int) llmRequest { + req := llmRequest{ + Model: "scaletest-model", + Messages: []llmRequestMessage{{Role: "user"}}, + Tools: []openAITool{{Type: "function", Function: openAIToolFunction{Name: executeToolName}}}, + } + for range completedToolCalls { + req.Messages = append(req.Messages, + llmRequestMessage{Role: "assistant"}, + llmRequestMessage{Role: "tool"}, + ) + } + return req +} diff --git a/scaletest/llmmock/server.go b/scaletest/llmmock/server.go index 8c9bdfe3c9..9c86cb7779 100644 --- a/scaletest/llmmock/server.go +++ b/scaletest/llmmock/server.go @@ -5,8 +5,11 @@ import ( "encoding/json" "errors" "fmt" + "iter" + "math/rand/v2" "net" "net/http" + "slices" "strings" "time" @@ -24,15 +27,34 @@ import ( "github.com/coder/coder/v2/coderd/tracing" ) +const ( + openAIDefaultResponseText = "This is a mock response from OpenAI." + openAIStopFinishReason = "stop" + openAIToolCallFinishReason = "tool_calls" + openAIResponsesDefaultResponseText = "This is a mock response from OpenAI Responses." + anthropicDefaultResponseText = "This is a mock response from Anthropic." + mockInputTokens = 10 + mockOutputTokens = 5 + // streamFixedWindowSize is the number of bytes per delta for fixed-size + // payload streams (when responsePayloadSize is set). + streamFixedWindowSize = 1024 +) + // Server wraps the LLM mock server and provides an HTTP API to retrieve requests. type Server struct { httpServer *http.Server httpListener net.Listener + httpCancel context.CancelFunc logger slog.Logger address string artificialLatency time.Duration + minStreamDuration time.Duration + maxStreamDuration time.Duration responsePayloadSize int + responsePayload string + toolCallsPerTurn int + toolCallCommand string tracerProvider trace.TracerProvider closeTracing func(context.Context) error @@ -42,35 +64,68 @@ type Config struct { Address string Logger slog.Logger ArtificialLatency time.Duration + MinStreamDuration time.Duration + MaxStreamDuration time.Duration ResponsePayloadSize int - - PprofEnable bool - PprofAddress string + ToolCallsPerTurn int + ToolCallCommand string TraceEnable bool } type llmRequest struct { - Model string `json:"model"` - Stream bool `json:"stream,omitempty"` + Model string `json:"model"` + Messages []llmRequestMessage `json:"messages,omitempty"` + Tools []openAITool `json:"tools,omitempty"` + Stream bool `json:"stream,omitempty"` +} + +// llmRequestMessage decodes only the request message fields the mock +// inspects. Content is intentionally omitted: providers send it as either a +// string or an array of content blocks, and the mock never reads it. +type llmRequestMessage struct { + Role string `json:"role"` } type openAIMessage struct { - Role string `json:"role"` - Content string `json:"content"` + Role string `json:"role"` + Content string `json:"content,omitempty"` + ToolCalls []openAIToolCall `json:"tool_calls,omitempty"` +} + +type openAITool struct { + Type string `json:"type"` + Function openAIToolFunction `json:"function"` +} + +type openAIToolFunction struct { + Name string `json:"name"` +} + +type openAIToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + Function openAIToolCallFunction `json:"function"` +} + +type openAIToolCallFunction struct { + Name string `json:"name,omitempty"` + Arguments string `json:"arguments,omitempty"` +} + +type openAIResponseChoice struct { + Index int `json:"index"` + Message openAIMessage `json:"message"` + FinishReason string `json:"finish_reason"` } type openAIResponse struct { - ID string `json:"id"` - Object string `json:"object"` - Created int64 `json:"created"` - Model string `json:"model"` - Choices []struct { - Index int `json:"index"` - Message openAIMessage `json:"message"` - FinishReason string `json:"finish_reason"` - } `json:"choices"` - Usage struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []openAIResponseChoice `json:"choices"` + Usage struct { PromptTokens int `json:"prompt_tokens"` CompletionTokens int `json:"completion_tokens"` TotalTokens int `json:"total_tokens"` @@ -119,7 +174,15 @@ func (s *Server) Start(ctx context.Context, cfg Config) error { s.address = cfg.Address s.logger = cfg.Logger s.artificialLatency = cfg.ArtificialLatency + s.minStreamDuration = cfg.MinStreamDuration + s.maxStreamDuration = cfg.MaxStreamDuration s.responsePayloadSize = cfg.ResponsePayloadSize + s.responsePayload = "" + if s.responsePayloadSize > 0 { + s.responsePayload = strings.Repeat("x", s.responsePayloadSize) + } + s.toolCallsPerTurn = cfg.ToolCallsPerTurn + s.toolCallCommand = cfg.ToolCallCommand if cfg.TraceEnable { otel.SetTextMapPropagator( @@ -148,6 +211,9 @@ func (s *Server) Start(ctx context.Context, cfg Config) error { } func (s *Server) Stop() error { + if s.httpCancel != nil { + s.httpCancel() + } if s.httpServer != nil { shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() @@ -169,6 +235,98 @@ func (s *Server) APIAddress() string { return fmt.Sprintf("http://%s", s.httpListener.Addr().String()) } +func (s *Server) responseText(fallback string) string { + if s.responsePayloadSize > 0 { + return s.responsePayload + } + return fallback +} + +func (s *Server) randomStreamDuration() time.Duration { + if s.minStreamDuration <= 0 || s.maxStreamDuration <= 0 { + return 0 + } + if s.minStreamDuration >= s.maxStreamDuration { + return s.minStreamDuration + } + + delta := s.maxStreamDuration - s.minStreamDuration + //nolint:gosec // This is a scaletest mock, not security-sensitive. + return s.minStreamDuration + time.Duration(rand.Int64N(int64(delta))) +} + +// streamContentChunks paces content across totalDuration. The fixed-size path +// slices content by byte offset, which is rune-safe only because that payload +// is ASCII (see responseText). +func (s *Server) streamContentChunks(ctx context.Context, totalDuration time.Duration, content string) iter.Seq[string] { + if totalDuration <= 0 || content == "" { + return func(yield func(string) bool) { + yield(content) + } + } + if s.responsePayloadSize > 0 { + n := (len(content) + streamFixedWindowSize - 1) / streamFixedWindowSize + return streamPacedChunks(ctx, totalDuration, n, func(yield func(string) bool) { + for i := range n { + start := i * streamFixedWindowSize + end := min(start+streamFixedWindowSize, len(content)) + if !yield(content[start:end]) { + return + } + } + }) + } + + chunks := streamWordChunks(content) + return streamPacedChunks(ctx, totalDuration, len(chunks), slices.Values(chunks)) +} + +func streamWordChunks(content string) []string { + chunks := strings.SplitAfter(content, " ") + if chunks[len(chunks)-1] == "" { + chunks = chunks[:len(chunks)-1] + } + return chunks +} + +// sleepContext blocks until d elapses or ctx is canceled. It reports whether +// the duration elapsed so callers can stop work when the context is done. +func sleepContext(ctx context.Context, d time.Duration) bool { + if d <= 0 { + return ctx.Err() == nil + } + timer := time.NewTimer(d) + defer timer.Stop() + select { + case <-ctx.Done(): + return false + case <-timer.C: + return true + } +} + +func streamPacedChunks(ctx context.Context, totalDuration time.Duration, n int, chunks iter.Seq[string]) iter.Seq[string] { + return func(yield func(string) bool) { + if n == 0 { + yield("") + return + } + + // Delay after every chunk, including the last, so the stream stays open + // for roughly totalDuration even when there is a single chunk (small + // --response-payload-size or single-word text). + delay := totalDuration / time.Duration(n) + for chunk := range chunks { + if !yield(chunk) { + return + } + if !sleepContext(ctx, delay) { + return + } + } + } +} + func (s *Server) startAPIServer(ctx context.Context) error { mux := http.NewServeMux() @@ -181,9 +339,14 @@ func (s *Server) startAPIServer(ctx context.Context) error { handler = s.tracingMiddleware(handler) } + baseCtx, httpCancel := context.WithCancel(ctx) + s.httpCancel = httpCancel s.httpServer = &http.Server{ Handler: handler, ReadHeaderTimeout: 10 * time.Second, + BaseContext: func(net.Listener) context.Context { + return baseCtx + }, } listener, err := net.Listen("tcp", s.address) @@ -222,59 +385,43 @@ func (s *Server) handleOpenAIWithLabels(w http.ResponseWriter, r *http.Request) return } - if s.artificialLatency > 0 { - time.Sleep(s.artificialLatency) + if !sleepContext(ctx, s.artificialLatency) { + return } - var resp openAIResponse - resp.ID = fmt.Sprintf("chatcmpl-%s", requestID.String()[:8]) - resp.Object = "chat.completion" - resp.Created = now.Unix() - resp.Model = req.Model + choice := s.buildOpenAIChoice(req) - var responseContent string - if s.responsePayloadSize > 0 { - pattern := "x" - repeated := strings.Repeat(pattern, s.responsePayloadSize) - responseContent = repeated[:s.responsePayloadSize] - } else { - responseContent = "This is a mock response from OpenAI." + resp := openAIResponse{ + ID: fmt.Sprintf("chatcmpl-%s", requestID.String()[:8]), + Object: "chat.completion", + Created: now.Unix(), + Model: req.Model, + Choices: []openAIResponseChoice{choice}, } - - resp.Choices = []struct { - Index int `json:"index"` - Message openAIMessage `json:"message"` - FinishReason string `json:"finish_reason"` - }{ - { - Index: 0, - Message: openAIMessage{ - Role: "assistant", - Content: responseContent, - }, - FinishReason: "stop", - }, - } - - resp.Usage.PromptTokens = 10 - resp.Usage.CompletionTokens = 5 - resp.Usage.TotalTokens = 15 - - responseBody, _ := json.Marshal(resp) + resp.Usage.PromptTokens = mockInputTokens + resp.Usage.CompletionTokens = mockOutputTokens + resp.Usage.TotalTokens = mockInputTokens + mockOutputTokens if req.Stream { s.sendOpenAIStream(ctx, w, resp) - } else { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - if _, err := w.Write(responseBody); err != nil { - s.logger.Error(ctx, "failed to write OpenAI response", - slog.F("request_id", requestID), - slog.Error(err), - slog.F("error_type", "write_error"), - slog.F("likely_cause", "network_error"), - ) - } + return + } + + responseBody, err := json.Marshal(resp) + if err != nil { + s.logger.Error(ctx, "failed to marshal OpenAI response", slog.Error(err)) + http.Error(w, "failed to marshal response", http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + if _, err := w.Write(responseBody); err != nil { + s.logger.Error(ctx, "failed to write OpenAI response", + slog.F("request_id", requestID), + slog.Error(err), + slog.F("error_type", "write_error"), + slog.F("likely_cause", "network_error"), + ) } } @@ -305,8 +452,8 @@ func (s *Server) handleResponsesWithLabels(w http.ResponseWriter, r *http.Reques return } - if s.artificialLatency > 0 { - time.Sleep(s.artificialLatency) + if !sleepContext(ctx, s.artificialLatency) { + return } var resp responsesResponse @@ -315,14 +462,7 @@ func (s *Server) handleResponsesWithLabels(w http.ResponseWriter, r *http.Reques resp.Created = now.Unix() resp.Model = req.Model - var responseContent string - if s.responsePayloadSize > 0 { - pattern := "x" - repeated := strings.Repeat(pattern, s.responsePayloadSize) - responseContent = repeated[:s.responsePayloadSize] - } else { - responseContent = "This is a mock response from OpenAI Responses." - } + assistantText := s.responseText(openAIResponsesDefaultResponseText) resp.Output = []struct { ID string `json:"id,omitempty"` @@ -343,31 +483,36 @@ func (s *Server) handleResponsesWithLabels(w http.ResponseWriter, r *http.Reques }{ { Type: "output_text", - Text: responseContent, + Text: assistantText, }, }, }, } - resp.Usage.InputTokens = 10 - resp.Usage.OutputTokens = 5 - resp.Usage.TotalTokens = 15 - - responseBody, _ := json.Marshal(resp) + resp.Usage.InputTokens = mockInputTokens + resp.Usage.OutputTokens = mockOutputTokens + resp.Usage.TotalTokens = mockInputTokens + mockOutputTokens if req.Stream { s.sendResponsesStream(ctx, w, resp) - } else { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - if _, err := w.Write(responseBody); err != nil { - s.logger.Error(ctx, "failed to write OpenAI responses response", - slog.F("request_id", requestID), - slog.Error(err), - slog.F("error_type", "write_error"), - slog.F("likely_cause", "network_error"), - ) - } + return + } + + responseBody, err := json.Marshal(resp) + if err != nil { + s.logger.Error(ctx, "failed to marshal OpenAI responses response", slog.Error(err)) + http.Error(w, "failed to marshal response", http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + if _, err := w.Write(responseBody); err != nil { + s.logger.Error(ctx, "failed to write OpenAI responses response", + slog.F("request_id", requestID), + slog.Error(err), + slog.F("error_type", "write_error"), + slog.F("likely_cause", "network_error"), + ) } } @@ -383,8 +528,8 @@ func (s *Server) handleAnthropicWithLabels(w http.ResponseWriter, r *http.Reques return } - if s.artificialLatency > 0 { - time.Sleep(s.artificialLatency) + if !sleepContext(ctx, s.artificialLatency) { + return } var resp anthropicResponse @@ -392,14 +537,7 @@ func (s *Server) handleAnthropicWithLabels(w http.ResponseWriter, r *http.Reques resp.Type = "message" resp.Role = "assistant" - var responseText string - if s.responsePayloadSize > 0 { - pattern := "x" - repeated := strings.Repeat(pattern, s.responsePayloadSize) - responseText = repeated[:s.responsePayloadSize] - } else { - responseText = "This is a mock response from Anthropic." - } + assistantText := s.responseText(anthropicDefaultResponseText) resp.Content = []struct { Type string `json:"type"` @@ -407,110 +545,39 @@ func (s *Server) handleAnthropicWithLabels(w http.ResponseWriter, r *http.Reques }{ { Type: "text", - Text: responseText, + Text: assistantText, }, } resp.Model = req.Model resp.StopReason = "end_turn" - resp.Usage.InputTokens = 10 - resp.Usage.OutputTokens = 5 - - responseBody, _ := json.Marshal(resp) + resp.Usage.InputTokens = mockInputTokens + resp.Usage.OutputTokens = mockOutputTokens if req.Stream { s.sendAnthropicStream(ctx, w, resp) - } else { - w.Header().Set("Content-Type", "application/json") - w.Header().Set("anthropic-version", "2023-06-01") - w.WriteHeader(http.StatusOK) - if _, err := w.Write(responseBody); err != nil { - s.logger.Error(ctx, "failed to write Anthropic response", - slog.F("request_id", requestID), - slog.Error(err), - slog.F("error_type", "write_error"), - slog.F("likely_cause", "network_error"), - ) - } + return } -} -func (s *Server) sendOpenAIStream(ctx context.Context, w http.ResponseWriter, resp openAIResponse) { - w.Header().Set("Content-Type", "text/event-stream") - w.Header().Set("Cache-Control", "no-cache") - w.Header().Set("Connection", "keep-alive") + responseBody, err := json.Marshal(resp) + if err != nil { + s.logger.Error(ctx, "failed to marshal Anthropic response", slog.Error(err)) + http.Error(w, "failed to marshal response", http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "application/json") + w.Header().Set("anthropic-version", "2023-06-01") w.WriteHeader(http.StatusOK) - - flusher, ok := w.(http.Flusher) - if !ok { - s.logger.Error(ctx, "responseWriter does not support flushing", - slog.F("response_id", resp.ID), + if _, err := w.Write(responseBody); err != nil { + s.logger.Error(ctx, "failed to write Anthropic response", + slog.F("request_id", requestID), + slog.Error(err), + slog.F("error_type", "write_error"), + slog.F("likely_cause", "network_error"), ) - return } - - writeChunk := func(data string) bool { - if _, err := fmt.Fprintf(w, "%s", data); err != nil { - s.logger.Error(ctx, "failed to write OpenAI stream chunk", - slog.F("response_id", resp.ID), - slog.Error(err), - slog.F("error_type", "write_error"), - slog.F("likely_cause", "network_error"), - ) - return false - } - flusher.Flush() - return true - } - - // Send initial chunk - chunk := map[string]interface{}{ - "id": resp.ID, - "object": "chat.completion.chunk", - "created": resp.Created, - "model": resp.Model, - "choices": []map[string]interface{}{ - { - "index": 0, - "delta": map[string]interface{}{ - "role": "assistant", - "content": resp.Choices[0].Message.Content, - }, - "finish_reason": nil, - }, - }, - } - chunkBytes, _ := json.Marshal(chunk) - if !writeChunk(fmt.Sprintf("data: %s\n\n", chunkBytes)) { - return - } - - // Send final chunk - finalChunk := map[string]interface{}{ - "id": resp.ID, - "object": "chat.completion.chunk", - "created": resp.Created, - "model": resp.Model, - "choices": []map[string]interface{}{ - { - "index": 0, - "delta": map[string]interface{}{}, - "finish_reason": resp.Choices[0].FinishReason, - }, - }, - } - finalChunkBytes, _ := json.Marshal(finalChunk) - if !writeChunk(fmt.Sprintf("data: %s\n\n", finalChunkBytes)) { - return - } - writeChunk("data: [DONE]\n\n") } func (s *Server) sendResponsesStream(ctx context.Context, w http.ResponseWriter, resp responsesResponse) { - w.Header().Set("Content-Type", "text/event-stream") - w.Header().Set("Cache-Control", "no-cache") - w.Header().Set("Connection", "keep-alive") - w.WriteHeader(http.StatusOK) - flusher, ok := w.(http.Flusher) if !ok { s.logger.Error(ctx, "responseWriter does not support flushing", @@ -519,6 +586,11 @@ func (s *Server) sendResponsesStream(ctx context.Context, w http.ResponseWriter, return } + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.WriteHeader(http.StatusOK) + writeChunk := func(data string) bool { if _, err := fmt.Fprintf(w, "%s", data); err != nil { s.logger.Error(ctx, "failed to write OpenAI responses stream chunk", @@ -533,26 +605,32 @@ func (s *Server) sendResponsesStream(ctx context.Context, w http.ResponseWriter, return true } - deltaChunk := map[string]interface{}{ - "id": resp.ID, - "object": "response.output_text.delta", - "created": resp.Created, - "model": resp.Model, - "output_index": 0, - "content_index": 0, - "delta": resp.Output[0].Content[0].Text, + text := resp.Output[0].Content[0].Text + for chunk := range s.streamContentChunks(ctx, s.randomStreamDuration(), text) { + deltaChunk := map[string]any{ + "id": resp.ID, + "object": "response.output_text.delta", + "created": resp.Created, + "model": resp.Model, + "output_index": 0, + "content_index": 0, + "delta": chunk, + } + deltaBytes, _ := json.Marshal(deltaChunk) + if !writeChunk(fmt.Sprintf("data: %s\n\n", deltaBytes)) { + return + } } - deltaBytes, _ := json.Marshal(deltaChunk) - if !writeChunk(fmt.Sprintf("data: %s\n\n", deltaBytes)) { + if ctx.Err() != nil { return } - finalChunk := map[string]interface{}{ + finalChunk := map[string]any{ "id": resp.ID, "object": "response.completed", "created": resp.Created, "model": resp.Model, - "response": map[string]interface{}{ + "response": map[string]any{ "id": resp.ID, "object": resp.Object, "created": resp.Created, @@ -565,16 +643,10 @@ func (s *Server) sendResponsesStream(ctx context.Context, w http.ResponseWriter, if !writeChunk(fmt.Sprintf("data: %s\n\n", finalBytes)) { return } - writeChunk("data: [DONE]\n\n") + _ = writeChunk("data: [DONE]\n\n") } func (s *Server) sendAnthropicStream(ctx context.Context, w http.ResponseWriter, resp anthropicResponse) { - w.Header().Set("Content-Type", "text/event-stream") - w.Header().Set("Cache-Control", "no-cache") - w.Header().Set("Connection", "keep-alive") - w.Header().Set("anthropic-version", "2023-06-01") - w.WriteHeader(http.StatusOK) - flusher, ok := w.(http.Flusher) if !ok { s.logger.Error(ctx, "responseWriter does not support flushing", @@ -583,6 +655,12 @@ func (s *Server) sendAnthropicStream(ctx context.Context, w http.ResponseWriter, return } + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.Header().Set("anthropic-version", "2023-06-01") + w.WriteHeader(http.StatusOK) + writeChunk := func(eventType string, data []byte) bool { if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", eventType, data); err != nil { s.logger.Error(ctx, "failed to write Anthropic stream chunk", @@ -598,9 +676,9 @@ func (s *Server) sendAnthropicStream(ctx context.Context, w http.ResponseWriter, } startEventType := "message_start" - startEvent := map[string]interface{}{ + startEvent := map[string]any{ "type": startEventType, - "message": map[string]interface{}{ + "message": map[string]any{ "id": resp.ID, "type": resp.Type, "role": resp.Role, @@ -612,14 +690,13 @@ func (s *Server) sendAnthropicStream(ctx context.Context, w http.ResponseWriter, return } - // Send content_block_start event contentStartEventType := "content_block_start" - contentStartEvent := map[string]interface{}{ + contentStartEvent := map[string]any{ "type": contentStartEventType, "index": 0, - "content_block": map[string]interface{}{ + "content_block": map[string]any{ "type": "text", - "text": resp.Content[0].Text, + "text": "", }, } contentStartBytes, _ := json.Marshal(contentStartEvent) @@ -627,24 +704,27 @@ func (s *Server) sendAnthropicStream(ctx context.Context, w http.ResponseWriter, return } - // Send content_block_delta event deltaEventType := "content_block_delta" - deltaEvent := map[string]interface{}{ - "type": deltaEventType, - "index": 0, - "delta": map[string]interface{}{ - "type": "text_delta", - "text": resp.Content[0].Text, - }, + for chunk := range s.streamContentChunks(ctx, s.randomStreamDuration(), resp.Content[0].Text) { + deltaEvent := map[string]any{ + "type": deltaEventType, + "index": 0, + "delta": map[string]any{ + "type": "text_delta", + "text": chunk, + }, + } + deltaBytes, _ := json.Marshal(deltaEvent) + if !writeChunk(deltaEventType, deltaBytes) { + return + } } - deltaBytes, _ := json.Marshal(deltaEvent) - if !writeChunk(deltaEventType, deltaBytes) { + if ctx.Err() != nil { return } - // Send content_block_stop event contentStopEventType := "content_block_stop" - contentStopEvent := map[string]interface{}{ + contentStopEvent := map[string]any{ "type": contentStopEventType, "index": 0, } @@ -653,11 +733,10 @@ func (s *Server) sendAnthropicStream(ctx context.Context, w http.ResponseWriter, return } - // Send message_delta event deltaMsgEventType := "message_delta" - deltaMsgEvent := map[string]interface{}{ + deltaMsgEvent := map[string]any{ "type": deltaMsgEventType, - "delta": map[string]interface{}{ + "delta": map[string]any{ "stop_reason": resp.StopReason, "stop_sequence": resp.StopSequence, }, @@ -668,13 +747,12 @@ func (s *Server) sendAnthropicStream(ctx context.Context, w http.ResponseWriter, return } - // Send message_stop event stopEventType := "message_stop" - stopEvent := map[string]interface{}{ + stopEvent := map[string]any{ "type": stopEventType, } stopBytes, _ := json.Marshal(stopEvent) - writeChunk(stopEventType, stopBytes) + _ = writeChunk(stopEventType, stopBytes) } func (s *Server) tracingMiddleware(next http.Handler) http.Handler {