diff --git a/pkg/responses/handler_test.go b/pkg/responses/handler_test.go index 255320ee4..cf30ae3b8 100644 --- a/pkg/responses/handler_test.go +++ b/pkg/responses/handler_test.go @@ -1260,6 +1260,186 @@ func TestHandler_CreateResponse_Streaming_Persistence(t *testing.T) { } } +func TestHandler_CreateResponse_Streaming_ToolCallArgumentChunks(t *testing.T) { + // Chat completion streams send the id and name only in the first chunk of + // each tool call. The argument chunks that follow carry just the index + // (this is what llama.cpp and OpenAI send). + chunk := func(delta string) string { + return "data: {\"id\":\"chatcmpl-1\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":" + delta + ",\"finish_reason\":null}]}\n\n" + } + mock := &mockSchedulerHTTP{ + streaming: true, + streamChunks: []string{ + chunk(`{"role":"assistant","content":"Checking the weather."}`), + chunk(`{"tool_calls":[{"index":0,"id":"call_a","type":"function","function":{"name":"get_weather"}}]}`), + chunk(`{"tool_calls":[{"index":0,"function":{"arguments":"{\"city\":"}}]}`), + chunk(`{"tool_calls":[{"index":0,"function":{"arguments":"\"Paris\"}"}}]}`), + chunk(`{"tool_calls":[{"index":1,"id":"call_b","type":"function","function":{"name":"get_time"}}]}`), + chunk(`{"tool_calls":[{"index":1,"function":{"arguments":"{\"tz\":\"CET\"}"}}]}`), + "data: {\"id\":\"chatcmpl-1\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"tool_calls\"}]}\n\n", + "data: [DONE]\n\n", + }, + } + + handler := newTestHandler(t, mock) + + reqBody := `{"model": "gpt-4", "input": "Weather and time in Paris?", "stream": true}` + req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(reqBody)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + + handler.handleCreate(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", w.Code, http.StatusOK) + } + + ids := handler.store.GetResponseIDs() + if len(ids) != 1 { + t.Fatalf("expected one stored response, got %d", len(ids)) + } + persisted, ok := handler.store.Get(ids[0]) + if !ok { + t.Fatal("stored response not found") + } + + type call struct{ callID, name, args string } + var got []call + for _, item := range persisted.Output { + if item.Type == ItemTypeFunctionCall { + got = append(got, call{item.CallID, item.Name, item.Arguments}) + } + } + want := []call{ + {"call_a", "get_weather", `{"city":"Paris"}`}, + {"call_b", "get_time", `{"tz":"CET"}`}, + } + if len(got) != len(want) { + t.Fatalf("function_call items = %+v, want %+v", got, want) + } + for i := range want { + if got[i] != want[i] { + t.Errorf("function_call[%d] = %+v, want %+v", i, got[i], want[i]) + } + } + + // Argument deltas must point at the output item they belong to. + for _, line := range strings.Split(w.Body.String(), "\n") { + data, found := strings.CutPrefix(line, "data: ") + if !found { + continue + } + if data == "[DONE]" { + continue + } + var ev StreamEvent + if err := json.Unmarshal([]byte(data), &ev); err != nil { + t.Fatalf("bad event %q: %v", data, err) + } + if ev.Type != EventFunctionCallArgsDelta { + continue + } + wantIndex := 1 // The assistant text item precedes the tool calls. + if strings.Contains(ev.Delta, "tz") { + wantIndex = 2 + } + if ev.OutputIndex != wantIndex { + t.Errorf("delta %q has output_index %d, want %d", ev.Delta, ev.OutputIndex, wantIndex) + } + } +} + +func TestStreamingResponseWriter_ToolCallBeforeTextKeepsOutputIndex(t *testing.T) { + // A tool call that starts before any assistant text must keep the + // output_index it was added with, and the stored output must list the + // items in that order. + w := httptest.NewRecorder() + resp := &Response{} + s := NewStreamingResponseWriter(w, resp, nil) + idx := 0 + s.handleToolCallDelta([]ChatToolCall{{Index: &idx, ID: "call_a", Function: ChatFunctionCall{Name: "get_weather"}}}) + s.handleToolCallDelta([]ChatToolCall{{Index: &idx, Function: ChatFunctionCall{Arguments: `{"city":`}}}) + s.handleContentDelta("Checking the weather.") + s.handleToolCallDelta([]ChatToolCall{{Index: &idx, Function: ChatFunctionCall{Arguments: `"Paris"}`}}}) + s.finalize() + + indices := map[string]map[int]bool{} + for _, line := range strings.Split(w.Body.String(), "\n") { + data, found := strings.CutPrefix(line, "data: ") + if !found { + continue + } + var ev StreamEvent + if err := json.Unmarshal([]byte(data), &ev); err != nil { + t.Fatalf("bad event %q: %v", data, err) + } + itemID := ev.ItemID + if ev.Item != nil { + itemID = ev.Item.ID + } + if itemID == "" { + continue + } + if indices[itemID] == nil { + indices[itemID] = map[int]bool{} + } + indices[itemID][ev.OutputIndex] = true + } + + if len(resp.Output) != 2 { + t.Fatalf("output = %+v, want a function call and a message", resp.Output) + } + if resp.Output[0].Type != ItemTypeFunctionCall || resp.Output[1].Type != ItemTypeMessage { + t.Errorf("output types = [%s %s], want [%s %s]", + resp.Output[0].Type, resp.Output[1].Type, ItemTypeFunctionCall, ItemTypeMessage) + } + for pos, item := range resp.Output { + got := indices[item.ID] + if len(got) != 1 || !got[pos] { + t.Errorf("%s item %s used output_index %v, want only %d", item.Type, item.ID, got, pos) + } + } +} + +func TestStreamingResponseWriter_ToolCallMatchedByIDThenIndex(t *testing.T) { + // A delta that adds the index to a call first seen by ID must extend + // that call, and later index-only deltas must go to the same call. + w := httptest.NewRecorder() + s := NewStreamingResponseWriter(w, &Response{}, nil) + idx := 0 + s.handleToolCallDelta([]ChatToolCall{{ID: "call_a", Function: ChatFunctionCall{Name: "get_weather"}}}) + s.handleToolCallDelta([]ChatToolCall{{Index: &idx, ID: "call_a", Function: ChatFunctionCall{Arguments: `{"city":`}}}) + s.handleToolCallDelta([]ChatToolCall{{Index: &idx, Function: ChatFunctionCall{Arguments: `"Paris"}`}}}) + + if len(s.toolCalls) != 1 { + t.Fatalf("tool calls = %+v, want one", s.toolCalls) + } + if got := s.toolCalls[0]; got.CallID != "call_a" || got.Name != "get_weather" || got.Arguments != `{"city":"Paris"}` { + t.Errorf("tool call = %+v", got) + } +} + +func TestStreamingResponseWriter_ToolCallsSharingIndexSplitByID(t *testing.T) { + // Some servers send index 0 for every parallel call and tell them apart + // by ID only. Those must stay separate calls. + w := httptest.NewRecorder() + s := NewStreamingResponseWriter(w, &Response{}, nil) + idx := 0 + s.handleToolCallDelta([]ChatToolCall{{Index: &idx, ID: "call_a", Function: ChatFunctionCall{Name: "get_weather", Arguments: `{"city":"Paris"}`}}}) + s.handleToolCallDelta([]ChatToolCall{{Index: &idx, ID: "call_b", Function: ChatFunctionCall{Name: "get_time", Arguments: `{"tz":`}}}) + s.handleToolCallDelta([]ChatToolCall{{Index: &idx, Function: ChatFunctionCall{Arguments: `"CET"}`}}}) + + if len(s.toolCalls) != 2 { + t.Fatalf("tool calls = %+v, want two", s.toolCalls) + } + if got := s.toolCalls[0]; got.CallID != "call_a" || got.Arguments != `{"city":"Paris"}` { + t.Errorf("tool call 0 = %+v", got) + } + if got := s.toolCalls[1]; got.CallID != "call_b" || got.Name != "get_time" || got.Arguments != `{"tz":"CET"}` { + t.Errorf("tool call 1 = %+v", got) + } +} + // Benchmark for response creation func BenchmarkHandler_CreateResponse(b *testing.B) { mock := &mockSchedulerHTTP{ diff --git a/pkg/responses/streaming.go b/pkg/responses/streaming.go index e4637fb58..60fe50a60 100644 --- a/pkg/responses/streaming.go +++ b/pkg/responses/streaming.go @@ -22,6 +22,14 @@ type StreamingResponseWriter struct { currentContentIdx int accumulatedContent strings.Builder toolCalls []OutputItem + // toolCallPos maps a streaming tool call index to its position in toolCalls. + toolCallPos map[int]int + + // Output indices are assigned when an item is added, in arrival order, and + // stay fixed for all of that item's events and its place in the output. + nextOutputIndex int + messageOutputIndex int + toolCallOutputIndices []int } // NewStreamingResponseWriter creates a new streaming response writer. @@ -267,6 +275,7 @@ func (s *StreamingResponseWriter) handleContentDelta(content string) { if s.currentItemID == "" { s.currentItemID = GenerateMessageID() s.currentContentIdx = 0 + s.messageOutputIndex = s.allocOutputIndex() // Send output_item.added item := &OutputItem{ @@ -284,7 +293,7 @@ func (s *StreamingResponseWriter) handleContentDelta(content string) { Type: EventOutputItemAdded, SequenceNumber: s.nextSeq(), Item: item, - OutputIndex: 0, + OutputIndex: s.messageOutputIndex, }) // Send content_part.added @@ -292,7 +301,7 @@ func (s *StreamingResponseWriter) handleContentDelta(content string) { Type: EventContentPartAdded, SequenceNumber: s.nextSeq(), ItemID: s.currentItemID, - OutputIndex: 0, + OutputIndex: s.messageOutputIndex, ContentIndex: 0, Part: &ContentPart{ Type: ContentTypeOutputText, @@ -310,7 +319,7 @@ func (s *StreamingResponseWriter) handleContentDelta(content string) { Type: EventOutputTextDelta, SequenceNumber: s.nextSeq(), ItemID: s.currentItemID, - OutputIndex: 0, + OutputIndex: s.messageOutputIndex, ContentIndex: 0, Delta: content, }) @@ -319,16 +328,37 @@ func (s *StreamingResponseWriter) handleContentDelta(content string) { // handleToolCallDelta handles tool call deltas from the chat completion stream. func (s *StreamingResponseWriter) handleToolCallDelta(toolCalls []ChatToolCall) { for _, tc := range toolCalls { - // Find or create the tool call item - var item *OutputItem - for i := range s.toolCalls { - if s.toolCalls[i].CallID == tc.ID { - item = &s.toolCalls[i] - break + // Find or create the tool call item. Argument deltas after the first + // one carry only the index (no ID), so match on the index first and + // fall back to the ID. + pos := -1 + if tc.Index != nil { + if p, ok := s.toolCallPos[*tc.Index]; ok { + pos = p + // Some servers reuse one index for several calls and tell them + // apart by ID, so an ID that differs starts a different call. + if tc.ID != "" && s.toolCalls[p].CallID != tc.ID { + pos = -1 + } + } + } + if pos < 0 && tc.ID != "" { + for i := range s.toolCalls { + if s.toolCalls[i].CallID == tc.ID { + pos = i + break + } } } - if item == nil { + var item *OutputItem + if pos >= 0 { + item = &s.toolCalls[pos] + if item.Name == "" { + item.Name = tc.Function.Name + } + s.rememberToolCallIndex(tc.Index, pos) + } else { // New tool call callID := tc.ID if callID == "" { @@ -343,14 +373,17 @@ func (s *StreamingResponseWriter) handleToolCallDelta(toolCalls []ChatToolCall) Status: StatusInProgress, } s.toolCalls = append(s.toolCalls, newItem) - item = &s.toolCalls[len(s.toolCalls)-1] + s.toolCallOutputIndices = append(s.toolCallOutputIndices, s.allocOutputIndex()) + pos = len(s.toolCalls) - 1 + item = &s.toolCalls[pos] + s.rememberToolCallIndex(tc.Index, pos) // Send output_item.added for function call s.sendEvent(EventOutputItemAdded, &StreamEvent{ Type: EventOutputItemAdded, SequenceNumber: s.nextSeq(), Item: item, - OutputIndex: len(s.toolCalls) - 1, + OutputIndex: s.toolCallOutputIndex(pos), }) } @@ -363,15 +396,41 @@ func (s *StreamingResponseWriter) handleToolCallDelta(toolCalls []ChatToolCall) Type: EventFunctionCallArgsDelta, SequenceNumber: s.nextSeq(), ItemID: item.ID, - OutputIndex: len(s.toolCalls) - 1, + OutputIndex: s.toolCallOutputIndex(pos), Delta: tc.Function.Arguments, }) } } } +// allocOutputIndex returns the output index for a newly added item. +func (s *StreamingResponseWriter) allocOutputIndex() int { + i := s.nextOutputIndex + s.nextOutputIndex++ + return i +} + +// toolCallOutputIndex returns the output index assigned to toolCalls[pos]. +func (s *StreamingResponseWriter) toolCallOutputIndex(pos int) int { + return s.toolCallOutputIndices[pos] +} + +// rememberToolCallIndex records which item a streaming tool call index refers to. +func (s *StreamingResponseWriter) rememberToolCallIndex(index *int, pos int) { + if index == nil { + return + } + if s.toolCallPos == nil { + s.toolCallPos = make(map[int]int) + } + s.toolCallPos[*index] = pos +} + // finalize completes the streaming response. func (s *StreamingResponseWriter) finalize() { + // Items go into the output at the index their events used. + output := make([]OutputItem, s.nextOutputIndex) + // Finalize any accumulated content if s.currentItemID != "" { finalText := s.accumulatedContent.String() @@ -381,7 +440,7 @@ func (s *StreamingResponseWriter) finalize() { Type: EventOutputTextDone, SequenceNumber: s.nextSeq(), ItemID: s.currentItemID, - OutputIndex: 0, + OutputIndex: s.messageOutputIndex, ContentIndex: 0, Part: &ContentPart{ Type: ContentTypeOutputText, @@ -395,7 +454,7 @@ func (s *StreamingResponseWriter) finalize() { Type: EventContentPartDone, SequenceNumber: s.nextSeq(), ItemID: s.currentItemID, - OutputIndex: 0, + OutputIndex: s.messageOutputIndex, ContentIndex: 0, Part: &ContentPart{ Type: ContentTypeOutputText, @@ -408,7 +467,7 @@ func (s *StreamingResponseWriter) finalize() { s.sendEvent(EventOutputItemDone, &StreamEvent{ Type: EventOutputItemDone, SequenceNumber: s.nextSeq(), - OutputIndex: 0, + OutputIndex: s.messageOutputIndex, Item: &OutputItem{ ID: s.currentItemID, Type: ItemTypeMessage, @@ -423,7 +482,7 @@ func (s *StreamingResponseWriter) finalize() { }) // Add to response output - s.response.Output = append(s.response.Output, OutputItem{ + output[s.messageOutputIndex] = OutputItem{ ID: s.currentItemID, Type: ItemTypeMessage, Role: "assistant", @@ -433,7 +492,7 @@ func (s *StreamingResponseWriter) finalize() { Annotations: []Annotation{}, }}, Status: StatusCompleted, - }) + } s.response.OutputText = finalText } @@ -444,7 +503,7 @@ func (s *StreamingResponseWriter) finalize() { Type: EventFunctionCallArgsDone, SequenceNumber: s.nextSeq(), ItemID: tc.ID, - OutputIndex: i, + OutputIndex: s.toolCallOutputIndex(i), Delta: tc.Arguments, }) @@ -453,13 +512,14 @@ func (s *StreamingResponseWriter) finalize() { s.sendEvent(EventOutputItemDone, &StreamEvent{ Type: EventOutputItemDone, SequenceNumber: s.nextSeq(), - OutputIndex: i, + OutputIndex: s.toolCallOutputIndex(i), Item: &tc, }) // Add to response output - s.response.Output = append(s.response.Output, tc) + output[s.toolCallOutputIndex(i)] = tc } + s.response.Output = append(s.response.Output, output...) // Update response status s.response.Status = StatusCompleted diff --git a/pkg/responses/transform.go b/pkg/responses/transform.go index cae624ee0..2b52d3f85 100644 --- a/pkg/responses/transform.go +++ b/pkg/responses/transform.go @@ -75,6 +75,9 @@ type ChatFunction struct { // ChatToolCall represents a tool call in chat format. type ChatToolCall struct { + // Index identifies the tool call in streaming deltas. Only the first + // delta of a call carries its ID; later argument deltas carry only Index. + Index *int `json:"index,omitempty"` ID string `json:"id"` Type string `json:"type"` Function ChatFunctionCall `json:"function"`