openai: finalize responses at the web search limit (#18328)

This commit is contained in:
Parth Sareen
2026-09-08 16:27:20 -07:00
committed by GitHub
parent 9160b3c0b8
commit cd1c5a145d
2 changed files with 193 additions and 30 deletions
+9 -1
View File
@@ -10,6 +10,7 @@ import (
"log/slog"
"math/rand"
"net/http"
"slices"
"strings"
"time"
@@ -863,9 +864,16 @@ func (w *WebSearchResponsesWriter) runLoop(ctx context.Context, initial api.Chat
}
calls = append(calls, responseCall)
resultContent := formatResponsesWebSearchResults(searchResponse.Results)
if loop == maxWebSearchLoops {
tools = slices.DeleteFunc(tools, func(tool api.Tool) bool {
return tool.Function.Name == "web_search"
})
resultContent += "\nThe web search limit for this response has been reached. Continue using the available results and state any limitations."
}
messages = append(messages,
buildWebSearchAssistantMessage(current, currentCall),
api.Message{Role: "tool", ToolCallID: currentCall.ID, Content: formatResponsesWebSearchResults(searchResponse.Results)},
api.Message{Role: "tool", ToolCallID: currentCall.ID, Content: resultContent},
)
var followUp api.ChatResponse
var followUpOutputStreamed bool
+184 -29
View File
@@ -3,8 +3,10 @@ package middleware
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"reflect"
"strings"
"testing"
@@ -684,38 +686,191 @@ func TestWebSearchResponsesWriterStreamingSplitInitialToolsDoNotLatchFinalText(t
}
}
func TestWebSearchResponsesWriterStreamingLoopExhaustionFails(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
stream := true
request := openai.ResponsesRequest{Model: "test-model", Stream: &stream, Tools: []openai.ResponsesTool{{Type: "web_search"}}}
inner := &ResponsesWriter{BaseWriter: BaseWriter{ResponseWriter: ctx.Writer}, converter: openai.NewResponsesStreamConverter("resp_test", "msg_test", request.Model, request), model: request.Model, stream: true, responseID: "resp_test", itemID: "msg_test", request: request}
searches := 0
writer := &WebSearchResponsesWriter{
BaseWriter: BaseWriter{ResponseWriter: ctx.Writer}, inner: inner, req: request,
chat: &api.ChatRequest{Model: request.Model, Tools: api.Tools{openai.WebSearchFunctionTool()}},
search: func(context.Context, string) (*api.WebSearchResponse, error) {
searches++
return &api.WebSearchResponse{}, nil
},
followUpStream: func(_ context.Context, _ []api.Message, _ api.Tools, yield func(api.ChatResponse) error) error {
call := api.ToolCall{ID: "call_next", Function: api.ToolCallFunction{Name: "web_search", Arguments: testArgs(map[string]any{"query": "again"})}}
if err := yield(api.ChatResponse{Message: api.Message{Role: "assistant", ToolCalls: []api.ToolCall{call}}}); err != nil {
return err
func TestWebSearchResponsesWriterFinalizesAtSearchLimit(t *testing.T) {
for _, test := range []struct {
name string
stream bool
clientTool bool
}{
{name: "non-streaming"},
{name: "streaming", stream: true},
{name: "non-streaming client tool", clientTool: true},
{name: "streaming client tool", stream: true, clientTool: true},
} {
t.Run(test.name, func(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
request := openai.ResponsesRequest{Model: "test-model", Stream: &test.stream, Tools: []openai.ResponsesTool{{Type: "web_search"}}}
chat := &api.ChatRequest{Model: request.Model, Tools: api.Tools{openai.WebSearchFunctionTool()}}
if test.clientTool {
chat.Tools = append(chat.Tools, api.Tool{Type: "function", Function: api.ToolFunction{Name: "get_weather"}})
}
return yield(api.ChatResponse{Done: true})
},
originalTools := append(api.Tools(nil), chat.Tools...)
inner := &ResponsesWriter{BaseWriter: BaseWriter{ResponseWriter: ctx.Writer}, converter: openai.NewResponsesStreamConverter("resp_test", "msg_test", request.Model, request), model: request.Model, stream: test.stream, responseID: "resp_test", itemID: "msg_test", request: request}
searches, followUps := 0, 0
searchCall := api.ToolCall{ID: "call_search", Function: api.ToolCallFunction{Name: "web_search", Arguments: testArgs(map[string]any{"query": "weather"})}}
writer := &WebSearchResponsesWriter{
BaseWriter: BaseWriter{ResponseWriter: ctx.Writer}, inner: inner, req: request, chat: chat,
search: func(context.Context, string) (*api.WebSearchResponse, error) {
searches++
return &api.WebSearchResponse{Results: []api.WebSearchResult{{Title: "Forecast", URL: "https://example.com/weather", Content: "Rain expected."}}}, nil
},
}
followUp := func(_ context.Context, messages []api.Message, tools api.Tools) (api.ChatResponse, error) {
followUps++
if len(messages) != searches*2 {
t.Fatalf("follow-up messages = %d, want %d", len(messages), searches*2)
}
for i := 1; i < len(messages); i += 2 {
if messages[i].Role != "tool" || messages[i].ToolCallID != searchCall.ID || !strings.Contains(messages[i].Content, "Rain expected.") {
t.Fatalf("search result %d lost: %#v", i, messages[i])
}
}
response := api.ChatResponse{Done: true, Message: api.Message{Role: "assistant"}, Metrics: api.Metrics{PromptEvalCount: 5, PromptEvalCachedCount: testIntPtr(2), EvalCount: 3}}
if searches < maxWebSearchLoops {
if !reflect.DeepEqual(tools, originalTools) {
t.Fatalf("tools removed before search limit: %#v", tools)
}
response.Message.ToolCalls = []api.ToolCall{searchCall}
return response, nil
}
if len(tools) != len(originalTools)-1 {
t.Fatalf("final follow-up tools = %#v, want only client tools", tools)
}
if !strings.Contains(messages[len(messages)-1].Content, "web search limit") {
t.Fatalf("final search result does not explain the limit: %#v", messages[len(messages)-1])
}
if test.clientTool {
if !reflect.DeepEqual(tools[0], originalTools[1]) {
t.Fatalf("client tool changed: %#v", tools)
}
response.Message.ToolCalls = []api.ToolCall{{ID: "call_weather", Function: api.ToolCallFunction{Name: "get_weather", Arguments: testArgs(map[string]any{"city": "SF"})}}}
} else {
response.Message.Content = "Rain expected, based on the available results."
}
return response, nil
}
writer.followUpChat = followUp
writer.followUpStream = func(ctx context.Context, messages []api.Message, tools api.Tools, yield func(api.ChatResponse) error) error {
response, err := followUp(ctx, messages, tools)
if err != nil {
return err
}
if err := yield(api.ChatResponse{Message: response.Message}); err != nil {
return err
}
return yield(api.ChatResponse{Done: true, Metrics: response.Metrics})
}
initial := api.ChatResponse{Done: true, Message: api.Message{Role: "assistant", ToolCalls: []api.ToolCall{searchCall}}, Metrics: api.Metrics{PromptEvalCount: 5, PromptEvalCachedCount: testIntPtr(2), EvalCount: 3}}
data, _ := json.Marshal(initial)
if _, err := writer.Write(data); err != nil {
t.Fatal(err)
}
if searches != maxWebSearchLoops || followUps != maxWebSearchLoops {
t.Fatalf("searches=%d follow-ups=%d, want %d each", searches, followUps, maxWebSearchLoops)
}
if !reflect.DeepEqual(chat.Tools, originalTools) {
t.Fatalf("original request tools mutated: %#v", chat.Tools)
}
body := recorder.Body.String()
var response openai.ResponsesResponse
if test.stream {
if strings.Count(body, "event: response.completed\n") != 1 || strings.Contains(body, "event: response.failed") {
t.Fatalf("unexpected terminal event: %s", body)
}
for _, block := range strings.Split(body, "\n\n") {
if data, ok := strings.CutPrefix(block, "event: response.completed\ndata: "); ok {
var event struct {
Response openai.ResponsesResponse `json:"response"`
}
if err := json.Unmarshal([]byte(data), &event); err != nil {
t.Fatal(err)
}
response = event.Response
}
}
} else if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil {
t.Fatal(err)
}
if recorder.Code != http.StatusOK || response.Status != "completed" || len(response.Output) != maxWebSearchLoops+1 {
t.Fatalf("unexpected final response: %s", body)
}
for _, item := range response.Output[:maxWebSearchLoops] {
if item.Type != "web_search_call" || item.Status != "completed" {
t.Fatalf("search output lost: %#v", item)
}
}
final := response.Output[maxWebSearchLoops]
if test.clientTool {
if final.Type != "function_call" || final.Name != "get_weather" || final.CallID != "call_weather" {
t.Fatalf("client tool output lost: %#v", final)
}
} else if final.Type != "message" || len(final.Content) != 1 || final.Content[0].Text != "Rain expected, based on the available results." {
t.Fatalf("final answer lost: %#v", final)
}
if response.Usage == nil || response.Usage.InputTokens != 20 || response.Usage.OutputTokens != 12 || response.Usage.InputTokensDetails.CachedTokens != 8 {
t.Fatalf("usage = %#v, want all four model responses", response.Usage)
}
})
}
}
initial := api.ChatResponse{Done: true, Message: api.Message{ToolCalls: []api.ToolCall{{ID: "call_1", Function: api.ToolCallFunction{Name: "web_search", Arguments: testArgs(map[string]any{"query": "first"})}}}}}
data, _ := json.Marshal(initial)
if _, err := writer.Write(data); err != nil {
t.Fatal(err)
}
body := recorder.Body.String()
if searches != maxWebSearchLoops || !strings.Contains(body, "response.failed") || !strings.Contains(body, "exceeded the maximum") || strings.Contains(body, "response.completed") {
t.Fatalf("loop exhaustion was not a terminal failure: searches=%d body=%s", searches, body)
func TestWebSearchResponsesWriterFinalizationFailure(t *testing.T) {
for _, stream := range []bool{false, true} {
for _, test := range []struct {
name string
err error
status int
message string
}{
{name: "model requests another search", status: http.StatusBadGateway, message: "web_search exceeded the maximum"},
{name: "model request fails", err: api.StatusError{StatusCode: http.StatusServiceUnavailable, ErrorMessage: "model unavailable"}, status: http.StatusServiceUnavailable, message: "model unavailable"},
} {
t.Run(fmt.Sprintf("%s/stream=%t", test.name, stream), func(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
request := openai.ResponsesRequest{Model: "test-model", Stream: &stream, Tools: []openai.ResponsesTool{{Type: "web_search"}}}
inner := &ResponsesWriter{BaseWriter: BaseWriter{ResponseWriter: ctx.Writer}, converter: openai.NewResponsesStreamConverter("resp_test", "msg_test", request.Model, request), model: request.Model, stream: stream, responseID: "resp_test", itemID: "msg_test", request: request}
searches, followUps := 0, 0
call := api.ToolCall{ID: "call_search", Function: api.ToolCallFunction{Name: "web_search", Arguments: testArgs(map[string]any{"query": "again"})}}
writer := &WebSearchResponsesWriter{
BaseWriter: BaseWriter{ResponseWriter: ctx.Writer}, inner: inner, req: request,
chat: &api.ChatRequest{Model: request.Model, Tools: api.Tools{openai.WebSearchFunctionTool()}},
search: func(context.Context, string) (*api.WebSearchResponse, error) {
searches++
return &api.WebSearchResponse{}, nil
},
followUpChat: func(context.Context, []api.Message, api.Tools) (api.ChatResponse, error) {
followUps++
if followUps == maxWebSearchLoops && test.err != nil {
return api.ChatResponse{}, test.err
}
return api.ChatResponse{Done: true, Message: api.Message{Role: "assistant", ToolCalls: []api.ToolCall{call}}}, nil
},
}
initial := api.ChatResponse{Done: true, Message: api.Message{Role: "assistant", ToolCalls: []api.ToolCall{call}}}
data, _ := json.Marshal(initial)
if _, err := writer.Write(data); err != nil {
t.Fatal(err)
}
if searches != maxWebSearchLoops || followUps != maxWebSearchLoops {
t.Fatalf("searches=%d follow-ups=%d, want %d each", searches, followUps, maxWebSearchLoops)
}
body := recorder.Body.String()
if !strings.Contains(body, test.message) {
t.Fatalf("finalization error lost: %s", body)
}
if stream {
if recorder.Code != http.StatusOK || strings.Count(body, "event: response.failed\n") != 1 || strings.Contains(body, "event: response.completed") {
t.Fatalf("expected a terminal failure: %s", body)
}
} else if recorder.Code != test.status {
t.Fatalf("status=%d, want %d: %s", recorder.Code, test.status, body)
}
})
}
}
}