Add Anthropic streaming support and fix provider factory

Switch anthropic protocol to native SDK instead of generic HTTP provider,
enabling streaming and proper tool call handling. Add ChatStream() method
using Messages.NewStreaming() for incremental response delivery.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
pmorgan 2026-02-22 11:57:53 -07:00
parent c51ceac70b
commit b9440c9f13
4 changed files with 258 additions and 14 deletions

View file

@ -93,6 +93,59 @@ func (p *Provider) Chat(
return parseResponse(resp), nil return parseResponse(resp), nil
} }
// ChatStream sends a streaming request to the Anthropic API.
// It calls onDelta for each text fragment as it arrives, then returns the
// fully accumulated response (identical to what Chat would return).
func (p *Provider) ChatStream(
ctx context.Context,
messages []Message,
tools []ToolDefinition,
model string,
options map[string]any,
onDelta func(delta string),
) (*LLMResponse, error) {
var opts []option.RequestOption
if p.tokenSource != nil {
tok, err := p.tokenSource()
if err != nil {
return nil, fmt.Errorf("refreshing token: %w", err)
}
opts = append(opts, option.WithAuthToken(tok))
}
params, err := buildParams(messages, tools, model, options)
if err != nil {
return nil, err
}
stream := p.client.Messages.NewStreaming(ctx, params, opts...)
var accumulated anthropic.Message
for stream.Next() {
event := stream.Current()
if err := accumulated.Accumulate(event); err != nil {
return nil, fmt.Errorf("accumulating stream event: %w", err)
}
// Deliver text deltas to the callback
if onDelta != nil {
switch e := event.AsAny().(type) {
case anthropic.ContentBlockDeltaEvent:
if td := e.Delta.AsTextDelta(); td.Text != "" {
onDelta(td.Text)
}
}
}
}
if err := stream.Err(); err != nil {
return nil, fmt.Errorf("claude streaming API call: %w", err)
}
return parseResponse(&accumulated), nil
}
func (p *Provider) GetDefaultModel() string { func (p *Provider) GetDefaultModel() string {
return "claude-sonnet-4.6" return "claude-sonnet-4.6"
} }
@ -110,15 +163,26 @@ func buildParams(
var system []anthropic.TextBlockParam var system []anthropic.TextBlockParam
var anthropicMessages []anthropic.MessageParam var anthropicMessages []anthropic.MessageParam
for _, msg := range messages { // Build messages, merging consecutive tool results into a single user
// message. The Anthropic API requires that ALL tool_result blocks for a
// given assistant tool_use turn appear in one user message immediately
// after the assistant message.
for i := 0; i < len(messages); i++ {
msg := messages[i]
switch msg.Role { switch msg.Role {
case "system": case "system":
system = append(system, anthropic.TextBlockParam{Text: msg.Content}) system = append(system, anthropic.TextBlockParam{Text: msg.Content})
case "user": case "user":
if msg.ToolCallID != "" { if msg.ToolCallID != "" {
anthropicMessages = append(anthropicMessages, // Tool result stored with "user" role — collect consecutive ones.
anthropic.NewUserMessage(anthropic.NewToolResultBlock(msg.ToolCallID, msg.Content, false)), var toolBlocks []anthropic.ContentBlockParamUnion
) for i < len(messages) && isToolResult(messages[i]) {
toolBlocks = append(toolBlocks,
anthropic.NewToolResultBlock(messages[i].ToolCallID, messages[i].Content, false))
i++
}
i-- // outer loop will increment
anthropicMessages = append(anthropicMessages, anthropic.NewUserMessage(toolBlocks...))
} else { } else {
anthropicMessages = append(anthropicMessages, anthropicMessages = append(anthropicMessages,
anthropic.NewUserMessage(anthropic.NewTextBlock(msg.Content)), anthropic.NewUserMessage(anthropic.NewTextBlock(msg.Content)),
@ -131,7 +195,14 @@ func buildParams(
blocks = append(blocks, anthropic.NewTextBlock(msg.Content)) blocks = append(blocks, anthropic.NewTextBlock(msg.Content))
} }
for _, tc := range msg.ToolCalls { for _, tc := range msg.ToolCalls {
blocks = append(blocks, anthropic.NewToolUseBlock(tc.ID, tc.Arguments, tc.Name)) args := tc.Arguments
if args == nil && tc.Function != nil && tc.Function.Arguments != "" {
_ = json.Unmarshal([]byte(tc.Function.Arguments), &args)
}
if args == nil {
args = map[string]any{}
}
blocks = append(blocks, anthropic.NewToolUseBlock(tc.ID, args, tc.Name))
} }
anthropicMessages = append(anthropicMessages, anthropic.NewAssistantMessage(blocks...)) anthropicMessages = append(anthropicMessages, anthropic.NewAssistantMessage(blocks...))
} else { } else {
@ -140,9 +211,15 @@ func buildParams(
) )
} }
case "tool": case "tool":
anthropicMessages = append(anthropicMessages, // Collect all consecutive tool results into one user message.
anthropic.NewUserMessage(anthropic.NewToolResultBlock(msg.ToolCallID, msg.Content, false)), var toolBlocks []anthropic.ContentBlockParamUnion
) for i < len(messages) && isToolResult(messages[i]) {
toolBlocks = append(toolBlocks,
anthropic.NewToolResultBlock(messages[i].ToolCallID, messages[i].Content, false))
i++
}
i-- // outer loop will increment
anthropicMessages = append(anthropicMessages, anthropic.NewUserMessage(toolBlocks...))
} }
} }
@ -244,6 +321,12 @@ func parseResponse(resp *anthropic.Message) *LLMResponse {
} }
} }
// isToolResult returns true if the message is a tool result, regardless of
// whether it's stored with "tool" role or "user" role with a ToolCallID.
func isToolResult(msg Message) bool {
return msg.Role == "tool" || (msg.Role == "user" && msg.ToolCallID != "")
}
func normalizeBaseURL(apiBase string) string { func normalizeBaseURL(apiBase string) string {
base := strings.TrimSpace(apiBase) base := strings.TrimSpace(apiBase)
if base == "" { if base == "" {

View file

@ -1,7 +1,9 @@
package anthropicprovider package anthropicprovider
import ( import (
"context"
"encoding/json" "encoding/json"
"fmt"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"sync/atomic" "sync/atomic"
@ -77,6 +79,55 @@ func TestBuildParams_ToolCallMessage(t *testing.T) {
} }
} }
func TestBuildParams_MultipleToolResults_MergedIntoOneMessage(t *testing.T) {
// When an assistant makes multiple tool calls, all tool results must be
// merged into a single user message for the Anthropic API.
messages := []Message{
{Role: "user", Content: "Check all endpoints"},
{
Role: "assistant",
ToolCalls: []ToolCall{
{ID: "call_1", Name: "web_fetch", Arguments: map[string]any{"url": "http://a"}},
{ID: "call_2", Name: "web_fetch", Arguments: map[string]any{"url": "http://b"}},
{ID: "call_3", Name: "web_fetch", Arguments: map[string]any{"url": "http://c"}},
},
},
{Role: "tool", Content: `{"status":"ok"}`, ToolCallID: "call_1"},
{Role: "tool", Content: `{"status":"ok"}`, ToolCallID: "call_2"},
{Role: "tool", Content: `{"status":"ok"}`, ToolCallID: "call_3"},
{Role: "assistant", Content: "All endpoints are healthy."},
}
params, err := buildParams(messages, nil, "claude-sonnet-4-6", map[string]any{})
if err != nil {
t.Fatalf("buildParams() error: %v", err)
}
// Expected: user, assistant(3 tool_use), user(3 tool_results merged), assistant
if len(params.Messages) != 4 {
t.Fatalf("len(Messages) = %d, want 4 (tool results should be merged)", len(params.Messages))
}
}
func TestBuildParams_SingleToolResult_StillWorks(t *testing.T) {
// Single tool call should still produce 3 messages (user, assistant, tool_result).
messages := []Message{
{Role: "user", Content: "What's the weather?"},
{
Role: "assistant",
ToolCalls: []ToolCall{
{ID: "call_1", Name: "get_weather", Arguments: map[string]any{"city": "SF"}},
},
},
{Role: "tool", Content: `{"temp": 72}`, ToolCallID: "call_1"},
}
params, err := buildParams(messages, nil, "claude-sonnet-4-6", map[string]any{})
if err != nil {
t.Fatalf("buildParams() error: %v", err)
}
if len(params.Messages) != 3 {
t.Fatalf("len(Messages) = %d, want 3", len(params.Messages))
}
}
func TestBuildParams_WithTools(t *testing.T) { func TestBuildParams_WithTools(t *testing.T) {
tools := []ToolDefinition{ tools := []ToolDefinition{
{ {
@ -262,6 +313,100 @@ func TestProvider_ChatUsesTokenSource(t *testing.T) {
} }
} }
func TestProvider_ChatStream_RoundTrip(t *testing.T) {
// SSE streaming mock: sends message_start, content_block_start,
// content_block_delta (x2), content_block_stop, message_delta, message_stop.
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/v1/messages" {
http.Error(w, "not found", http.StatusNotFound)
return
}
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
flusher, ok := w.(http.Flusher)
if !ok {
http.Error(w, "streaming unsupported", http.StatusInternalServerError)
return
}
events := []string{
`event: message_start
data: {"type":"message_start","message":{"id":"msg_test","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4.6","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":0}}}`,
`event: content_block_start
data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`,
`event: content_block_delta
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}`,
`event: content_block_delta
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":" World"}}`,
`event: content_block_stop
data: {"type":"content_block_stop","index":0}`,
`event: message_delta
data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":5}}`,
`event: message_stop
data: {"type":"message_stop"}`,
}
for _, event := range events {
fmt.Fprintf(w, "%s\n\n", event)
flusher.Flush()
}
}))
defer server.Close()
provider := NewProviderWithClient(createAnthropicTestClient(server.URL, "test-token"))
var deltas []string
resp, err := provider.ChatStream(
t.Context(),
[]Message{{Role: "user", Content: "Hello"}},
nil,
"claude-sonnet-4.6",
map[string]any{"max_tokens": 1024},
func(delta string) {
deltas = append(deltas, delta)
},
)
if err != nil {
t.Fatalf("ChatStream() error: %v", err)
}
if resp.Content != "Hello World" {
t.Errorf("Content = %q, want %q", resp.Content, "Hello World")
}
if resp.FinishReason != "stop" {
t.Errorf("FinishReason = %q, want %q", resp.FinishReason, "stop")
}
if len(deltas) != 2 {
t.Errorf("len(deltas) = %d, want 2", len(deltas))
}
if len(deltas) >= 2 {
if deltas[0] != "Hello" {
t.Errorf("deltas[0] = %q, want %q", deltas[0], "Hello")
}
if deltas[1] != " World" {
t.Errorf("deltas[1] = %q, want %q", deltas[1], " World")
}
}
if resp.Usage.PromptTokens != 10 {
t.Errorf("PromptTokens = %d, want 10", resp.Usage.PromptTokens)
}
}
func TestProvider_ImplementsStreamingProvider(t *testing.T) {
// Verify that Provider satisfies the StreamingProvider interface
var _ interface {
ChatStream(
ctx context.Context,
messages []Message,
tools []ToolDefinition,
model string,
options map[string]any,
onDelta func(string),
) (*LLMResponse, error)
} = &Provider{}
}
func createAnthropicTestClient(baseURL, token string) *anthropic.Client { func createAnthropicTestClient(baseURL, token string) *anthropic.Client {
c := anthropic.NewClient( c := anthropic.NewClient(
anthropicoption.WithAuthToken(token), anthropicoption.WithAuthToken(token),

View file

@ -108,15 +108,12 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err
} }
return provider, modelID, nil return provider, modelID, nil
} }
// Use API key with HTTP API // Use API key with native Anthropic SDK
apiBase := cfg.APIBase
if apiBase == "" {
apiBase = "https://api.anthropic.com/v1"
}
if cfg.APIKey == "" { if cfg.APIKey == "" {
return nil, "", fmt.Errorf("api_key is required for anthropic protocol (model: %s)", cfg.Model) return nil, "", fmt.Errorf("api_key is required for anthropic protocol (model: %s)", cfg.Model)
} }
return NewHTTPProviderWithMaxTokensField(cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField), modelID, nil apiBase := cfg.APIBase
return NewClaudeProviderWithBaseURL(cfg.APIKey, apiBase), modelID, nil
case "antigravity": case "antigravity":
return NewAntigravityProvider(), modelID, nil return NewAntigravityProvider(), modelID, nil

View file

@ -30,6 +30,25 @@ type LLMProvider interface {
GetDefaultModel() string GetDefaultModel() string
} }
// StreamCallback receives partial content as it arrives from the LLM.
type StreamCallback func(delta string)
// StreamingProvider is an optional interface for providers that support streaming.
// Providers that implement this can send partial responses as they arrive,
// which is important for large context windows where non-streaming calls
// can take 30+ seconds.
type StreamingProvider interface {
LLMProvider
ChatStream(
ctx context.Context,
messages []Message,
tools []ToolDefinition,
model string,
options map[string]any,
onDelta StreamCallback,
) (*LLMResponse, error)
}
// FailoverReason classifies why an LLM request failed for fallback decisions. // FailoverReason classifies why an LLM request failed for fallback decisions.
type FailoverReason string type FailoverReason string