anthropic: enable websearch (#14246)
This commit is contained in:
+817
-17
@@ -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(¤tToolCall)
|
||||
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(¤tToolCall)
|
||||
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_")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
Reference in New Issue
Block a user