diff --git a/openai_response.go b/openai_response.go index a73bab2..9c58a18 100644 --- a/openai_response.go +++ b/openai_response.go @@ -187,10 +187,11 @@ func (p *OpenAIResponseProvider) handleNonStreamingResponse(w http.ResponseWrite // handleResponsesStreamingResponse sends a streaming SSE response for Responses API func (p *OpenAIResponseProvider) handleResponsesStreamingResponse(w http.ResponseWriter, response responses.Response) { - 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) + events, err := buildResponseStream(response) + if err != nil { + http.Error(w, fmt.Sprintf("Invalid streaming response fixture: %v", err), http.StatusInternalServerError) + return + } flusher, ok := w.(http.Flusher) if !ok { @@ -198,123 +199,198 @@ func (p *OpenAIResponseProvider) handleResponsesStreamingResponse(w http.Respons return } - // Helper to send a chunk - sendChunk := func(chunk map[string]interface{}) { - jsonBytes, err := json.Marshal(chunk) + 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) + + for _, event := range events { + fmt.Fprintf(w, "event: %s\n", event.eventType) + fmt.Fprintf(w, "data: %s\n\n", event.data) + flusher.Flush() + } + + fmt.Fprint(w, "data: [DONE]\n\n") + flusher.Flush() +} + +// streamEvent separates semantic event construction from SSE framing. +type streamEvent struct { + eventType string + data json.RawMessage +} + +type outputItemStreamPayload struct { + Type string `json:"type"` + SequenceNumber int64 `json:"sequence_number"` + OutputIndex int64 `json:"output_index"` + Item any `json:"item"` +} + +type contentPartStreamPayload struct { + Type string `json:"type"` + SequenceNumber int64 `json:"sequence_number"` + ItemID string `json:"item_id"` + OutputIndex int64 `json:"output_index"` + ContentIndex int64 `json:"content_index"` + Part any `json:"part"` +} + +type messageStreamItem struct { + ID string `json:"id"` + Type string `json:"type"` + Status string `json:"status"` + Role string `json:"role"` + Content any `json:"content"` +} + +type functionCallStreamItem struct { + ID string `json:"id"` + Type string `json:"type"` + Status string `json:"status"` + CallID string `json:"call_id"` + Name string `json:"name"` + Namespace string `json:"namespace,omitempty"` + Arguments string `json:"arguments"` +} + +type outputTextStreamPart struct { + Type string `json:"type"` + Text string `json:"text"` + Annotations []responses.ResponseOutputTextAnnotationUnion `json:"annotations"` + Logprobs []responses.ResponseOutputTextLogprob `json:"logprobs,omitempty"` +} + +// buildResponseStream translates one final Responses API fixture into the +// lifecycle events emitted by the real streaming API. It never changes response. +func buildResponseStream(response responses.Response) ([]streamEvent, error) { + if response.ID == "" { + return nil, fmt.Errorf("response id is required") + } + + inProgress := response + inProgress.Status = responses.ResponseStatusInProgress + inProgress.Output = []responses.ResponseOutputItemUnion{} + + completed := response + completed.Status = responses.ResponseStatusCompleted + + sequence := int64(0) + events := make([]streamEvent, 0) + add := func(eventType string, payload any) error { + data, err := json.Marshal(payload) if err != nil { - fmt.Printf("Failed to encode chunk: %v\n", err) - return + return fmt.Errorf("marshal %s: %w", eventType, err) } - fmt.Fprintf(w, "data: %s\n\n", jsonBytes) - flusher.Flush() + events = append(events, streamEvent{eventType: eventType, data: data}) + sequence++ + return nil } - // Send response.created event - chunk := map[string]interface{}{ - "type": "response.created", + if err := add("response.created", responses.ResponseCreatedEvent{SequenceNumber: sequence, Response: inProgress}); err != nil { + return nil, err + } + if err := add("response.in_progress", responses.ResponseInProgressEvent{SequenceNumber: sequence, Response: inProgress}); err != nil { + return nil, err } - sendChunk(chunk) - for _, outputItem := range response.Output { + for outputIndex, outputItem := range response.Output { switch outputItem.Type { case "message": message := outputItem.AsMessage() - - for _, contentItem := range message.Content { - if contentItem.Type == "output_text" && contentItem.Text != "" { - text := contentItem.Text - // Send delta chunks of 10 characters at a time - chunkSize := 10 - for i := 0; i < len(text); i += chunkSize { - end := i + chunkSize - if end > len(text) { - end = len(text) - } - delta := text[i:end] - - chunk := map[string]interface{}{ - "type": "response.output_text.delta", - "delta": delta, - } - sendChunk(chunk) - } - - // Send done event with complete text - chunk := map[string]interface{}{ - "type": "response.output_text.done", - "json": map[string]interface{}{ - "text": text, - }, - "text": text, - } - sendChunk(chunk) + if message.ID == "" { + return nil, fmt.Errorf("output[%d] message id is required", outputIndex) + } + for contentIndex, contentItem := range message.Content { + if contentItem.Type != "output_text" { + return nil, fmt.Errorf("output[%d].content[%d] type %q is unsupported", outputIndex, contentIndex, contentItem.Type) } } - case "function_call": - funcCall := outputItem.AsFunctionCall() - // Send function call arguments delta - if funcCall.Arguments != "" { - chunk := map[string]interface{}{ - "type": "response.function_call_arguments.delta", - "name": funcCall.Name, - "arguments": funcCall.Arguments, + addedItem := messageStreamItem{ID: message.ID, Type: "message", Status: "in_progress", Role: "assistant", Content: []any{}} + if err := add("response.output_item.added", outputItemStreamPayload{Type: "response.output_item.added", SequenceNumber: sequence, OutputIndex: int64(outputIndex), Item: addedItem}); err != nil { + return nil, err + } + + for contentIndex, contentItem := range message.Content { + content := contentItem.AsOutputText() + addedPart := outputTextStreamPart{Type: "output_text", Text: "", Annotations: []responses.ResponseOutputTextAnnotationUnion{}} + if err := add("response.content_part.added", contentPartStreamPayload{Type: "response.content_part.added", SequenceNumber: sequence, ItemID: message.ID, OutputIndex: int64(outputIndex), ContentIndex: int64(contentIndex), Part: addedPart}); err != nil { + return nil, err + } + if err := add("response.output_text.delta", responses.ResponseTextDeltaEvent{SequenceNumber: sequence, ItemID: message.ID, OutputIndex: int64(outputIndex), ContentIndex: int64(contentIndex), Delta: content.Text, Logprobs: []responses.ResponseTextDeltaEventLogprob{}}); err != nil { + return nil, err + } + if err := add("response.output_text.done", responses.ResponseTextDoneEvent{SequenceNumber: sequence, ItemID: message.ID, OutputIndex: int64(outputIndex), ContentIndex: int64(contentIndex), Text: content.Text, Logprobs: []responses.ResponseTextDoneEventLogprob{}}); err != nil { + return nil, err + } + donePart := outputTextStreamPart{Type: "output_text", Text: content.Text, Annotations: content.Annotations, Logprobs: content.Logprobs} + if err := add("response.content_part.done", contentPartStreamPayload{Type: "response.content_part.done", SequenceNumber: sequence, ItemID: message.ID, OutputIndex: int64(outputIndex), ContentIndex: int64(contentIndex), Part: donePart}); err != nil { + return nil, err } - sendChunk(chunk) } - // Send function call arguments done - chunk := map[string]interface{}{ - "type": "response.function_call_arguments.done", - "name": funcCall.Name, - "arguments": funcCall.Arguments, - "call_id": funcCall.CallID, + doneItem := messageStreamItem{ID: message.ID, Type: "message", Status: "completed", Role: "assistant", Content: message.Content} + if err := add("response.output_item.done", outputItemStreamPayload{Type: "response.output_item.done", SequenceNumber: sequence, OutputIndex: int64(outputIndex), Item: doneItem}); err != nil { + return nil, err } - sendChunk(chunk) - - case "function_call_output": - callID := outputItem.CallID - output := "" - - // The Output field is a union type, try to extract string value - if outputItem.Output.OfString != "" { - output = outputItem.Output.OfString - } else { - // If it's not a string, marshal it to JSON - // These are meant for complex OpenAI built-in tools, unlikely to be used in mocks - if outputBytes, err := json.Marshal(outputItem.Output); err == nil { - output = string(outputBytes) - } + + case "function_call": + functionCall := outputItem.AsFunctionCall() + var rawFunctionCall struct { + Namespace string `json:"namespace"` + } + if err := json.Unmarshal([]byte(outputItem.RawJSON()), &rawFunctionCall); err != nil { + return nil, fmt.Errorf("decode output[%d] function call: %w", outputIndex, err) + } + if functionCall.ID == "" { + return nil, fmt.Errorf("output[%d] function call id is required", outputIndex) + } + if functionCall.CallID == "" || functionCall.Name == "" { + return nil, fmt.Errorf("output[%d] function call requires call_id and name", outputIndex) + } + if !json.Valid([]byte(functionCall.Arguments)) { + return nil, fmt.Errorf("output[%d] function call arguments must be valid JSON", outputIndex) } - if output != "" { - // Send function call output delta - chunk := map[string]interface{}{ - "type": "response.function_call_output.delta", - "call_id": callID, - "output": output, - } - sendChunk(chunk) + addedItem := functionCallStreamItem{ID: functionCall.ID, Type: "function_call", Status: "in_progress", CallID: functionCall.CallID, Name: functionCall.Name, Namespace: rawFunctionCall.Namespace, Arguments: ""} + if err := add("response.output_item.added", outputItemStreamPayload{Type: "response.output_item.added", SequenceNumber: sequence, OutputIndex: int64(outputIndex), Item: addedItem}); err != nil { + return nil, err + } + if err := add("response.function_call_arguments.delta", responses.ResponseFunctionCallArgumentsDeltaEvent{SequenceNumber: sequence, ItemID: functionCall.ID, OutputIndex: int64(outputIndex), Delta: functionCall.Arguments}); err != nil { + return nil, err + } + if err := add("response.function_call_arguments.done", responses.ResponseFunctionCallArgumentsDoneEvent{SequenceNumber: sequence, ItemID: functionCall.ID, OutputIndex: int64(outputIndex), Arguments: functionCall.Arguments, Name: functionCall.Name}); err != nil { + return nil, err + } + doneItem := functionCallStreamItem{ID: functionCall.ID, Type: "function_call", Status: "completed", CallID: functionCall.CallID, Name: functionCall.Name, Namespace: rawFunctionCall.Namespace, Arguments: functionCall.Arguments} + if err := add("response.output_item.done", outputItemStreamPayload{Type: "response.output_item.done", SequenceNumber: sequence, OutputIndex: int64(outputIndex), Item: doneItem}); err != nil { + return nil, err + } - // Send function call output done - chunk = map[string]interface{}{ - "type": "response.function_call_output.done", - "call_id": callID, - "output": output, - } - sendChunk(chunk) + case "tool_search_call": + rawToolSearchCall := json.RawMessage(outputItem.RawJSON()) + var toolSearchCall struct { + ID string `json:"id"` + CallID string `json:"call_id"` + } + if err := json.Unmarshal(rawToolSearchCall, &toolSearchCall); err != nil { + return nil, fmt.Errorf("decode output[%d] tool search call: %w", outputIndex, err) + } + if toolSearchCall.ID == "" || toolSearchCall.CallID == "" { + return nil, fmt.Errorf("output[%d] tool search call requires id and call_id", outputIndex) + } + if err := add("response.output_item.done", outputItemStreamPayload{Type: "response.output_item.done", SequenceNumber: sequence, OutputIndex: int64(outputIndex), Item: rawToolSearchCall}); err != nil { + return nil, err } + + default: + return nil, fmt.Errorf("output[%d] type %q is unsupported for streaming", outputIndex, outputItem.Type) } } - // Send response.completed event - chunk = map[string]interface{}{ - "type": "response.completed", - "response": response, + if err := add("response.completed", responses.ResponseCompletedEvent{SequenceNumber: sequence, Response: completed}); err != nil { + return nil, err } - sendChunk(chunk) - - // Send DONE - fmt.Fprint(w, "data: [DONE]\n\n") - flusher.Flush() + return events, nil } diff --git a/openai_response_test.go b/openai_response_test.go new file mode 100644 index 0000000..9d30746 --- /dev/null +++ b/openai_response_test.go @@ -0,0 +1,257 @@ +package mockllm + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/openai/openai-go/v3/responses" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func responseFixture(t *testing.T, outputJSON string) responses.Response { + t.Helper() + data := `{ + "id":"resp_test","object":"response","created_at":1,"model":"gpt-5.2-codex", + "status":"completed","output":` + outputJSON + `, + "parallel_tool_calls":false,"tools":[] + }` + var response responses.Response + require.NoError(t, json.Unmarshal([]byte(data), &response)) + return response +} + +func decodeStreamEvents(t *testing.T, events []streamEvent) []map[string]any { + t.Helper() + decoded := make([]map[string]any, len(events)) + for i, event := range events { + require.NoError(t, json.Unmarshal(event.data, &decoded[i])) + assert.Equal(t, event.eventType, decoded[i]["type"]) + assert.Equal(t, float64(i), decoded[i]["sequence_number"]) + + // Verify that every payload is recognized by the pinned official SDK's + // Responses API stream-event union. + var official responses.ResponseStreamEventUnion + require.NoError(t, json.Unmarshal(event.data, &official)) + assert.Equal(t, event.eventType, official.Type) + } + return decoded +} + +func eventTypes(events []streamEvent) []string { + types := make([]string, len(events)) + for i, event := range events { + types[i] = event.eventType + } + return types +} + +func TestBuildResponseStreamAssistantLifecycle(t *testing.T) { + response := responseFixture(t, `[{"id":"msg_1","type":"message","status":"completed","role":"assistant","content":[{"type":"output_text","text":"Hello, δΈ–η•Œ πŸ‘‹","annotations":[]}]}]`) + original, err := json.Marshal(response) + require.NoError(t, err) + + events, err := buildResponseStream(response) + require.NoError(t, err) + assert.Equal(t, []string{ + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.delta", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.completed", + }, eventTypes(events)) + decoded := decodeStreamEvents(t, events) + + for _, i := range []int{0, 1} { + snapshot := decoded[i]["response"].(map[string]any) + assert.Equal(t, "in_progress", snapshot["status"]) + assert.Empty(t, snapshot["output"]) + } + + added := decoded[2]["item"].(map[string]any) + assert.Equal(t, "msg_1", added["id"]) + assert.Equal(t, "in_progress", added["status"]) + assert.Empty(t, added["content"]) + + for _, i := range []int{3, 4, 5, 6} { + assert.Equal(t, "msg_1", decoded[i]["item_id"]) + assert.Equal(t, float64(0), decoded[i]["output_index"]) + assert.Equal(t, float64(0), decoded[i]["content_index"]) + } + assert.Equal(t, "Hello, δΈ–η•Œ πŸ‘‹", decoded[4]["delta"]) + assert.Equal(t, "Hello, δΈ–η•Œ πŸ‘‹", decoded[5]["text"]) + + done := decoded[7]["item"].(map[string]any) + assert.Equal(t, "completed", done["status"]) + require.Len(t, done["content"], 1) + assert.Equal(t, "Hello, δΈ–η•Œ πŸ‘‹", done["content"].([]any)[0].(map[string]any)["text"]) + + completed := decoded[8]["response"].(map[string]any) + assert.Equal(t, "completed", completed["status"]) + assert.Len(t, completed["output"], 1) + var originalWire map[string]any + require.NoError(t, json.Unmarshal(original, &originalWire)) + assert.Equal(t, originalWire["output"], completed["output"], "completed response must retain the fixture's final output") + + after, err := json.Marshal(response) + require.NoError(t, err) + assert.JSONEq(t, string(original), string(after), "building snapshots must not mutate the fixture") +} + +func TestBuildResponseStreamFunctionCallLifecycle(t *testing.T) { + response := responseFixture(t, `[{"id":"fc_1","type":"function_call","status":"completed","call_id":"call_1","namespace":"mcp__tools","name":"calculate","arguments":"{\"n\":2}"}]`) + events, err := buildResponseStream(response) + require.NoError(t, err) + assert.Equal(t, []string{ + "response.created", + "response.in_progress", + "response.output_item.added", + "response.function_call_arguments.delta", + "response.function_call_arguments.done", + "response.output_item.done", + "response.completed", + }, eventTypes(events)) + decoded := decodeStreamEvents(t, events) + + added := decoded[2]["item"].(map[string]any) + assert.Equal(t, "in_progress", added["status"]) + assert.Equal(t, "", added["arguments"]) + assert.Equal(t, "mcp__tools", added["namespace"]) + assert.Equal(t, "fc_1", decoded[3]["item_id"]) + assert.Equal(t, float64(0), decoded[3]["output_index"]) + assert.Equal(t, `{"n":2}`, decoded[3]["delta"]) + assert.NotContains(t, decoded[3], "arguments") + assert.Equal(t, `{"n":2}`, decoded[4]["arguments"]) + assert.Equal(t, "calculate", decoded[4]["name"]) + assert.Equal(t, "fc_1", decoded[4]["item_id"]) + assert.Equal(t, float64(0), decoded[4]["output_index"]) + done := decoded[5]["item"].(map[string]any) + assert.Equal(t, "completed", done["status"]) + assert.Equal(t, "mcp__tools", done["namespace"]) + + for _, event := range events { + assert.NotContains(t, event.eventType, "function_call_output") + } +} + +func TestBuildResponseStreamToolSearchCall(t *testing.T) { + response := responseFixture(t, `[{"id":"ts_1","type":"tool_search_call","status":"completed","call_id":"search_1","execution":"client","arguments":{"query":"mcp__tools add_numbers","limit":10}}]`) + events, err := buildResponseStream(response) + require.NoError(t, err) + assert.Equal(t, []string{ + "response.created", + "response.in_progress", + "response.output_item.done", + "response.completed", + }, eventTypes(events)) + + var done map[string]any + require.NoError(t, json.Unmarshal(events[2].data, &done)) + assert.Equal(t, float64(2), done["sequence_number"]) + item := done["item"].(map[string]any) + assert.Equal(t, "tool_search_call", item["type"]) + assert.Equal(t, "ts_1", item["id"]) + assert.Equal(t, "search_1", item["call_id"]) + assert.Equal(t, "client", item["execution"]) + assert.Equal(t, "mcp__tools add_numbers", item["arguments"].(map[string]any)["query"]) +} + +func TestBuildResponseStreamEmptyText(t *testing.T) { + response := responseFixture(t, `[{"id":"msg_empty","type":"message","status":"completed","role":"assistant","content":[{"type":"output_text","text":"","annotations":[]}]}]`) + events, err := buildResponseStream(response) + require.NoError(t, err) + decoded := decodeStreamEvents(t, events) + assert.Contains(t, decoded[4], "delta") + assert.Equal(t, "", decoded[4]["delta"]) + assert.Contains(t, decoded[5], "text") + assert.Equal(t, "", decoded[5]["text"]) +} + +func TestBuildResponseStreamMultipleOutputsAndContent(t *testing.T) { + response := responseFixture(t, `[ + {"id":"msg_0","type":"message","status":"completed","role":"assistant","content":[ + {"type":"output_text","text":"zero","annotations":[]}, + {"type":"output_text","text":"one","annotations":[]} + ]}, + {"id":"fc_1","type":"function_call","status":"completed","call_id":"call_1","name":"tool","arguments":"{}"}, + {"id":"msg_2","type":"message","status":"completed","role":"assistant","content":[ + {"type":"output_text","text":"two","annotations":[]} + ]} + ]`) + events, err := buildResponseStream(response) + require.NoError(t, err) + decoded := decodeStreamEvents(t, events) + + activeItems := map[string]bool{} + activeParts := map[string]bool{} + seenDeltaCoordinates := [][2]int{} + for _, event := range decoded { + typ := event["type"].(string) + switch typ { + case "response.output_item.added": + activeItems[event["item"].(map[string]any)["id"].(string)] = true + case "response.content_part.added": + key := event["item_id"].(string) + ":" + eventNumber(event, "content_index") + assert.True(t, activeItems[event["item_id"].(string)]) + activeParts[key] = true + case "response.output_text.delta": + key := event["item_id"].(string) + ":" + eventNumber(event, "content_index") + assert.True(t, activeItems[event["item_id"].(string)], "delta requires an active output item") + assert.True(t, activeParts[key], "delta requires an active content part") + seenDeltaCoordinates = append(seenDeltaCoordinates, [2]int{int(event["output_index"].(float64)), int(event["content_index"].(float64))}) + case "response.content_part.done": + key := event["item_id"].(string) + ":" + eventNumber(event, "content_index") + delete(activeParts, key) + case "response.output_item.done": + delete(activeItems, event["item"].(map[string]any)["id"].(string)) + } + } + assert.Equal(t, [][2]int{{0, 0}, {0, 1}, {2, 0}}, seenDeltaCoordinates) + assert.Empty(t, activeItems) + assert.Empty(t, activeParts) +} + +func eventNumber(event map[string]any, key string) string { + data, _ := json.Marshal(event[key]) + return string(data) +} + +func TestStreamingHandlerSSEFramingAndValidation(t *testing.T) { + valid := responseFixture(t, `[{"id":"msg_1","type":"message","status":"completed","role":"assistant","content":[{"type":"output_text","text":"ok","annotations":[]}]}]`) + recorder := httptest.NewRecorder() + provider := NewOpenAIResponseProvider(nil) + provider.handleResponsesStreamingResponse(recorder, valid) + + result := recorder.Result() + assert.Equal(t, http.StatusOK, result.StatusCode) + assert.Equal(t, "text/event-stream", result.Header.Get("Content-Type")) + body := recorder.Body.String() + assert.Contains(t, body, "event: response.created\ndata: {") + assert.True(t, strings.HasSuffix(body, "data: [DONE]\n\n")) + + malformed := responseFixture(t, `[{"type":"message","status":"completed","role":"assistant","content":[]}]`) + recorder = httptest.NewRecorder() + provider.handleResponsesStreamingResponse(recorder, malformed) + assert.Equal(t, http.StatusInternalServerError, recorder.Code) + assert.Contains(t, recorder.Body.String(), "output[0] message id is required") + assert.NotEqual(t, "text/event-stream", recorder.Header().Get("Content-Type"), "validation must happen before SSE headers") +} + +func TestBuildResponseStreamRejectsUnsupportedOutput(t *testing.T) { + response := responseFixture(t, `[{"type":"reasoning","id":"reason_1","status":"completed","summary":[]}]`) + _, err := buildResponseStream(response) + require.Error(t, err) + assert.Contains(t, err.Error(), `type "reasoning" is unsupported for streaming`) + + response = responseFixture(t, `[{"type":"function_call_output","call_id":"call_1","output":"done"}]`) + _, err = buildResponseStream(response) + require.Error(t, err) + assert.Contains(t, err.Error(), `type "function_call_output" is unsupported for streaming`) +} diff --git a/server_test.go b/server_test.go index 9e93ab6..370737f 100644 --- a/server_test.go +++ b/server_test.go @@ -321,7 +321,7 @@ func TestOpenAIResponseMock(t *testing.T) { loadTestData(t, "openai_response_mock.json", &haikuResponse) haikuMock := newOpenAIResponseMock(t, "haiku-response", mockllm.MatchTypeContains, haikuRequest, haikuResponse) - // Setup function output response for streaming + // The first tool-round-trip request returns only the model's function call. funcRequest := responses.ResponseNewParams{ Model: openai.ChatModelGPT4, Input: responses.ResponseNewParamsInputUnion{ @@ -333,9 +333,19 @@ func TestOpenAIResponseMock(t *testing.T) { loadTestData(t, "openai_response_function_mock.json", &funcResponse) funcMock := newOpenAIResponseMock(t, "function-response", mockllm.MatchTypeContains, funcRequest, funcResponse) + // The follow-up request is selected by the function_call_output in the last + // input-list position, not by the original prompt earlier in the list. + funcOutputRequest := responses.ResponseNewParams{ + Model: openai.ChatModelGPT4, + Input: responses.ResponseNewParamsInputUnion{ + OfString: openai.String("call_func_123"), + }, + } + funcOutputMock := newOpenAIResponseMock(t, "function-output-response", mockllm.MatchTypeContains, funcOutputRequest, haikuResponse) + // Setup both mocks on the same server and use request matching to determine which response to return config := mockllm.Config{ - OpenAIResponse: []mockllm.OpenAIResponseMock{haikuMock, funcMock}, + OpenAIResponse: []mockllm.OpenAIResponseMock{haikuMock, funcMock, funcOutputMock}, } baseURL, cleanup := startTestServer(t, config) @@ -356,33 +366,26 @@ func TestOpenAIResponseMock(t *testing.T) { assert.Contains(t, outputText, "Kagent finds its mind") }) - // Test streaming response (function output) - t.Run("streaming", func(t *testing.T) { - stream := client.Responses.NewStreaming(t.Context(), funcRequest) - - var receivedEvents []string - var functionCallFound, functionOutputFound bool - - for stream.Next() { - data := stream.Current() - if data.Type != "" { - receivedEvents = append(receivedEvents, data.Type) - if data.Type == "response.function_call_arguments.done" { - functionCallFound = true - } - if data.Type == "response.function_call_output.done" { - functionOutputFound = true - } - } - } - - if err := stream.Err(); err != nil { - t.Fatalf("Stream error: %v", err) + t.Run("tool round trip matches only last list item", func(t *testing.T) { + first, err := client.Responses.New(t.Context(), funcRequest) + require.NoError(t, err) + require.Len(t, first.Output, 1) + assert.Equal(t, "function_call", first.Output[0].Type) + + followUp := responses.ResponseNewParams{ + Model: openai.ChatModelGPT4, + Input: responses.ResponseNewParamsInputUnion{ + OfInputItemList: responses.ResponseInputParam{ + responses.ResponseInputItemParamOfMessage("Calculate 2+2", responses.EasyInputMessageRoleUser), + responses.ResponseInputItemParamOfFunctionCall("{\"expression\":\"2+2\"}", "call_func_123", "calculate"), + responses.ResponseInputItemParamOfFunctionCallOutput("call_func_123", "4"), + }, + }, } - - assert.Greater(t, len(receivedEvents), 0) - assert.True(t, functionCallFound, "Should receive function_call_arguments.done event") - assert.True(t, functionOutputFound, "Should receive function_call_output.done event") + second, err := client.Responses.New(t.Context(), followUp) + require.NoError(t, err) + assert.Equal(t, "resp_123", second.ID, "the original prompt must not select the first mock again") + assert.Contains(t, second.OutputText(), "Kagent finds its mind") }) } diff --git a/testdata/openai_response_function_mock.json b/testdata/openai_response_function_mock.json index efd866c..95966f8 100644 --- a/testdata/openai_response_function_mock.json +++ b/testdata/openai_response_function_mock.json @@ -7,14 +7,10 @@ { "type": "function_call", "id": "call_func_123", + "status": "completed", "name": "calculate", "arguments": "{\"expression\": \"2+2\"}", "call_id": "call_func_123" - }, - { - "type": "function_call_output", - "call_id": "call_func_123", - "output": "4" } ], "error": {},