anthropic: enable websearch (#14246)

This commit is contained in:
Parth Sareen
2026-02-13 19:20:46 -08:00
committed by GitHub
parent f0a07a353b
commit 5f5ef20131
6 changed files with 4276 additions and 31 deletions
+817 -17
View File
@@ -2,15 +2,22 @@ package middleware
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/ollama/ollama/anthropic"
"github.com/ollama/ollama/api"
"github.com/ollama/ollama/envconfig"
internalcloud "github.com/ollama/ollama/internal/cloud"
"github.com/ollama/ollama/logutil"
)
// AnthropicWriter wraps the response writer to transform Ollama responses to Anthropic format
@@ -18,7 +25,6 @@ type AnthropicWriter struct {
BaseWriter
stream bool
id string
model string
converter *anthropic.StreamConverter
}
@@ -31,7 +37,7 @@ func (w *AnthropicWriter) writeError(data []byte) (int, error) {
}
w.ResponseWriter.Header().Set("Content-Type", "application/json")
err := json.NewEncoder(w.ResponseWriter).Encode(anthropic.NewError(w.ResponseWriter.Status(), errData.Error))
err := json.NewEncoder(w.ResponseWriter).Encode(anthropic.NewError(w.Status(), errData.Error))
if err != nil {
return 0, err
}
@@ -40,18 +46,7 @@ func (w *AnthropicWriter) writeError(data []byte) (int, error) {
}
func (w *AnthropicWriter) writeEvent(eventType string, data any) error {
d, err := json.Marshal(data)
if err != nil {
return err
}
_, err = w.ResponseWriter.Write([]byte(fmt.Sprintf("event: %s\ndata: %s\n\n", eventType, d)))
if err != nil {
return err
}
if f, ok := w.ResponseWriter.(http.Flusher); ok {
f.Flush()
}
return nil
return writeSSE(w.ResponseWriter, eventType, data)
}
func (w *AnthropicWriter) writeResponse(data []byte) (int, error) {
@@ -65,6 +60,7 @@ func (w *AnthropicWriter) writeResponse(data []byte) (int, error) {
w.ResponseWriter.Header().Set("Content-Type", "text/event-stream")
events := w.converter.Process(chatResponse)
logutil.Trace("anthropic middleware: stream chunk", "resp", anthropic.TraceChatResponse(chatResponse), "events", len(events))
for _, event := range events {
if err := w.writeEvent(event.Event, event.Data); err != nil {
return 0, err
@@ -75,6 +71,7 @@ func (w *AnthropicWriter) writeResponse(data []byte) (int, error) {
w.ResponseWriter.Header().Set("Content-Type", "application/json")
response := anthropic.ToMessagesResponse(w.id, chatResponse)
logutil.Trace("anthropic middleware: converted response", "resp", anthropic.TraceMessagesResponse(response))
return len(data), json.NewEncoder(w.ResponseWriter).Encode(response)
}
@@ -87,9 +84,743 @@ func (w *AnthropicWriter) Write(data []byte) (int, error) {
return w.writeResponse(data)
}
// WebSearchAnthropicWriter intercepts responses containing web_search tool calls,
// executes the search, re-invokes the model with results, and assembles the
// Anthropic-format response (server_tool_use + web_search_tool_result + text).
type WebSearchAnthropicWriter struct {
BaseWriter
newLoopContext func() (context.Context, context.CancelFunc)
inner *AnthropicWriter
req anthropic.MessagesRequest // original Anthropic request
chatReq *api.ChatRequest // converted Ollama request (for followup calls)
stream bool
estimatedInputTokens int
terminalSent bool
observedPromptEvalCount int
observedEvalCount int
loopInFlight bool
loopBaseInputTok int
loopBaseOutputTok int
loopResultCh chan webSearchLoopResult
streamMessageStarted bool
streamHasOpenBlock bool
streamOpenBlockIndex int
streamNextIndex int
}
const maxWebSearchLoops = 3
type webSearchLoopResult struct {
response anthropic.MessagesResponse
loopErr *webSearchLoopError
}
type webSearchLoopError struct {
code string
query string
usage anthropic.Usage
err error
}
func (e *webSearchLoopError) Error() string {
if e.err == nil {
return e.code
}
return fmt.Sprintf("%s: %v", e.code, e.err)
}
func (w *WebSearchAnthropicWriter) Write(data []byte) (int, error) {
if w.terminalSent {
return len(data), nil
}
code := w.Status()
if code != http.StatusOK {
return w.inner.writeError(data)
}
var chatResponse api.ChatResponse
if err := json.Unmarshal(data, &chatResponse); err != nil {
return 0, err
}
w.recordObservedUsage(chatResponse.Metrics)
if w.stream && w.loopInFlight {
if !chatResponse.Done {
return len(data), nil
}
if err := w.writeLoopResult(); err != nil {
return len(data), err
}
return len(data), nil
}
webSearchCall, hasWebSearch, hasOtherTools := findWebSearchToolCall(chatResponse.Message.ToolCalls)
logutil.Trace("anthropic middleware: upstream chunk",
"resp", anthropic.TraceChatResponse(chatResponse),
"web_search", hasWebSearch,
"other_tools", hasOtherTools,
)
if hasWebSearch && hasOtherTools {
// Prefer web_search if both server and client tools are present in one chunk.
slog.Debug("preferring web_search tool call over client tool calls in mixed tool response")
}
if !hasWebSearch {
if w.stream {
if err := w.writePassthroughStreamChunk(chatResponse); err != nil {
return 0, err
}
return len(data), nil
}
return w.inner.writeResponse(data)
}
if w.stream {
// Let the original generation continue to completion while web search runs in parallel.
logutil.Trace("anthropic middleware: starting async web_search loop",
"tool_call", anthropic.TraceToolCall(webSearchCall),
"resp", anthropic.TraceChatResponse(chatResponse),
)
w.startLoopWorker(chatResponse, webSearchCall)
if chatResponse.Done {
if err := w.writeLoopResult(); err != nil {
return len(data), err
}
}
return len(data), nil
}
loopCtx, cancel := w.startLoopContext()
defer cancel()
initialUsage := anthropic.Usage{
InputTokens: max(w.observedPromptEvalCount, chatResponse.Metrics.PromptEvalCount),
OutputTokens: max(w.observedEvalCount, chatResponse.Metrics.EvalCount),
}
logutil.Trace("anthropic middleware: starting sync web_search loop",
"tool_call", anthropic.TraceToolCall(webSearchCall),
"resp", anthropic.TraceChatResponse(chatResponse),
"usage", initialUsage,
)
response, loopErr := w.runWebSearchLoop(loopCtx, chatResponse, webSearchCall, initialUsage)
if loopErr != nil {
return len(data), w.sendError(loopErr.code, loopErr.query, loopErr.usage)
}
if err := w.writeTerminalResponse(response); err != nil {
return 0, err
}
return len(data), nil
}
func (w *WebSearchAnthropicWriter) runWebSearchLoop(ctx context.Context, initialResponse api.ChatResponse, initialToolCall api.ToolCall, initialUsage anthropic.Usage) (anthropic.MessagesResponse, *webSearchLoopError) {
followUpMessages := make([]api.Message, 0, len(w.chatReq.Messages)+maxWebSearchLoops*2)
followUpMessages = append(followUpMessages, w.chatReq.Messages...)
followUpTools := append(api.Tools(nil), w.chatReq.Tools...)
usage := initialUsage
logutil.TraceContext(ctx, "anthropic middleware: web_search loop init",
"model", w.req.Model,
"tool_call", anthropic.TraceToolCall(initialToolCall),
"messages", len(followUpMessages),
"tools", len(followUpTools),
"max_loops", maxWebSearchLoops,
)
currentResponse := initialResponse
currentToolCall := initialToolCall
var serverContent []anthropic.ContentBlock
if !isCloudModelName(w.req.Model) {
logutil.TraceContext(ctx, "anthropic middleware: web_search execution blocked", "reason", "non_cloud_model")
return anthropic.MessagesResponse{}, &webSearchLoopError{
code: "web_search_not_supported_for_local_models",
query: extractQueryFromToolCall(&initialToolCall),
usage: usage,
}
}
for loop := 1; loop <= maxWebSearchLoops; loop++ {
query := extractQueryFromToolCall(&currentToolCall)
logutil.TraceContext(ctx, "anthropic middleware: web_search loop iteration",
"loop", loop,
"query", anthropic.TraceTruncateString(query),
"messages", len(followUpMessages),
)
if query == "" {
return anthropic.MessagesResponse{}, &webSearchLoopError{
code: "invalid_request",
query: "",
usage: usage,
}
}
const defaultMaxResults = 5
searchResp, err := anthropic.WebSearch(ctx, query, defaultMaxResults)
if err != nil {
logutil.TraceContext(ctx, "anthropic middleware: web_search request failed",
"loop", loop,
"query", query,
"error", err,
)
return anthropic.MessagesResponse{}, &webSearchLoopError{
code: "unavailable",
query: query,
usage: usage,
err: err,
}
}
logutil.TraceContext(ctx, "anthropic middleware: web_search results",
"loop", loop,
"results", len(searchResp.Results),
)
toolUseID := loopServerToolUseID(w.inner.id, loop)
searchResults := anthropic.ConvertOllamaToAnthropicResults(searchResp)
serverContent = append(serverContent,
anthropic.ContentBlock{
Type: "server_tool_use",
ID: toolUseID,
Name: "web_search",
Input: map[string]any{"query": query},
},
anthropic.ContentBlock{
Type: "web_search_tool_result",
ToolUseID: toolUseID,
Content: searchResults,
},
)
assistantMsg := buildWebSearchAssistantMessage(currentResponse, currentToolCall)
toolResultMsg := api.Message{
Role: "tool",
Content: formatWebSearchResultsForToolMessage(searchResp.Results),
ToolCallID: currentToolCall.ID,
}
followUpMessages = append(followUpMessages, assistantMsg, toolResultMsg)
followUpResponse, err := w.callFollowUpChat(ctx, followUpMessages, followUpTools)
if err != nil {
logutil.TraceContext(ctx, "anthropic middleware: followup /api/chat failed",
"loop", loop,
"query", query,
"error", err,
)
return anthropic.MessagesResponse{}, &webSearchLoopError{
code: "api_error",
query: query,
usage: usage,
err: err,
}
}
logutil.TraceContext(ctx, "anthropic middleware: followup response",
"loop", loop,
"resp", anthropic.TraceChatResponse(followUpResponse),
)
usage.InputTokens += followUpResponse.Metrics.PromptEvalCount
usage.OutputTokens += followUpResponse.Metrics.EvalCount
nextToolCall, hasWebSearch, hasOtherTools := findWebSearchToolCall(followUpResponse.Message.ToolCalls)
if hasWebSearch && hasOtherTools {
// Prefer web_search if both server and client tools are present in one chunk.
slog.Debug("preferring web_search tool call over client tool calls in mixed followup response")
}
if !hasWebSearch {
finalResponse := w.combineServerAndFinalContent(serverContent, followUpResponse, usage)
logutil.TraceContext(ctx, "anthropic middleware: web_search loop complete",
"loop", loop,
"resp", anthropic.TraceMessagesResponse(finalResponse),
)
return finalResponse, nil
}
currentResponse = followUpResponse
currentToolCall = nextToolCall
}
maxLoopQuery := extractQueryFromToolCall(&currentToolCall)
maxLoopToolUseID := loopServerToolUseID(w.inner.id, maxWebSearchLoops+1)
serverContent = append(serverContent,
anthropic.ContentBlock{
Type: "server_tool_use",
ID: maxLoopToolUseID,
Name: "web_search",
Input: map[string]any{"query": maxLoopQuery},
},
anthropic.ContentBlock{
Type: "web_search_tool_result",
ToolUseID: maxLoopToolUseID,
Content: anthropic.WebSearchToolResultError{
Type: "web_search_tool_result_error",
ErrorCode: "max_uses_exceeded",
},
},
)
maxResponse := anthropic.MessagesResponse{
ID: w.inner.id,
Type: "message",
Role: "assistant",
Model: w.req.Model,
Content: serverContent,
StopReason: "end_turn",
Usage: usage,
}
logutil.TraceContext(ctx, "anthropic middleware: web_search loop max reached",
"resp", anthropic.TraceMessagesResponse(maxResponse),
)
return maxResponse, nil
}
func (w *WebSearchAnthropicWriter) startLoopWorker(initialResponse api.ChatResponse, initialToolCall api.ToolCall) {
if w.loopInFlight {
return
}
initialUsage := anthropic.Usage{
InputTokens: max(w.observedPromptEvalCount, initialResponse.Metrics.PromptEvalCount),
OutputTokens: max(w.observedEvalCount, initialResponse.Metrics.EvalCount),
}
w.loopBaseInputTok = initialUsage.InputTokens
w.loopBaseOutputTok = initialUsage.OutputTokens
w.loopResultCh = make(chan webSearchLoopResult, 1)
w.loopInFlight = true
logutil.Trace("anthropic middleware: loop worker started",
"usage", initialUsage,
"tool_call", anthropic.TraceToolCall(initialToolCall),
)
go func() {
ctx, cancel := w.startLoopContext()
defer cancel()
response, loopErr := w.runWebSearchLoop(ctx, initialResponse, initialToolCall, initialUsage)
w.loopResultCh <- webSearchLoopResult{
response: response,
loopErr: loopErr,
}
}()
}
func (w *WebSearchAnthropicWriter) writeLoopResult() error {
if w.loopResultCh == nil {
return w.sendError("api_error", "", w.currentObservedUsage())
}
result := <-w.loopResultCh
w.loopResultCh = nil
w.loopInFlight = false
if result.loopErr != nil {
logutil.Trace("anthropic middleware: loop worker returned error",
"code", result.loopErr.code,
"query", result.loopErr.query,
"usage", result.loopErr.usage,
"error", result.loopErr.err,
)
usage := result.loopErr.usage
w.applyObservedUsageDeltaToUsage(&usage)
return w.sendError(result.loopErr.code, result.loopErr.query, usage)
}
logutil.Trace("anthropic middleware: loop worker done", "resp", anthropic.TraceMessagesResponse(result.response))
w.applyObservedUsageDelta(&result.response)
return w.writeTerminalResponse(result.response)
}
func (w *WebSearchAnthropicWriter) applyObservedUsageDelta(response *anthropic.MessagesResponse) {
w.applyObservedUsageDeltaToUsage(&response.Usage)
}
func (w *WebSearchAnthropicWriter) recordObservedUsage(metrics api.Metrics) {
if metrics.PromptEvalCount > w.observedPromptEvalCount {
w.observedPromptEvalCount = metrics.PromptEvalCount
}
if metrics.EvalCount > w.observedEvalCount {
w.observedEvalCount = metrics.EvalCount
}
}
func (w *WebSearchAnthropicWriter) applyObservedUsageDeltaToUsage(usage *anthropic.Usage) {
if deltaIn := w.observedPromptEvalCount - w.loopBaseInputTok; deltaIn > 0 {
usage.InputTokens += deltaIn
}
if deltaOut := w.observedEvalCount - w.loopBaseOutputTok; deltaOut > 0 {
usage.OutputTokens += deltaOut
}
}
func (w *WebSearchAnthropicWriter) currentObservedUsage() anthropic.Usage {
return anthropic.Usage{
InputTokens: w.observedPromptEvalCount,
OutputTokens: w.observedEvalCount,
}
}
func (w *WebSearchAnthropicWriter) startLoopContext() (context.Context, context.CancelFunc) {
if w.newLoopContext != nil {
return w.newLoopContext()
}
return context.WithTimeout(context.Background(), 5*time.Minute)
}
func (w *WebSearchAnthropicWriter) combineServerAndFinalContent(serverContent []anthropic.ContentBlock, finalResponse api.ChatResponse, usage anthropic.Usage) anthropic.MessagesResponse {
converted := anthropic.ToMessagesResponse(w.inner.id, finalResponse)
content := make([]anthropic.ContentBlock, 0, len(serverContent)+len(converted.Content))
content = append(content, serverContent...)
content = append(content, converted.Content...)
return anthropic.MessagesResponse{
ID: w.inner.id,
Type: "message",
Role: "assistant",
Model: w.req.Model,
Content: content,
StopReason: converted.StopReason,
StopSequence: converted.StopSequence,
Usage: usage,
}
}
func buildWebSearchAssistantMessage(response api.ChatResponse, webSearchCall api.ToolCall) api.Message {
assistantMsg := api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{webSearchCall},
}
if response.Message.Content != "" {
assistantMsg.Content = response.Message.Content
}
if response.Message.Thinking != "" {
assistantMsg.Thinking = response.Message.Thinking
}
return assistantMsg
}
func formatWebSearchResultsForToolMessage(results []anthropic.OllamaWebSearchResult) string {
var resultText strings.Builder
for _, r := range results {
fmt.Fprintf(&resultText, "Title: %s\nURL: %s\n", r.Title, r.URL)
if r.Content != "" {
fmt.Fprintf(&resultText, "Content: %s\n", r.Content)
}
resultText.WriteString("\n")
}
return resultText.String()
}
func findWebSearchToolCall(toolCalls []api.ToolCall) (api.ToolCall, bool, bool) {
var webSearchCall api.ToolCall
hasWebSearch := false
hasOtherTools := false
for _, toolCall := range toolCalls {
if toolCall.Function.Name == "web_search" {
if !hasWebSearch {
webSearchCall = toolCall
hasWebSearch = true
}
continue
}
hasOtherTools = true
}
return webSearchCall, hasWebSearch, hasOtherTools
}
func loopServerToolUseID(messageID string, loop int) string {
base := serverToolUseID(messageID)
if loop <= 1 {
return base
}
return fmt.Sprintf("%s_%d", base, loop)
}
func (w *WebSearchAnthropicWriter) callFollowUpChat(ctx context.Context, messages []api.Message, tools api.Tools) (api.ChatResponse, error) {
streaming := false
followUp := api.ChatRequest{
Model: w.chatReq.Model,
Messages: messages,
Stream: &streaming,
Tools: tools,
Options: w.chatReq.Options,
}
body, err := json.Marshal(followUp)
if err != nil {
return api.ChatResponse{}, err
}
chatURL := envconfig.Host().String() + "/api/chat"
logutil.TraceContext(ctx, "anthropic middleware: followup request",
"url", chatURL,
"req", anthropic.TraceChatRequest(&followUp),
)
httpReq, err := http.NewRequestWithContext(ctx, "POST", chatURL, bytes.NewReader(body))
if err != nil {
return api.ChatResponse{}, err
}
httpReq.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(httpReq)
if err != nil {
return api.ChatResponse{}, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
respBody, _ := io.ReadAll(resp.Body)
logutil.TraceContext(ctx, "anthropic middleware: followup non-200 response",
"status", resp.StatusCode,
"response", strings.TrimSpace(string(respBody)),
)
return api.ChatResponse{}, fmt.Errorf("followup /api/chat returned status %d: %s", resp.StatusCode, strings.TrimSpace(string(respBody)))
}
var chatResp api.ChatResponse
if err := json.NewDecoder(resp.Body).Decode(&chatResp); err != nil {
return api.ChatResponse{}, err
}
logutil.TraceContext(ctx, "anthropic middleware: followup decoded", "resp", anthropic.TraceChatResponse(chatResp))
return chatResp, nil
}
func (w *WebSearchAnthropicWriter) writePassthroughStreamChunk(chatResponse api.ChatResponse) error {
events := w.inner.converter.Process(chatResponse)
for _, event := range events {
switch e := event.Data.(type) {
case anthropic.MessageStartEvent:
w.streamMessageStarted = true
case anthropic.ContentBlockStartEvent:
w.streamHasOpenBlock = true
w.streamOpenBlockIndex = e.Index
if e.Index+1 > w.streamNextIndex {
w.streamNextIndex = e.Index + 1
}
case anthropic.ContentBlockStopEvent:
if w.streamHasOpenBlock && w.streamOpenBlockIndex == e.Index {
w.streamHasOpenBlock = false
}
if e.Index+1 > w.streamNextIndex {
w.streamNextIndex = e.Index + 1
}
case anthropic.MessageStopEvent:
w.terminalSent = true
}
if err := writeSSE(w.ResponseWriter, event.Event, event.Data); err != nil {
return err
}
}
return nil
}
func (w *WebSearchAnthropicWriter) ensureStreamMessageStart(usage anthropic.Usage) error {
if w.streamMessageStarted {
return nil
}
inputTokens := usage.InputTokens
if inputTokens == 0 {
inputTokens = w.estimatedInputTokens
}
if err := writeSSE(w.ResponseWriter, "message_start", anthropic.MessageStartEvent{
Type: "message_start",
Message: anthropic.MessagesResponse{
ID: w.inner.id,
Type: "message",
Role: "assistant",
Model: w.req.Model,
Content: []anthropic.ContentBlock{},
Usage: anthropic.Usage{
InputTokens: inputTokens,
},
},
}); err != nil {
return err
}
w.streamMessageStarted = true
return nil
}
func (w *WebSearchAnthropicWriter) closeOpenStreamBlock() error {
if !w.streamHasOpenBlock {
return nil
}
if err := writeSSE(w.ResponseWriter, "content_block_stop", anthropic.ContentBlockStopEvent{
Type: "content_block_stop",
Index: w.streamOpenBlockIndex,
}); err != nil {
return err
}
if w.streamOpenBlockIndex+1 > w.streamNextIndex {
w.streamNextIndex = w.streamOpenBlockIndex + 1
}
w.streamHasOpenBlock = false
return nil
}
func (w *WebSearchAnthropicWriter) writeStreamContentBlocks(content []anthropic.ContentBlock) error {
for _, block := range content {
index := w.streamNextIndex
if block.Type == "text" {
emptyText := ""
if err := writeSSE(w.ResponseWriter, "content_block_start", anthropic.ContentBlockStartEvent{
Type: "content_block_start",
Index: index,
ContentBlock: anthropic.ContentBlock{
Type: "text",
Text: &emptyText,
},
}); err != nil {
return err
}
text := ""
if block.Text != nil {
text = *block.Text
}
if err := writeSSE(w.ResponseWriter, "content_block_delta", anthropic.ContentBlockDeltaEvent{
Type: "content_block_delta",
Index: index,
Delta: anthropic.Delta{
Type: "text_delta",
Text: text,
},
}); err != nil {
return err
}
} else {
if err := writeSSE(w.ResponseWriter, "content_block_start", anthropic.ContentBlockStartEvent{
Type: "content_block_start",
Index: index,
ContentBlock: block,
}); err != nil {
return err
}
}
if err := writeSSE(w.ResponseWriter, "content_block_stop", anthropic.ContentBlockStopEvent{
Type: "content_block_stop",
Index: index,
}); err != nil {
return err
}
w.streamNextIndex++
}
return nil
}
func (w *WebSearchAnthropicWriter) writeTerminalResponse(response anthropic.MessagesResponse) error {
if w.terminalSent {
return nil
}
if !w.stream {
w.ResponseWriter.Header().Set("Content-Type", "application/json")
if err := json.NewEncoder(w.ResponseWriter).Encode(response); err != nil {
return err
}
w.terminalSent = true
return nil
}
if err := w.ensureStreamMessageStart(response.Usage); err != nil {
return err
}
if err := w.closeOpenStreamBlock(); err != nil {
return err
}
if err := w.writeStreamContentBlocks(response.Content); err != nil {
return err
}
if err := writeSSE(w.ResponseWriter, "message_delta", anthropic.MessageDeltaEvent{
Type: "message_delta",
Delta: anthropic.MessageDelta{
StopReason: response.StopReason,
},
Usage: anthropic.DeltaUsage{
InputTokens: response.Usage.InputTokens,
OutputTokens: response.Usage.OutputTokens,
},
}); err != nil {
return err
}
if err := writeSSE(w.ResponseWriter, "message_stop", anthropic.MessageStopEvent{
Type: "message_stop",
}); err != nil {
return err
}
w.terminalSent = true
return nil
}
// streamResponse emits a complete MessagesResponse as SSE events.
func (w *WebSearchAnthropicWriter) streamResponse(response anthropic.MessagesResponse) error {
return w.writeTerminalResponse(response)
}
func (w *WebSearchAnthropicWriter) webSearchErrorResponse(errorCode, query string, usage anthropic.Usage) anthropic.MessagesResponse {
toolUseID := serverToolUseID(w.inner.id)
return anthropic.MessagesResponse{
ID: w.inner.id,
Type: "message",
Role: "assistant",
Model: w.req.Model,
Content: []anthropic.ContentBlock{
{
Type: "server_tool_use",
ID: toolUseID,
Name: "web_search",
Input: map[string]any{"query": query},
},
{
Type: "web_search_tool_result",
ToolUseID: toolUseID,
Content: anthropic.WebSearchToolResultError{
Type: "web_search_tool_result_error",
ErrorCode: errorCode,
},
},
},
StopReason: "end_turn",
Usage: usage,
}
}
// sendError sends a web search error response.
func (w *WebSearchAnthropicWriter) sendError(errorCode, query string, usage anthropic.Usage) error {
response := w.webSearchErrorResponse(errorCode, query, usage)
logutil.Trace("anthropic middleware: web_search error", "code", errorCode, "query", query, "usage", usage)
return w.writeTerminalResponse(response)
}
// AnthropicMessagesMiddleware handles Anthropic Messages API requests
func AnthropicMessagesMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
requestCtx := c.Request.Context()
var req anthropic.MessagesRequest
err := c.ShouldBindJSON(&req)
if err != nil {
@@ -134,11 +865,10 @@ func AnthropicMessagesMiddleware() gin.HandlerFunc {
// Estimate input tokens for streaming (actual count not available until generation completes)
estimatedTokens := anthropic.EstimateInputTokens(req)
w := &AnthropicWriter{
innerWriter := &AnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: c.Writer},
stream: req.Stream,
id: messageID,
model: req.Model,
converter: anthropic.NewStreamConverter(messageID, req.Model, estimatedTokens),
}
@@ -148,8 +878,78 @@ func AnthropicMessagesMiddleware() gin.HandlerFunc {
c.Writer.Header().Set("Connection", "keep-alive")
}
c.Writer = w
if hasWebSearchTool(req.Tools) {
// Guard against runtime cloud-disable policy (OLLAMA_NO_CLOUD/server.json)
// for cloud models. Local models may still receive web_search tool definitions;
// execution is validated when the model actually emits a web_search tool call.
if isCloudModelName(req.Model) {
if disabled, _ := internalcloud.Status(); disabled {
c.AbortWithStatusJSON(http.StatusForbidden, anthropic.NewError(http.StatusForbidden, internalcloud.DisabledError("web search is unavailable")))
return
}
}
c.Writer = &WebSearchAnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: c.Writer},
newLoopContext: func() (context.Context, context.CancelFunc) {
return context.WithTimeout(requestCtx, 5*time.Minute)
},
inner: innerWriter,
req: req,
chatReq: chatReq,
stream: req.Stream,
estimatedInputTokens: estimatedTokens,
}
} else {
c.Writer = innerWriter
}
c.Next()
}
}
// hasWebSearchTool checks if the request tools include a web_search tool
func hasWebSearchTool(tools []anthropic.Tool) bool {
for _, tool := range tools {
if strings.HasPrefix(tool.Type, "web_search") {
return true
}
}
return false
}
func isCloudModelName(name string) bool {
return strings.HasSuffix(name, ":cloud") || strings.HasSuffix(name, "-cloud")
}
// extractQueryFromToolCall extracts the search query from a web_search tool call
func extractQueryFromToolCall(tc *api.ToolCall) string {
q, ok := tc.Function.Arguments.Get("query")
if !ok {
return ""
}
if s, ok := q.(string); ok {
return s
}
return ""
}
// writeSSE writes a Server-Sent Event
func writeSSE(w http.ResponseWriter, eventType string, data any) error {
d, err := json.Marshal(data)
if err != nil {
return err
}
if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", eventType, d); err != nil {
return err
}
if f, ok := w.(http.Flusher); ok {
f.Flush()
}
return nil
}
// serverToolUseID derives a server tool use ID from a message ID
func serverToolUseID(messageID string) string {
return "srvtoolu_" + strings.TrimPrefix(messageID, "msg_")
}
+2372
View File
@@ -605,3 +605,2375 @@ func TestAnthropicMessagesMiddleware_SetsRelaxThinkingFlag(t *testing.T) {
t.Error("expected relax_thinking flag to be set in context")
}
}
// Web Search Tests
func TestHasWebSearchTool(t *testing.T) {
tests := []struct {
name string
tools []anthropic.Tool
expected bool
}{
{
name: "no tools",
tools: nil,
expected: false,
},
{
name: "regular tool only",
tools: []anthropic.Tool{
{Type: "custom", Name: "get_weather"},
},
expected: false,
},
{
name: "web search tool",
tools: []anthropic.Tool{
{Type: "web_search_20250305", Name: "web_search"},
},
expected: true,
},
{
name: "mixed tools",
tools: []anthropic.Tool{
{Type: "custom", Name: "get_weather"},
{Type: "web_search_20250305", Name: "web_search"},
},
expected: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := hasWebSearchTool(tt.tools)
if result != tt.expected {
t.Errorf("expected %v, got %v", tt.expected, result)
}
})
}
}
func TestExtractQueryFromToolCall(t *testing.T) {
tests := []struct {
name string
tc *api.ToolCall
expected string
}{
{
name: "valid query",
tc: &api.ToolCall{
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "test search"),
},
},
expected: "test search",
},
{
name: "empty arguments",
tc: &api.ToolCall{
Function: api.ToolCallFunction{
Name: "web_search",
},
},
expected: "",
},
{
name: "no query key",
tc: &api.ToolCall{
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("other", "value"),
},
},
expected: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := extractQueryFromToolCall(tt.tc)
if result != tt.expected {
t.Errorf("expected %q, got %q", tt.expected, result)
}
})
}
}
// makeArgs is a test helper that creates ToolCallFunctionArguments
func makeArgs(key string, value any) api.ToolCallFunctionArguments {
args := api.NewToolCallFunctionArguments()
args.Set(key, value)
return args
}
// --- Web Search Integration Tests ---
// TestWebSearchServerToolUseID tests the ID derivation logic.
func TestWebSearchServerToolUseID(t *testing.T) {
tests := []struct {
msgID string
expected string
}{
{"msg_abc123", "srvtoolu_abc123"},
{"msg_", "srvtoolu_"},
{"nomsgprefix", "srvtoolu_nomsgprefix"},
}
for _, tt := range tests {
got := serverToolUseID(tt.msgID)
if got != tt.expected {
t.Errorf("serverToolUseID(%q) = %q, want %q", tt.msgID, got, tt.expected)
}
}
}
// TestWebSearchNoWebSearchTool verifies that when there is no web_search tool,
// requests pass through to the normal AnthropicWriter without interception.
func TestWebSearchNoWebSearchTool(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
Content: "Normal response",
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{"model":"test-model","max_tokens":100,"messages":[{"role":"user","content":"Hello"}]}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if result.Type != "message" {
t.Errorf("expected type 'message', got %q", result.Type)
}
if len(result.Content) != 1 || result.Content[0].Type != "text" {
t.Fatalf("expected single text block, got %d blocks", len(result.Content))
}
if *result.Content[0].Text != "Normal response" {
t.Errorf("expected text 'Normal response', got %q", *result.Content[0].Text)
}
}
// TestWebSearchToolPresent_ModelDoesNotCallIt_NonStreaming verifies that when
// the web_search tool is present but the model does not call it, the response
// passes through normally (non-streaming case).
func TestWebSearchToolPresent_ModelDoesNotCallIt_NonStreaming(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
Content: "I can answer that without searching.",
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 12, EvalCount: 8},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"What is 2+2?"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if result.Type != "message" {
t.Errorf("expected type 'message', got %q", result.Type)
}
if len(result.Content) != 1 || result.Content[0].Type != "text" {
t.Fatalf("expected single text block, got %+v", result.Content)
}
if *result.Content[0].Text != "I can answer that without searching." {
t.Errorf("unexpected text: %q", *result.Content[0].Text)
}
if result.StopReason != "end_turn" {
t.Errorf("expected stop_reason 'end_turn', got %q", result.StopReason)
}
}
// TestWebSearchToolPresent_ModelDoesNotCallIt_Streaming verifies the streaming
// pass-through case when the model does not invoke web_search.
func TestWebSearchToolPresent_ModelDoesNotCallIt_Streaming(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
// Simulate streaming: two partial chunks then a final chunk
chunks := []api.ChatResponse{
{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "Hello "},
Done: false,
},
{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "world"},
Done: false,
},
{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: ""},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
},
}
c.Writer.WriteHeader(http.StatusOK)
for _, chunk := range chunks {
data, _ := json.Marshal(chunk)
_, _ = c.Writer.Write(data)
}
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"stream":true,
"messages":[{"role":"user","content":"Hi"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
// Parse SSE events
events := parseSSEEvents(t, resp.Body.String())
// Should have standard streaming event flow
if len(events) == 0 {
t.Fatal("expected SSE events, got none")
}
// First event should be message_start
if events[0].event != "message_start" {
t.Errorf("first event should be message_start, got %q", events[0].event)
}
// Should have content_block_start for text
hasTextStart := false
hasTextDelta := false
hasMessageStop := false
for _, e := range events {
if e.event == "content_block_start" {
var cbs anthropic.ContentBlockStartEvent
if err := json.Unmarshal([]byte(e.data), &cbs); err == nil {
if cbs.ContentBlock.Type == "text" {
hasTextStart = true
}
}
}
if e.event == "content_block_delta" {
var cbd anthropic.ContentBlockDeltaEvent
if err := json.Unmarshal([]byte(e.data), &cbd); err == nil {
if cbd.Delta.Type == "text_delta" {
hasTextDelta = true
}
}
}
if e.event == "message_stop" {
hasMessageStop = true
}
}
if !hasTextStart {
t.Error("expected content_block_start with text type")
}
if !hasTextDelta {
t.Error("expected content_block_delta with text_delta")
}
if !hasMessageStop {
t.Error("expected message_stop event")
}
}
// TestWebSearchToolPresent_ModelCallsIt_NonStreaming tests the full web search flow
// in non-streaming mode. It mocks the followup /api/chat call using a local HTTP server.
func TestWebSearchToolPresent_ModelCallsIt_NonStreaming(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
// Create a mock Ollama server that responds to the followup /api/chat call
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
Content: "Based on my search, the answer is 42.",
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 50, EvalCount: 20},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
// Set OLLAMA_HOST to our mock server so the followup call goes there
t.Setenv("OLLAMA_HOST", followupServer.URL)
// Also mock the web search API
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Test Result", URL: "https://example.com/result", Content: "Some content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
// Point DoWebSearch at our mock search server
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_001",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "meaning of life"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 15, EvalCount: 3},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"What is the meaning of life?"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v\nbody: %s", err, resp.Body.String())
}
if result.Type != "message" {
t.Errorf("expected type 'message', got %q", result.Type)
}
if result.Role != "assistant" {
t.Errorf("expected role 'assistant', got %q", result.Role)
}
// Should have 3 blocks: server_tool_use + web_search_tool_result + text
if len(result.Content) != 3 {
t.Fatalf("expected 3 content blocks, got %d: %+v", len(result.Content), result.Content)
}
if result.Content[0].Type != "server_tool_use" {
t.Errorf("expected first block type 'server_tool_use', got %q", result.Content[0].Type)
}
if result.Content[0].Name != "web_search" {
t.Errorf("expected name 'web_search', got %q", result.Content[0].Name)
}
if result.Content[1].Type != "web_search_tool_result" {
t.Errorf("expected second block type 'web_search_tool_result', got %q", result.Content[1].Type)
}
if result.Content[1].ToolUseID != result.Content[0].ID {
t.Errorf("tool_use_id mismatch: %q != %q", result.Content[1].ToolUseID, result.Content[0].ID)
}
if result.Content[2].Type != "text" {
t.Errorf("expected third block type 'text', got %q", result.Content[2].Type)
}
if result.Content[2].Text == nil || *result.Content[2].Text == "" {
t.Error("expected non-empty text in third block")
}
if result.StopReason != "end_turn" {
t.Errorf("expected stop_reason 'end_turn', got %q", result.StopReason)
}
}
// TestWebSearchToolPresent_ModelCallsIt_Streaming tests the streaming SSE output
// when the model calls web_search with mocked search and followup endpoints.
func TestWebSearchToolPresent_ModelCallsIt_Streaming(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
// Mock followup /api/chat server
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "Here are the latest news."},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 40, EvalCount: 15},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
// Mock web search API
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "News Result", URL: "https://example.com/news", Content: "Breaking news"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
// Simulate buffered streaming: non-final chunk then final with tool call
chunks := []api.ChatResponse{
{
Model: "test-model",
Message: api.Message{Role: "assistant"},
Done: false,
},
{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_002",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "latest news"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 2},
},
}
c.Writer.WriteHeader(http.StatusOK)
for _, chunk := range chunks {
data, _ := json.Marshal(chunk)
_, _ = c.Writer.Write(data)
}
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"stream":true,
"messages":[{"role":"user","content":"What is the latest news?"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
events := parseSSEEvents(t, resp.Body.String())
// Success path: 10 events (3 blocks: server_tool_use, web_search_tool_result, text with delta)
expectedEventTypes := []string{
"message_start",
"content_block_start", // server_tool_use
"content_block_stop",
"content_block_start", // web_search_tool_result
"content_block_stop",
"content_block_start", // text (empty)
"content_block_delta", // text_delta with actual content
"content_block_stop",
"message_delta",
"message_stop",
}
if len(events) != len(expectedEventTypes) {
t.Fatalf("expected %d events, got %d.\nEvents: %v", len(expectedEventTypes), len(events), eventNames(events))
}
for i, expected := range expectedEventTypes {
if events[i].event != expected {
t.Errorf("event[%d]: expected %q, got %q", i, expected, events[i].event)
}
}
// Verify text delta has the followup model's content
var textDelta anthropic.ContentBlockDeltaEvent
if err := json.Unmarshal([]byte(events[6].data), &textDelta); err != nil {
t.Fatalf("failed to parse text delta: %v", err)
}
if textDelta.Delta.Type != "text_delta" {
t.Errorf("expected delta type 'text_delta', got %q", textDelta.Delta.Type)
}
if textDelta.Delta.Text != "Here are the latest news." {
t.Errorf("expected followup text, got %q", textDelta.Delta.Text)
}
}
// TestWebSearchStreamResponse tests the streamResponse method directly by constructing
// a WebSearchAnthropicWriter and calling streamResponse with a known response.
func TestWebSearchStreamResponse(t *testing.T) {
gin.SetMode(gin.TestMode)
text := "Here is the answer."
response := anthropic.MessagesResponse{
ID: "msg_test123",
Type: "message",
Role: "assistant",
Model: "test-model",
Content: []anthropic.ContentBlock{
{
Type: "server_tool_use",
ID: "srvtoolu_test123",
Name: "web_search",
Input: map[string]any{"query": "test query"},
},
{
Type: "web_search_tool_result",
ToolUseID: "srvtoolu_test123",
Content: []anthropic.WebSearchResult{
{Type: "web_search_result", URL: "https://example.com", Title: "Example"},
},
},
{
Type: "text",
Text: &text,
},
},
StopReason: "end_turn",
Usage: anthropic.Usage{InputTokens: 20, OutputTokens: 10},
}
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
innerWriter := &AnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: ginCtx.Writer},
stream: true,
id: "msg_test123",
}
wsWriter := &WebSearchAnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: ginCtx.Writer},
inner: innerWriter,
stream: true,
req: anthropic.MessagesRequest{Model: "test-model"},
}
if err := wsWriter.streamResponse(response); err != nil {
t.Fatalf("streamResponse error: %v", err)
}
events := parseSSEEvents(t, rec.Body.String())
// Verify full event sequence
expectedEventTypes := []string{
"message_start",
"content_block_start", // server_tool_use (index 0)
"content_block_stop", // index 0
"content_block_start", // web_search_tool_result (index 1)
"content_block_stop", // index 1
"content_block_start", // text (index 2)
"content_block_delta", // text_delta
"content_block_stop", // index 2
"message_delta",
"message_stop",
}
if len(events) != len(expectedEventTypes) {
t.Fatalf("expected %d events, got %d.\nEvents: %v", len(expectedEventTypes), len(events), eventNames(events))
}
for i, expected := range expectedEventTypes {
if events[i].event != expected {
t.Errorf("event[%d]: expected %q, got %q", i, expected, events[i].event)
}
}
// Verify message_start content
var msgStart anthropic.MessageStartEvent
if err := json.Unmarshal([]byte(events[0].data), &msgStart); err != nil {
t.Fatalf("failed to parse message_start: %v", err)
}
if msgStart.Message.ID != "msg_test123" {
t.Errorf("expected message ID 'msg_test123', got %q", msgStart.Message.ID)
}
if msgStart.Message.Role != "assistant" {
t.Errorf("expected role 'assistant', got %q", msgStart.Message.Role)
}
if len(msgStart.Message.Content) != 0 {
t.Errorf("expected empty content in message_start, got %d blocks", len(msgStart.Message.Content))
}
// Verify content_block_start for server_tool_use (event index 1)
var toolStart anthropic.ContentBlockStartEvent
if err := json.Unmarshal([]byte(events[1].data), &toolStart); err != nil {
t.Fatalf("failed to parse server_tool_use start: %v", err)
}
if toolStart.Index != 0 {
t.Errorf("expected index 0, got %d", toolStart.Index)
}
if toolStart.ContentBlock.Type != "server_tool_use" {
t.Errorf("expected type 'server_tool_use', got %q", toolStart.ContentBlock.Type)
}
if toolStart.ContentBlock.ID != "srvtoolu_test123" {
t.Errorf("expected ID 'srvtoolu_test123', got %q", toolStart.ContentBlock.ID)
}
// Verify content_block_start for web_search_tool_result (event index 3)
var searchStart anthropic.ContentBlockStartEvent
if err := json.Unmarshal([]byte(events[3].data), &searchStart); err != nil {
t.Fatalf("failed to parse web_search_tool_result start: %v", err)
}
if searchStart.Index != 1 {
t.Errorf("expected index 1, got %d", searchStart.Index)
}
if searchStart.ContentBlock.Type != "web_search_tool_result" {
t.Errorf("expected type 'web_search_tool_result', got %q", searchStart.ContentBlock.Type)
}
// Verify text block: content_block_start (event index 5)
var textStart anthropic.ContentBlockStartEvent
if err := json.Unmarshal([]byte(events[5].data), &textStart); err != nil {
t.Fatalf("failed to parse text start: %v", err)
}
if textStart.Index != 2 {
t.Errorf("expected index 2, got %d", textStart.Index)
}
if textStart.ContentBlock.Type != "text" {
t.Errorf("expected type 'text', got %q", textStart.ContentBlock.Type)
}
// Text in start should be empty
if textStart.ContentBlock.Text == nil || *textStart.ContentBlock.Text != "" {
t.Errorf("expected empty text in content_block_start, got %v", textStart.ContentBlock.Text)
}
// Verify text delta (event index 6)
var textDelta anthropic.ContentBlockDeltaEvent
if err := json.Unmarshal([]byte(events[6].data), &textDelta); err != nil {
t.Fatalf("failed to parse text delta: %v", err)
}
if textDelta.Index != 2 {
t.Errorf("expected index 2, got %d", textDelta.Index)
}
if textDelta.Delta.Type != "text_delta" {
t.Errorf("expected delta type 'text_delta', got %q", textDelta.Delta.Type)
}
if textDelta.Delta.Text != "Here is the answer." {
t.Errorf("expected delta text 'Here is the answer.', got %q", textDelta.Delta.Text)
}
// Verify message_delta (event index 8)
var msgDelta anthropic.MessageDeltaEvent
if err := json.Unmarshal([]byte(events[8].data), &msgDelta); err != nil {
t.Fatalf("failed to parse message_delta: %v", err)
}
if msgDelta.Delta.StopReason != "end_turn" {
t.Errorf("expected stop_reason 'end_turn', got %q", msgDelta.Delta.StopReason)
}
if msgDelta.Usage.InputTokens != 20 {
t.Errorf("expected input_tokens 20, got %d", msgDelta.Usage.InputTokens)
}
if msgDelta.Usage.OutputTokens != 10 {
t.Errorf("expected output_tokens 10, got %d", msgDelta.Usage.OutputTokens)
}
}
// TestWebSearchSendError_NonStreaming tests sendError produces correct response shape.
func TestWebSearchSendError_NonStreaming(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
innerWriter := &AnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: ginCtx.Writer},
stream: false,
id: "msg_err001",
}
wsWriter := &WebSearchAnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: ginCtx.Writer},
inner: innerWriter,
stream: false,
req: anthropic.MessagesRequest{Model: "test-model"},
}
errorUsage := anthropic.Usage{InputTokens: 7, OutputTokens: 2}
if err := wsWriter.sendError("unavailable", "test query", errorUsage); err != nil {
t.Fatalf("sendError error: %v", err)
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(rec.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v\nbody: %s", err, rec.Body.String())
}
if result.Type != "message" {
t.Errorf("expected type 'message', got %q", result.Type)
}
if result.ID != "msg_err001" {
t.Errorf("expected ID 'msg_err001', got %q", result.ID)
}
// Should have exactly 2 blocks: server_tool_use + web_search_tool_result
if len(result.Content) != 2 {
t.Fatalf("expected 2 content blocks, got %d", len(result.Content))
}
// Block 0: server_tool_use
if result.Content[0].Type != "server_tool_use" {
t.Errorf("expected 'server_tool_use', got %q", result.Content[0].Type)
}
expectedToolID := "srvtoolu_err001"
if result.Content[0].ID != expectedToolID {
t.Errorf("expected ID %q, got %q", expectedToolID, result.Content[0].ID)
}
if result.Content[0].Name != "web_search" {
t.Errorf("expected name 'web_search', got %q", result.Content[0].Name)
}
// Verify input contains the query
inputMap, ok := result.Content[0].Input.(map[string]any)
if !ok {
t.Fatalf("expected Input to be map, got %T", result.Content[0].Input)
}
if inputMap["query"] != "test query" {
t.Errorf("expected query 'test query', got %v", inputMap["query"])
}
// Block 1: web_search_tool_result with error
if result.Content[1].Type != "web_search_tool_result" {
t.Errorf("expected 'web_search_tool_result', got %q", result.Content[1].Type)
}
if result.Content[1].ToolUseID != expectedToolID {
t.Errorf("expected tool_use_id %q, got %q", expectedToolID, result.Content[1].ToolUseID)
}
// The Content field should be a WebSearchToolResultError
contentJSON, _ := json.Marshal(result.Content[1].Content)
var errContent anthropic.WebSearchToolResultError
if err := json.Unmarshal(contentJSON, &errContent); err != nil {
t.Fatalf("failed to parse error content: %v\nraw: %s", err, string(contentJSON))
}
if errContent.Type != "web_search_tool_result_error" {
t.Errorf("expected error type 'web_search_tool_result_error', got %q", errContent.Type)
}
if errContent.ErrorCode != "unavailable" {
t.Errorf("expected error_code 'unavailable', got %q", errContent.ErrorCode)
}
if result.StopReason != "end_turn" {
t.Errorf("expected stop_reason 'end_turn', got %q", result.StopReason)
}
if result.Usage != errorUsage {
t.Errorf("expected usage %+v, got %+v", errorUsage, result.Usage)
}
}
// TestWebSearchSendError_Streaming tests sendError in streaming mode produces proper SSE.
func TestWebSearchSendError_Streaming(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
innerWriter := &AnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: ginCtx.Writer},
stream: true,
id: "msg_err002",
}
wsWriter := &WebSearchAnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: ginCtx.Writer},
inner: innerWriter,
stream: true,
req: anthropic.MessagesRequest{Model: "test-model"},
}
errorUsage := anthropic.Usage{InputTokens: 9, OutputTokens: 4}
if err := wsWriter.sendError("invalid_request", "bad query", errorUsage); err != nil {
t.Fatalf("sendError error: %v", err)
}
events := parseSSEEvents(t, rec.Body.String())
// Error response has 2 blocks: server_tool_use + web_search_tool_result
// Expected events: message_start,
// content_block_start(server_tool_use), content_block_stop,
// content_block_start(web_search_tool_result), content_block_stop,
// message_delta, message_stop
expectedEventTypes := []string{
"message_start",
"content_block_start",
"content_block_stop",
"content_block_start",
"content_block_stop",
"message_delta",
"message_stop",
}
if len(events) != len(expectedEventTypes) {
t.Fatalf("expected %d events, got %d.\nEvents: %v", len(expectedEventTypes), len(events), eventNames(events))
}
for i, expected := range expectedEventTypes {
if events[i].event != expected {
t.Errorf("event[%d]: expected %q, got %q", i, expected, events[i].event)
}
}
// Verify the server_tool_use block
var toolStart anthropic.ContentBlockStartEvent
if err := json.Unmarshal([]byte(events[1].data), &toolStart); err != nil {
t.Fatalf("failed to parse server_tool_use start: %v", err)
}
if toolStart.ContentBlock.Type != "server_tool_use" {
t.Errorf("expected 'server_tool_use', got %q", toolStart.ContentBlock.Type)
}
// Verify the web_search_tool_result block
var resultStart anthropic.ContentBlockStartEvent
if err := json.Unmarshal([]byte(events[3].data), &resultStart); err != nil {
t.Fatalf("failed to parse web_search_tool_result start: %v", err)
}
if resultStart.ContentBlock.Type != "web_search_tool_result" {
t.Errorf("expected 'web_search_tool_result', got %q", resultStart.ContentBlock.Type)
}
var msgDelta anthropic.MessageDeltaEvent
if err := json.Unmarshal([]byte(events[5].data), &msgDelta); err != nil {
t.Fatalf("failed to parse message_delta: %v", err)
}
if msgDelta.Usage.InputTokens != errorUsage.InputTokens || msgDelta.Usage.OutputTokens != errorUsage.OutputTokens {
t.Fatalf("expected usage %+v in message_delta, got %+v", errorUsage, msgDelta.Usage)
}
}
// TestWebSearchSendError_EmptyQuery tests sendError with an empty query.
func TestWebSearchSendError_EmptyQuery(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
ginCtx, _ := gin.CreateTestContext(rec)
innerWriter := &AnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: ginCtx.Writer},
stream: false,
id: "msg_empty001",
}
wsWriter := &WebSearchAnthropicWriter{
BaseWriter: BaseWriter{ResponseWriter: ginCtx.Writer},
inner: innerWriter,
stream: false,
req: anthropic.MessagesRequest{Model: "test-model"},
}
if err := wsWriter.sendError("invalid_request", "", anthropic.Usage{}); err != nil {
t.Fatalf("sendError error: %v", err)
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(rec.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if len(result.Content) != 2 {
t.Fatalf("expected 2 content blocks, got %d", len(result.Content))
}
// Verify the input has empty query
inputMap, ok := result.Content[0].Input.(map[string]any)
if !ok {
t.Fatalf("expected Input to be map, got %T", result.Content[0].Input)
}
if inputMap["query"] != "" {
t.Errorf("expected empty query, got %v", inputMap["query"])
}
}
// --- SSE parsing helpers ---
type sseEvent struct {
event string
data string
}
// parseSSEEvents parses Server-Sent Events from a string.
func parseSSEEvents(t *testing.T, body string) []sseEvent {
t.Helper()
var events []sseEvent
var currentEvent string
var currentData strings.Builder
for _, line := range strings.Split(body, "\n") {
if strings.HasPrefix(line, "event: ") {
currentEvent = strings.TrimPrefix(line, "event: ")
} else if strings.HasPrefix(line, "data: ") {
currentData.WriteString(strings.TrimPrefix(line, "data: "))
} else if line == "" && currentEvent != "" {
events = append(events, sseEvent{event: currentEvent, data: currentData.String()})
currentEvent = ""
currentData.Reset()
}
}
return events
}
// eventNames returns a list of event type names for debugging.
func eventNames(events []sseEvent) []string {
names := make([]string, len(events))
for i, e := range events {
names[i] = e.event
}
return names
}
// TestWebSearchCloudModelGating tests web_search behavior across model types.
func TestWebSearchCloudModelGating(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
t.Run("local model allowed when web_search is not called", func(t *testing.T) {
handlerCalled := false
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
handlerCalled = true
resp := api.ChatResponse{
Model: "llama3.2",
Message: api.Message{Role: "assistant", Content: "hello"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{"model":"llama3.2","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"tools":[{"type":"web_search_20250305","name":"web_search"}]}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Errorf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
if !handlerCalled {
t.Error("handler should be called for local model when web_search is not called")
}
})
t.Run("local model emits web_search and gets structured error", func(t *testing.T) {
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "llama3.2",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_local_ws",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "hello"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 8, EvalCount: 2},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{"model":"llama3.2","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"tools":[{"type":"web_search_20250305","name":"web_search"}]}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if len(result.Content) != 2 {
t.Fatalf("expected 2 content blocks for local model web_search error, got %d", len(result.Content))
}
contentJSON, _ := json.Marshal(result.Content[1].Content)
var errContent anthropic.WebSearchToolResultError
if err := json.Unmarshal(contentJSON, &errContent); err != nil {
t.Fatalf("failed to parse web_search error content: %v", err)
}
if errContent.ErrorCode != "web_search_not_supported_for_local_models" {
t.Fatalf("expected web_search_not_supported_for_local_models, got %q", errContent.ErrorCode)
}
})
t.Run("model ending in cloud without cloud suffix treated as local", func(t *testing.T) {
handlerCalled := false
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
handlerCalled = true
resp := api.ChatResponse{
Model: "notreallycloud",
Message: api.Message{Role: "assistant", Content: "hello"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{"model":"notreallycloud","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"tools":[{"type":"web_search_20250305","name":"web_search"}]}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if !handlerCalled {
t.Error("handler should be called for non-cloud model when web_search is not called")
}
if resp.Code != http.StatusOK {
t.Errorf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
})
t.Run("cloud model with size tag allowed", func(t *testing.T) {
handlerCalled := false
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
handlerCalled = true
resp := api.ChatResponse{
Model: "gpt-oss:120b",
Message: api.Message{Role: "assistant", Content: "hello"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{"model":"gpt-oss:120b-cloud","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"tools":[{"type":"web_search_20250305","name":"web_search"}]}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if !handlerCalled {
t.Error("handler should be called for cloud model")
}
if resp.Code != http.StatusOK {
t.Errorf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
})
t.Run("cloud model allowed", func(t *testing.T) {
handlerCalled := false
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
handlerCalled = true
resp := api.ChatResponse{
Model: "kimi-k2.5",
Message: api.Message{Role: "assistant", Content: "hello"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{"model":"kimi-k2.5:cloud","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"tools":[{"type":"web_search_20250305","name":"web_search"}]}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if !handlerCalled {
t.Error("handler should be called for cloud model")
}
if resp.Code != http.StatusOK {
t.Errorf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
})
t.Run("cloud disabled blocks web search for cloud model", func(t *testing.T) {
t.Setenv("OLLAMA_NO_CLOUD", "1")
handlerCalled := false
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
handlerCalled = true
})
body := `{"model":"kimi-k2.5:cloud","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"tools":[{"type":"web_search_20250305","name":"web_search"}]}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusForbidden {
t.Fatalf("expected 403, got %d: %s", resp.Code, resp.Body.String())
}
if handlerCalled {
t.Fatal("handler should not be called when cloud is disabled")
}
var errResp anthropic.ErrorResponse
if err := json.Unmarshal(resp.Body.Bytes(), &errResp); err != nil {
t.Fatalf("failed to parse error response: %v", err)
}
if !strings.Contains(errResp.Error.Message, "ollama cloud is disabled") {
t.Fatalf("expected cloud disabled error, got: %q", errResp.Error.Message)
}
})
t.Run("cloud disabled does not block local model if web_search is not called", func(t *testing.T) {
t.Setenv("OLLAMA_NO_CLOUD", "1")
handlerCalled := false
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
handlerCalled = true
resp := api.ChatResponse{
Model: "llama3.2",
Message: api.Message{Role: "assistant", Content: "hello"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{"model":"llama3.2","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"tools":[{"type":"web_search_20250305","name":"web_search"}]}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
if !handlerCalled {
t.Fatal("handler should be called for local model when web_search is not called")
}
})
}
func TestWebSearchDoesNotRequireAuthorizationHeaderForMockEndpoint(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
var authHeader string
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
authHeader = r.Header.Get("Authorization")
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "done"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 5, EvalCount: 2},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_auth",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "auth test"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 4, EvalCount: 1},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"test auth"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
if authHeader != "" {
t.Fatalf("expected no Authorization header for mock web search endpoint, got %q", authHeader)
}
}
// TestWebSearchSearchAPIError tests that a failing search API returns a proper error response.
func TestWebSearchSearchAPIError(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
// Mock search server that returns 500
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "internal error", http.StatusInternalServerError)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_err",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "test"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 2},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"test"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
// Error response: server_tool_use + web_search_tool_result with error
if len(result.Content) != 2 {
t.Fatalf("expected 2 content blocks for error, got %d", len(result.Content))
}
if result.Content[0].Type != "server_tool_use" {
t.Errorf("expected 'server_tool_use', got %q", result.Content[0].Type)
}
if result.Content[1].Type != "web_search_tool_result" {
t.Errorf("expected 'web_search_tool_result', got %q", result.Content[1].Type)
}
if result.Usage.InputTokens != 10 || result.Usage.OutputTokens != 2 {
t.Fatalf("expected usage input=10 output=2, got %+v", result.Usage)
}
}
func TestWebSearchStreamingImmediateTakeover(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "After search."},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 20, EvalCount: 10},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
chunks := []api.ChatResponse{
{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "Preface "},
Done: false,
},
{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_stream_1",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "latest updates"),
},
},
},
},
Done: false,
},
{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "ignored chunk"},
Done: false,
},
{
Model: "test-model",
Message: api.Message{Role: "assistant"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 9, EvalCount: 4},
},
}
c.Writer.WriteHeader(http.StatusOK)
for _, chunk := range chunks {
data, _ := json.Marshal(chunk)
_, _ = c.Writer.Write(data)
}
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"stream":true,
"messages":[{"role":"user","content":"Find updates"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
events := parseSSEEvents(t, resp.Body.String())
if countEventsByName(events, "message_start") != 1 {
t.Fatalf("expected exactly one message_start, got %d", countEventsByName(events, "message_start"))
}
if countEventsByName(events, "message_stop") != 1 {
t.Fatalf("expected exactly one message_stop, got %d", countEventsByName(events, "message_stop"))
}
textDeltas := collectTextDeltas(t, events)
if !containsString(textDeltas, "Preface ") {
t.Fatalf("expected passthrough text delta, got %v", textDeltas)
}
if !containsString(textDeltas, "After search.") {
t.Fatalf("expected post-search text delta, got %v", textDeltas)
}
if containsString(textDeltas, "ignored chunk") {
t.Fatalf("unexpected text from chunks after takeover: %v", textDeltas)
}
}
func TestWebSearchStreamingUsageUsesObservedChunkMetrics(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "After search."},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 20, EvalCount: 7},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
chunks := []api.ChatResponse{
{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "Preface "},
Done: false,
Metrics: api.Metrics{PromptEvalCount: 12, EvalCount: 4},
},
{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_stream_usage",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "latest updates"),
},
},
},
},
Done: false,
Metrics: api.Metrics{PromptEvalCount: 0, EvalCount: 0},
},
{
Model: "test-model",
Message: api.Message{Role: "assistant"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 12, EvalCount: 4},
},
}
c.Writer.WriteHeader(http.StatusOK)
for _, chunk := range chunks {
data, _ := json.Marshal(chunk)
_, _ = c.Writer.Write(data)
}
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"stream":true,
"messages":[{"role":"user","content":"Find updates"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
events := parseSSEEvents(t, resp.Body.String())
var messageDelta anthropic.MessageDeltaEvent
found := false
for _, event := range events {
if event.event != "message_delta" {
continue
}
if err := json.Unmarshal([]byte(event.data), &messageDelta); err != nil {
t.Fatalf("failed to unmarshal message_delta: %v", err)
}
found = true
break
}
if !found {
t.Fatal("expected message_delta event")
}
if messageDelta.Usage.InputTokens != 32 {
t.Fatalf("expected aggregated input tokens 32 (12 passthrough + 20 followup), got %d", messageDelta.Usage.InputTokens)
}
if messageDelta.Usage.OutputTokens != 11 {
t.Fatalf("expected aggregated output tokens 11 (4 passthrough + 7 followup), got %d", messageDelta.Usage.OutputTokens)
}
}
func TestWebSearchMixedToolCallsPreferWebSearch(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "Search answer."},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 11, EvalCount: 6},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_other",
Function: api.ToolCallFunction{
Name: "get_weather",
Arguments: makeArgs("location", "SF"),
},
},
{
ID: "call_ws_mixed",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "latest weather"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 2},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"Weather?"}],
"tools":[
{"type":"web_search_20250305","name":"web_search"},
{"type":"custom","name":"get_weather","input_schema":{"type":"object"}}
]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if len(result.Content) < 3 {
t.Fatalf("expected at least 3 blocks, got %d", len(result.Content))
}
if result.Content[0].Type != "server_tool_use" {
t.Fatalf("expected server_tool_use first, got %q", result.Content[0].Type)
}
if result.Content[1].Type != "web_search_tool_result" {
t.Fatalf("expected web_search_tool_result second, got %q", result.Content[1].Type)
}
for _, block := range result.Content {
if block.Type == "tool_use" && block.Name == "get_weather" {
t.Fatalf("did not expect get_weather tool_use in mixed web_search-preferred path: %+v", result.Content)
}
}
}
func TestWebSearchFollowupClientToolStopReasonToolUse(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_weather_final",
Function: api.ToolCallFunction{
Name: "get_weather",
Arguments: makeArgs("location", "New York"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 25, EvalCount: 7},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_tool_use",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "forecast"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 15, EvalCount: 3},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"Do I need an umbrella?"}],
"tools":[
{"type":"web_search_20250305","name":"web_search"},
{"type":"custom","name":"get_weather","input_schema":{"type":"object"}}
]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if result.StopReason != "tool_use" {
t.Fatalf("expected stop_reason tool_use, got %q", result.StopReason)
}
if len(result.Content) < 3 {
t.Fatalf("expected server blocks + tool_use, got %d blocks", len(result.Content))
}
last := result.Content[len(result.Content)-1]
if last.Type != "tool_use" {
t.Fatalf("expected final block tool_use, got %q", last.Type)
}
if last.Name != "get_weather" {
t.Fatalf("expected final tool name get_weather, got %q", last.Name)
}
if result.Usage.InputTokens != 40 || result.Usage.OutputTokens != 10 {
t.Fatalf("unexpected aggregated usage: %+v", result.Usage)
}
}
func TestWebSearchMultiIterationLoop(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
followupCall := 0
followupDecodeErr := false
missingWebSearchTool := false
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var followupReq api.ChatRequest
if err := json.NewDecoder(r.Body).Decode(&followupReq); err != nil {
followupDecodeErr = true
http.Error(w, "bad request", http.StatusBadRequest)
return
}
hasWebSearchTool := false
for _, tool := range followupReq.Tools {
if tool.Function.Name == "web_search" {
hasWebSearchTool = true
break
}
}
if !hasWebSearchTool {
missingWebSearchTool = true
}
followupCall++
switch followupCall {
case 1:
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_2",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "loop query 2"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 20, EvalCount: 2},
}
_ = json.NewEncoder(w).Encode(resp)
case 2:
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "Final answer after 2 searches."},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 30, EvalCount: 3},
}
_ = json.NewEncoder(w).Encode(resp)
default:
t.Fatalf("unexpected extra followup call: %d", followupCall)
}
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_1",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "loop query 1"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 1},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"do multiple searches"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
if followupCall != 2 {
t.Fatalf("expected 2 followup calls, got %d", followupCall)
}
if followupDecodeErr {
t.Fatal("failed to decode followup request body")
}
if missingWebSearchTool {
t.Fatal("expected followup requests to retain web_search tool definition")
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
serverToolUses := 0
webResults := 0
for _, block := range result.Content {
if block.Type == "server_tool_use" {
serverToolUses++
}
if block.Type == "web_search_tool_result" {
webResults++
}
}
if serverToolUses != 2 || webResults != 2 {
t.Fatalf("expected two search iterations, got server_tool_use=%d web_search_tool_result=%d", serverToolUses, webResults)
}
if result.Usage.InputTokens != 60 || result.Usage.OutputTokens != 6 {
t.Fatalf("unexpected aggregated usage: %+v", result.Usage)
}
}
func TestWebSearchLoopMaxLimit(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
followupCall := 0
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
followupCall++
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_loop_limit",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "loop query next"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 7, EvalCount: 2},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_initial",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "loop query 1"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 5, EvalCount: 1},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"keep searching"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
if followupCall != 3 {
t.Fatalf("expected 3 followup calls before max loop error, got %d", followupCall)
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
last := result.Content[len(result.Content)-1]
if last.Type != "web_search_tool_result" {
t.Fatalf("expected last block web_search_tool_result, got %q", last.Type)
}
contentJSON, _ := json.Marshal(last.Content)
var errContent anthropic.WebSearchToolResultError
if err := json.Unmarshal(contentJSON, &errContent); err != nil {
t.Fatalf("failed to parse web search error content: %v", err)
}
if errContent.ErrorCode != "max_uses_exceeded" {
t.Fatalf("expected max_uses_exceeded error, got %q", errContent.ErrorCode)
}
if result.StopReason != "end_turn" {
t.Fatalf("expected end_turn, got %q", result.StopReason)
}
}
func TestWebSearchStreamingFinalStopReasonToolUse(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_weather_stream",
Function: api.ToolCallFunction{
Name: "get_weather",
Arguments: makeArgs("location", "Seattle"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 14, EvalCount: 5},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
chunks := []api.ChatResponse{
{
Model: "test-model",
Message: api.Message{Role: "assistant", Content: "Let me check. "},
Done: false,
},
{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_stream_tool_use",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "weather seattle"),
},
},
},
},
Done: false,
},
{
Model: "test-model",
Message: api.Message{Role: "assistant"},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 3},
},
}
c.Writer.WriteHeader(http.StatusOK)
for _, chunk := range chunks {
data, _ := json.Marshal(chunk)
_, _ = c.Writer.Write(data)
}
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"stream":true,
"messages":[{"role":"user","content":"Should I take a jacket?"}],
"tools":[
{"type":"web_search_20250305","name":"web_search"},
{"type":"custom","name":"get_weather","input_schema":{"type":"object"}}
]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
events := parseSSEEvents(t, resp.Body.String())
if countEventsByName(events, "message_start") != 1 {
t.Fatalf("expected exactly one message_start, got %d", countEventsByName(events, "message_start"))
}
var messageDelta anthropic.MessageDeltaEvent
foundMessageDelta := false
foundToolUse := false
for _, event := range events {
if event.event == "message_delta" {
foundMessageDelta = true
if err := json.Unmarshal([]byte(event.data), &messageDelta); err != nil {
t.Fatalf("failed to unmarshal message_delta: %v", err)
}
}
if event.event == "content_block_start" {
var start anthropic.ContentBlockStartEvent
if err := json.Unmarshal([]byte(event.data), &start); err != nil {
t.Fatalf("failed to unmarshal content_block_start: %v", err)
}
if start.ContentBlock.Type == "tool_use" && start.ContentBlock.Name == "get_weather" {
foundToolUse = true
}
}
}
if !foundMessageDelta {
t.Fatal("expected message_delta event")
}
if messageDelta.Delta.StopReason != "tool_use" {
t.Fatalf("expected stop_reason tool_use, got %q", messageDelta.Delta.StopReason)
}
if !foundToolUse {
t.Fatal("expected tool_use content block for get_weather")
}
}
func TestWebSearchFollowupNon200ReturnsApiError(t *testing.T) {
gin.SetMode(gin.TestMode)
enableCloudForTest(t)
followupServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "boom", http.StatusInternalServerError)
}))
defer followupServer.Close()
t.Setenv("OLLAMA_HOST", followupServer.URL)
searchServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := anthropic.OllamaWebSearchResponse{
Results: []anthropic.OllamaWebSearchResult{
{Title: "Result", URL: "https://example.com", Content: "content"},
},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer searchServer.Close()
originalEndpoint := anthropic.WebSearchEndpoint
anthropic.WebSearchEndpoint = searchServer.URL
defer func() { anthropic.WebSearchEndpoint = originalEndpoint }()
router := gin.New()
router.Use(AnthropicMessagesMiddleware())
router.POST("/v1/messages", func(c *gin.Context) {
resp := api.ChatResponse{
Model: "test-model",
Message: api.Message{
Role: "assistant",
ToolCalls: []api.ToolCall{
{
ID: "call_ws_non200",
Function: api.ToolCallFunction{
Name: "web_search",
Arguments: makeArgs("query", "test"),
},
},
},
},
Done: true,
DoneReason: "stop",
Metrics: api.Metrics{PromptEvalCount: 9, EvalCount: 1},
}
data, _ := json.Marshal(resp)
c.Writer.WriteHeader(http.StatusOK)
_, _ = c.Writer.Write(data)
})
body := `{
"model":"test-model:cloud",
"max_tokens":100,
"messages":[{"role":"user","content":"test"}],
"tools":[{"type":"web_search_20250305","name":"web_search"}]
}`
req, _ := http.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
if resp.Code != http.StatusOK {
t.Fatalf("expected 200, got %d: %s", resp.Code, resp.Body.String())
}
var result anthropic.MessagesResponse
if err := json.Unmarshal(resp.Body.Bytes(), &result); err != nil {
t.Fatalf("unmarshal error: %v", err)
}
if len(result.Content) != 2 {
t.Fatalf("expected 2 blocks in error response, got %d", len(result.Content))
}
contentJSON, _ := json.Marshal(result.Content[1].Content)
var errContent anthropic.WebSearchToolResultError
if err := json.Unmarshal(contentJSON, &errContent); err != nil {
t.Fatalf("failed to parse error content: %v", err)
}
if errContent.ErrorCode != "api_error" {
t.Fatalf("expected api_error, got %q", errContent.ErrorCode)
}
if result.Usage.InputTokens != 9 || result.Usage.OutputTokens != 1 {
t.Fatalf("expected usage input=9 output=1, got %+v", result.Usage)
}
}
func countEventsByName(events []sseEvent, eventName string) int {
count := 0
for _, event := range events {
if event.event == eventName {
count++
}
}
return count
}
func collectTextDeltas(t *testing.T, events []sseEvent) []string {
t.Helper()
var deltas []string
for _, event := range events {
if event.event != "content_block_delta" {
continue
}
var delta anthropic.ContentBlockDeltaEvent
if err := json.Unmarshal([]byte(event.data), &delta); err != nil {
t.Fatalf("failed to unmarshal content_block_delta: %v", err)
}
if delta.Delta.Type == "text_delta" {
deltas = append(deltas, delta.Delta.Text)
}
}
return deltas
}
func containsString(values []string, target string) bool {
for _, value := range values {
if value == target {
return true
}
}
return false
}
+22
View File
@@ -0,0 +1,22 @@
package middleware
import (
"testing"
"github.com/ollama/ollama/envconfig"
)
func setTestHome(t *testing.T, home string) {
t.Helper()
t.Setenv("HOME", home)
t.Setenv("USERPROFILE", home)
envconfig.ReloadServerConfig()
}
// enableCloudForTest sets HOME to a clean temp dir and clears OLLAMA_NO_CLOUD
// so that cloud features are enabled for the duration of the test.
func enableCloudForTest(t *testing.T) {
t.Helper()
t.Setenv("OLLAMA_NO_CLOUD", "")
setTestHome(t, t.TempDir())
}