Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
180 changes: 180 additions & 0 deletions pkg/responses/handler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand Down
Loading
Loading