From 5e2dee44a87cc891bab18e8b53f141e6529efdaa Mon Sep 17 00:00:00 2001 From: Lsong Date: Thu, 10 Sep 2026 20:02:25 +0800 Subject: [PATCH] refactor: simplify agent provider core --- agent/agent.go | 73 ++++--- agent/agent_loop_test.go | 52 ++--- agent/manager.go | 48 ++--- agent/manager_test.go | 58 +++++- anthropic/client.go | 64 ++++-- anthropic/client_test.go | 36 ++-- anthropic/client_url_test.go | 18 ++ anthropic/openai.go | 359 ++++++++++++++++++++------------- anthropic/openai_test.go | 85 ++++++++ anthropic/types.go | 96 +++++---- openai/client.go | 40 ++-- openai/client_refactor_test.go | 49 +++++ 12 files changed, 661 insertions(+), 317 deletions(-) create mode 100644 anthropic/client_url_test.go create mode 100644 anthropic/openai_test.go create mode 100644 openai/client_refactor_test.go diff --git a/agent/agent.go b/agent/agent.go index aceba95..06159ee 100644 --- a/agent/agent.go +++ b/agent/agent.go @@ -13,35 +13,62 @@ import ( "github.com/lsongdev/miya-agents/tools" ) -type LLM interface { - CreateChatCompletionStream(context.Context, *openai.ChatCompletionRequest) (<-chan openai.ChatCompletionResponse, error) -} +// StreamFunc is the only model capability required by the agent loop. +type StreamFunc func(context.Context, *openai.ChatCompletionRequest) (<-chan openai.ChatCompletionResponse, error) type Agent struct { Name string Config *config.ProfileConfig - LLM LLM - // tools - toolsMap map[string]openai.Tool - toolsDefs []openai.ToolDef + Stream StreamFunc + tools []openai.Tool +} + +func New(name string, cfg *config.ProfileConfig, stream StreamFunc) *Agent { + return &Agent{Name: name, Config: cfg, Stream: stream} +} + +// Use adds tools to the agent and returns it for chaining. +func (a *Agent) Use(tools ...openai.Tool) *Agent { + a.tools = append(a.tools, tools...) + return a +} + +func (a *Agent) tool(name string) (openai.Tool, bool) { + for _, tool := range a.tools { + if tool.Def().Function.Name == name { + return tool, true + } + } + return nil, false +} + +func (a *Agent) toolDefs() []openai.ToolDef { + defs := make([]openai.ToolDef, len(a.tools)) + for i, tool := range a.tools { + defs[i] = tool.Def() + } + return defs } func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink EventSink) error { + if a.Stream == nil { + return fmt.Errorf("agent has no model stream") + } for { req := openai.ChatCompletionRequest{ Model: a.Config.ModelName, Messages: sess.Messages, - Tools: a.toolsDefs, + Tools: a.toolDefs(), Stream: true, } - resp, err := a.LLM.CreateChatCompletionStream(ctx, &req) + resp, err := a.Stream(ctx, &req) if err != nil { - return fmt.Errorf("failed to create chat completion stream: %w", err) + return fmt.Errorf("create model stream: %w", err) } builder := openai.NewMessageBuilder() for chunk := range resp { if chunk.Error != nil { - return fmt.Errorf("API error: %s", chunk.Error.Message) + return fmt.Errorf("model stream: %s", chunk.Error.Message) } m := chunk.GetMessage() if m == nil { @@ -64,10 +91,9 @@ func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink Ev } respMessage := builder.Build() if respMessage.IsEmpty() { - return fmt.Errorf("chat completion stream closed without a response") + return fmt.Errorf("model stream closed without a response") } sess.AppendResponse(respMessage) - // finish if !respMessage.HasToolCall() { if err := sink.Usage(UsageEvent{}); err != nil { return err @@ -81,12 +107,12 @@ func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink Ev } return nil } - // Execute tool calls + for _, tc := range respMessage.ToolCalls { if tc.ID == "" { return fmt.Errorf("tool call %q is missing an id", tc.Function.Name) } - tool, ok := a.toolsMap[tc.Function.Name] + tool, ok := a.tool(tc.Function.Name) if err := sink.ToolCallStart(ToolCallEvent{ ID: tc.ID, Name: tc.Function.Name, @@ -156,12 +182,6 @@ func emitAttachedFileResult(sink EventSink, result string) (string, bool, error) return fmt.Sprintf("Attached %s (%s, %d bytes) as %s.", attachment.Name, attachment.MimeType, attachment.Size, attachment.URI), true, nil } -func (a *Agent) AddTool(tool openai.Tool) { - d := tool.Def() - a.toolsMap[d.Function.Name] = tool - a.toolsDefs = append(a.toolsDefs, d) -} - func (a *Agent) NewSession() *session.Session { s := session.New(a.Name) prompt := a.readSystemPrompt() @@ -196,7 +216,7 @@ func (a *Agent) BuildTools() { if workspace != "" { _ = os.MkdirAll(workspace, 0755) } - var tools = []openai.Tool{ + a.Use( &tools.WebFetchTool{}, &tools.WebSearchTool{}, &tools.ReadFileTool{Workspace: workspace}, @@ -209,11 +229,6 @@ func (a *Agent) BuildTools() { DefaultTimeout: tools.ExecDefaultTimeoutSeconds, RestrictToWorkspace: true, }, - &tools.SkillsTool{ - Workspace: filepath.Join(config.ConfigPath, "skills"), - }, - } - for _, t := range tools { - a.AddTool(t) - } + &tools.SkillsTool{Workspace: filepath.Join(config.ConfigPath, "skills")}, + ) } diff --git a/agent/agent_loop_test.go b/agent/agent_loop_test.go index e288732..078dec7 100644 --- a/agent/agent_loop_test.go +++ b/agent/agent_loop_test.go @@ -12,17 +12,15 @@ import ( "github.com/lsongdev/miya-agents/session" ) -type fakeLLM struct { - chunks []openai.ChatCompletionResponse -} - -func (m *fakeLLM) CreateChatCompletionStream(context.Context, *openai.ChatCompletionRequest) (<-chan openai.ChatCompletionResponse, error) { - ch := make(chan openai.ChatCompletionResponse, len(m.chunks)) - for _, chunk := range m.chunks { - ch <- chunk +func fakeStream(chunks ...openai.ChatCompletionResponse) StreamFunc { + return func(context.Context, *openai.ChatCompletionRequest) (<-chan openai.ChatCompletionResponse, error) { + ch := make(chan openai.ChatCompletionResponse, len(chunks)) + for _, chunk := range chunks { + ch <- chunk + } + close(ch) + return ch, nil } - close(ch) - return ch, nil } type discardSink struct{} @@ -37,10 +35,7 @@ func (discardSink) Usage(UsageEvent) error { return nil } func (discardSink) Done() error { return nil } func TestRunAgentLoopRejectsEmptyStream(t *testing.T) { - ag := &Agent{ - Config: &config.ProfileConfig{ModelName: "test"}, - LLM: &fakeLLM{}, - } + ag := New("test", &config.ProfileConfig{ModelName: "test"}, fakeStream()) err := ag.RunAgentLoop(context.Background(), session.New("test"), discardSink{}) if err == nil || !strings.Contains(err.Error(), "closed without a response") { @@ -50,13 +45,10 @@ func TestRunAgentLoopRejectsEmptyStream(t *testing.T) { func TestRunAgentLoopRejectsInterruptedStream(t *testing.T) { message := openai.ChatCompletionMessage{Role: openai.RoleAssistant, Content: "partial"} - ag := &Agent{ - Config: &config.ProfileConfig{ModelName: "test"}, - LLM: &fakeLLM{chunks: []openai.ChatCompletionResponse{ - {Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}}, - {Error: &openai.Error{Type: "stream_error", Message: "connection reset"}}, - }}, - } + ag := New("test", &config.ProfileConfig{ModelName: "test"}, fakeStream( + openai.ChatCompletionResponse{Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}}, + openai.ChatCompletionResponse{Error: &openai.Error{Type: "stream_error", Message: "connection reset"}}, + )) sess := session.New("test") err := ag.RunAgentLoop(context.Background(), sess, discardSink{}) @@ -75,12 +67,9 @@ func TestRunAgentLoopRejectsToolCallWithoutID(t *testing.T) { Function: openai.FunctionCall{Name: "read_file", Arguments: `{}`}, }}, } - ag := &Agent{ - Config: &config.ProfileConfig{ModelName: "test"}, - LLM: &fakeLLM{chunks: []openai.ChatCompletionResponse{{ - Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}, - }}}, - } + ag := New("test", &config.ProfileConfig{ModelName: "test"}, fakeStream( + openai.ChatCompletionResponse{Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}}, + )) err := ag.RunAgentLoop(context.Background(), session.New("test"), discardSink{}) if err == nil || !strings.Contains(err.Error(), "missing an id") { @@ -98,12 +87,9 @@ func TestRunAgentLoopReturnsSaveError(t *testing.T) { t.Cleanup(func() { config.ConfigPath = oldConfigPath }) message := openai.ChatCompletionMessage{Role: openai.RoleAssistant, Content: "done"} - ag := &Agent{ - Config: &config.ProfileConfig{ModelName: "test"}, - LLM: &fakeLLM{chunks: []openai.ChatCompletionResponse{{ - Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}, - }}}, - } + ag := New("test", &config.ProfileConfig{ModelName: "test"}, fakeStream( + openai.ChatCompletionResponse{Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}}, + )) err := ag.RunAgentLoop(context.Background(), session.New("test"), discardSink{}) if err == nil || !strings.Contains(err.Error(), "save session") { diff --git a/agent/manager.go b/agent/manager.go index dd74d02..64c9fce 100644 --- a/agent/manager.go +++ b/agent/manager.go @@ -11,6 +11,7 @@ import ( "time" "github.com/lsongdev/miya-agents/acp" + "github.com/lsongdev/miya-agents/anthropic" "github.com/lsongdev/miya-agents/config" "github.com/lsongdev/miya-agents/openai" "github.com/lsongdev/miya-agents/session" @@ -30,38 +31,39 @@ func NewAgentManager(config *config.Config) *Manager { } } -func (m *Manager) UseAgent(name string) (a *Agent, err error) { - ac, ok := m.config.Profiles[name] +func (m *Manager) UseAgent(name string) (*Agent, error) { + profile, ok := m.config.Profiles[name] if !ok { - err = fmt.Errorf("agent not found: %s", name) - return + return nil, fmt.Errorf("agent not found: %s", name) } - pc, ok := m.config.Providers[ac.Provider] + provider, ok := m.config.Providers[profile.Provider] if !ok { - err = fmt.Errorf("provider not found: %s", ac.Provider) - return + return nil, fmt.Errorf("provider not found: %s", profile.Provider) } - llm, err := openai.NewClient(&openai.Configuration{ - API: pc.APIBase, - APIKey: pc.APIKey, - }) - if err != nil { - return - } - a = &Agent{ - Name: name, - LLM: llm, - Config: ac, - toolsMap: make(map[string]openai.Tool), - toolsDefs: []openai.ToolDef{}, + + var stream StreamFunc + switch strings.ToLower(strings.TrimSpace(provider.Type)) { + case "", "openai": + client, err := openai.NewClient(&openai.Configuration{API: provider.APIBase, APIKey: provider.APIKey}) + if err != nil { + return nil, err + } + stream = client.CreateChatCompletionStream + case "anthropic": + client := anthropic.NewClient(&anthropic.Configuration{API: provider.APIBase, APIKey: provider.APIKey}) + stream = client.CreateChatCompletionStream + default: + return nil, fmt.Errorf("unsupported provider type %q", provider.Type) } + + a := New(name, profile, stream) a.BuildTools() mcpManager := tools.NewMcpManager(m.config.McpServers) for _, tool := range mcpManager.Tools { - a.AddTool(tool) + a.Use(tool) } - a.AddTool(tools.NewSubagentTool(m)) - return + a.Use(tools.NewSubagentTool(m)) + return a, nil } func (m *Manager) defaultAgentName() (string, error) { diff --git a/agent/manager_test.go b/agent/manager_test.go index 4051977..fbc19a2 100644 --- a/agent/manager_test.go +++ b/agent/manager_test.go @@ -11,6 +11,7 @@ import ( "github.com/lsongdev/miya-agents/acp" "github.com/lsongdev/miya-agents/config" "github.com/lsongdev/miya-agents/mcp" + "github.com/lsongdev/miya-agents/openai" "github.com/lsongdev/miya-agents/session" ) @@ -164,8 +165,61 @@ func TestUseAgentIncludesConfiguredMCPTools(t *testing.T) { if err != nil { t.Fatalf("UseAgent: %v", err) } - if _, ok := ag.toolsMap["mcp_coffee_queryShopList"]; !ok { - t.Fatalf("missing MCP tool; tools = %#v", ag.toolsMap) + if _, ok := ag.tool("mcp_coffee_queryShopList"); !ok { + t.Fatalf("missing MCP tool; tools = %#v", ag.tools) + } +} + +func TestUseAgentUsesAnthropicProvider(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/messages" { + t.Fatalf("path = %q", r.URL.Path) + } + if got := r.Header.Get("x-api-key"); got != "test-key" { + t.Fatalf("x-api-key = %q", got) + } + var req map[string]any + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + t.Fatal(err) + } + if req["model"] != "claude-test" { + t.Fatalf("model = %#v", req["model"]) + } + + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_1\",\"type\":\"message\",\"role\":\"assistant\",\"model\":\"claude-test\",\"usage\":{\"input_tokens\":1,\"output_tokens\":0}}}\n\n") + fmt.Fprint(w, "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n") + fmt.Fprint(w, "event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"},\"usage\":{\"output_tokens\":1}}\n\n") + fmt.Fprint(w, "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n") + })) + defer server.Close() + + m := NewAgentManager(&config.Config{ + Profiles: map[string]*config.ProfileConfig{ + "default": {Provider: "claude", ModelName: "claude-test", Workspace: t.TempDir()}, + }, + Providers: map[string]*config.ProviderConfig{ + "claude": {Type: "anthropic", APIBase: server.URL, APIKey: "test-key"}, + }, + }) + ag, err := m.UseAgent("default") + if err != nil { + t.Fatal(err) + } + stream, err := ag.Stream(context.Background(), &openai.ChatCompletionRequest{ + Model: "claude-test", Messages: []openai.ChatCompletionMessage{openai.UserMessage("hi")}, Stream: true, + }) + if err != nil { + t.Fatal(err) + } + builder := openai.NewMessageBuilder() + for chunk := range stream { + if message := chunk.GetMessage(); message != nil { + builder.Update(*message) + } + } + if got := builder.Build().Content; got != "hello" { + t.Fatalf("content = %q", got) } } diff --git a/anthropic/client.go b/anthropic/client.go index 3468782..4d1ae6b 100644 --- a/anthropic/client.go +++ b/anthropic/client.go @@ -71,7 +71,12 @@ func (c *Client) applyHeaders(req *http.Request) error { // body. Endpoint may be absolute or relative to the configured API URL. func (c *Client) NewRequest(ctx context.Context, method, endpoint string, body io.Reader) (*http.Request, error) { if !strings.HasPrefix(endpoint, "http://") && !strings.HasPrefix(endpoint, "https://") { - endpoint = strings.TrimRight(c.config.API, "/") + "/" + strings.TrimLeft(endpoint, "/") + base := strings.TrimRight(c.config.API, "/") + endpoint = strings.TrimLeft(endpoint, "/") + if strings.HasSuffix(base, "/v1") { + endpoint = strings.TrimPrefix(endpoint, "v1/") + } + endpoint = base + "/" + endpoint } request, err := http.NewRequestWithContext(ctx, method, endpoint, body) if err != nil { @@ -93,7 +98,7 @@ func (c *Client) Do(request *http.Request) (*http.Response, error) { return c.client.Do(request) } -func (c *Client) makeRequest(ctx context.Context, method, path string, body interface{}) (*http.Response, error) { +func (c *Client) makeRequest(ctx context.Context, method, path string, body any) (*http.Response, error) { var bodyReader io.Reader if body != nil { data, err := json.Marshal(body) @@ -123,7 +128,7 @@ func (c *Client) makeRequest(ctx context.Context, method, path string, body inte } func (c *Client) Models(ctx context.Context) ([]Model, error) { - resp, err := c.makeRequest(ctx, "GET", "/v1/models", nil) + resp, err := c.makeRequest(ctx, http.MethodGet, "/v1/models", nil) if err != nil { return nil, fmt.Errorf("failed to fetch models: %w", err) } @@ -140,7 +145,7 @@ func (c *Client) Models(ctx context.Context) ([]Model, error) { // CreateMessage sends a non-streaming message request. func (c *Client) CreateMessage(ctx context.Context, req *Request) (*Response, error) { - resp, err := c.makeRequest(ctx, "POST", "/v1/messages", req) + resp, err := c.makeRequest(ctx, http.MethodPost, "/v1/messages", req) if err != nil { return nil, err } @@ -158,17 +163,16 @@ func (c *Client) CreateMessage(ctx context.Context, req *Request) (*Response, er return &result, nil } -// MessageStream represents a streaming response channel. +// MessageStream ends when Events closes. type MessageStream struct { Events chan Event - Done chan struct{} } -// CreateMessageStream sends a streaming message request and returns a channel of SSE events. +// CreateMessageStream sends a streaming message request and returns Anthropic SSE events. func (c *Client) CreateMessageStream(ctx context.Context, req *Request) (*MessageStream, error) { - req.Stream = true - - data, err := json.Marshal(req) + streamReq := *req + streamReq.Stream = true + data, err := json.Marshal(&streamReq) if err != nil { return nil, fmt.Errorf("json marshal error: %w", err) } @@ -183,14 +187,18 @@ func (c *Client) CreateMessageStream(ctx context.Context, req *Request) (*Messag return nil, err } - out := &MessageStream{ - Events: make(chan Event), - Done: make(chan struct{}), - } - + out := &MessageStream{Events: make(chan Event)} go func() { defer close(out.Events) - defer close(out.Done) + emitError := func(err error) { + if err == nil || ctx.Err() != nil { + return + } + select { + case out.Events <- Event{Type: "error", Error: &APIError{Type: "stream_error", Message: err.Error()}}: + case <-ctx.Done(): + } + } for { select { @@ -198,15 +206,33 @@ func (c *Client) CreateMessageStream(ctx context.Context, req *Request) (*Messag return case evt, ok := <-stream.Events: if !ok { + if err, ok := <-stream.Err(); ok && err != nil { + emitError(err) + } else { + emitError(io.ErrUnexpectedEOF) + } return } var event Event if err := json.Unmarshal([]byte(evt.Data), &event); err != nil { - continue + emitError(fmt.Errorf("decode stream event: %w", err)) + return } event.Type = evt.Type - out.Events <- event - case <-stream.Err(): + select { + case out.Events <- event: + case <-ctx.Done(): + return + } + if event.Type == "message_stop" { + return + } + case err, ok := <-stream.Err(): + if ok && err != nil { + emitError(err) + } else { + emitError(io.ErrUnexpectedEOF) + } return } } diff --git a/anthropic/client_test.go b/anthropic/client_test.go index 8d761d9..6da8cf1 100644 --- a/anthropic/client_test.go +++ b/anthropic/client_test.go @@ -141,7 +141,7 @@ func TestCreateMessage(t *testing.T) { resp, err := c.CreateMessage(context.Background(), &Request{ Model: "claude-3", MaxTokens: 100, - Messages: []Message{{Role: "user", Content: "hi"}}, + Messages: []Message{TextMessage("user", "hi")}, }) if err != nil { t.Fatalf("unexpected error: %v", err) @@ -171,7 +171,7 @@ func TestCreateMessage_Error(t *testing.T) { _, err := c.CreateMessage(context.Background(), &Request{ Model: "claude-3", MaxTokens: 100, - Messages: []Message{{Role: "user", Content: "hi"}}, + Messages: []Message{TextMessage("user", "hi")}, }) if err == nil { t.Fatal("expected error, got nil") @@ -186,7 +186,7 @@ func TestCreateMessage_ConnectionError(t *testing.T) { _, err := c.CreateMessage(context.Background(), &Request{ Model: "claude-3", MaxTokens: 100, - Messages: []Message{{Role: "user", Content: "hi"}}, + Messages: []Message{TextMessage("user", "hi")}, }) if err == nil { t.Fatal("expected error, got nil") @@ -214,7 +214,7 @@ func TestCreateMessage_MultipleContentBlocks(t *testing.T) { resp, err := c.CreateMessage(context.Background(), &Request{ Model: "claude-3", MaxTokens: 100, - Messages: []Message{{Role: "user", Content: "explain quantum physics"}}, + Messages: []Message{TextMessage("user", "explain quantum physics")}, }) if err != nil { t.Fatalf("unexpected error: %v", err) @@ -274,14 +274,18 @@ func TestCreateMessageStream(t *testing.T) { defer server.Close() c := NewClient(&Configuration{API: server.URL, APIKey: "test-key"}) - ms, err := c.CreateMessageStream(context.Background(), &Request{ + req := &Request{ Model: "claude-3", MaxTokens: 100, - Messages: []Message{{Role: "user", Content: "hi"}}, - }) + Messages: []Message{TextMessage("user", "hi")}, + } + ms, err := c.CreateMessageStream(context.Background(), req) if err != nil { t.Fatalf("unexpected error: %v", err) } + if req.Stream { + t.Fatal("CreateMessageStream mutated request") + } var texts []string var eventTypes []string @@ -307,7 +311,7 @@ func TestCreateMessageStream(t *testing.T) { expectedTypes := []string{"message_start", "content_block_delta", "content_block_delta", "message_delta", "message_stop"} if len(eventTypes) != len(expectedTypes) { - t.Fatalf("expected %d events, got %d", len(expectedTypes), len(eventTypes)) + t.Fatalf("expected %d events, got %d: %v", len(expectedTypes), len(eventTypes), eventTypes) } for i, et := range expectedTypes { if eventTypes[i] != et { @@ -329,7 +333,7 @@ func TestCreateMessageStream_ContextCancel(t *testing.T) { ms, err := c.CreateMessageStream(ctx, &Request{ Model: "claude-3", MaxTokens: 100, - Messages: []Message{{Role: "user", Content: "hi"}}, + Messages: []Message{TextMessage("user", "hi")}, }) if err != nil { t.Fatalf("unexpected error: %v", err) @@ -347,7 +351,7 @@ func TestCreateMessageStream_ConnectionError(t *testing.T) { _, err := c.CreateMessageStream(context.Background(), &Request{ Model: "claude-3", MaxTokens: 100, - Messages: []Message{{Role: "user", Content: "hi"}}, + Messages: []Message{TextMessage("user", "hi")}, }) if err == nil { t.Fatal("expected error, got nil") @@ -366,18 +370,18 @@ func TestCreateMessageStream_JSONError(t *testing.T) { ms, err := c.CreateMessageStream(context.Background(), &Request{ Model: "claude-3", MaxTokens: 100, - Messages: []Message{{Role: "user", Content: "hi"}}, + Messages: []Message{TextMessage("user", "hi")}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } - var count int - for range ms.Events { - count++ + event, ok := <-ms.Events + if !ok || event.Type != "error" || event.Error == nil { + t.Fatalf("event = %#v", event) } - if count != 1 { - t.Fatalf("expected 1 event (message_stop), got %d", count) + if _, ok := <-ms.Events; ok { + t.Fatal("stream continued after decode error") } } diff --git a/anthropic/client_url_test.go b/anthropic/client_url_test.go new file mode 100644 index 0000000..7096d79 --- /dev/null +++ b/anthropic/client_url_test.go @@ -0,0 +1,18 @@ +package anthropic + +import ( + "context" + "net/http" + "testing" +) + +func TestNewRequestAcceptsVersionedAPIBase(t *testing.T) { + client := NewClient(&Configuration{API: "https://api.anthropic.com/v1", APIKey: "test-key"}) + request, err := client.NewRequest(context.Background(), http.MethodGet, "/v1/models", nil) + if err != nil { + t.Fatal(err) + } + if got := request.URL.String(); got != "https://api.anthropic.com/v1/models" { + t.Fatalf("URL = %q", got) + } +} diff --git a/anthropic/openai.go b/anthropic/openai.go index 4af7bd0..92c238b 100644 --- a/anthropic/openai.go +++ b/anthropic/openai.go @@ -10,137 +10,211 @@ import ( "github.com/lsongdev/miya-agents/openai" ) -// ToRequest converts an OpenAI chat completion request to Anthropic messages format. +// NewAnthropicRequestFromChatCompletionRequest converts the agent's OpenAI-shaped +// conversation into a native Anthropic Messages request. func NewAnthropicRequestFromChatCompletionRequest(req *openai.ChatCompletionRequest) *Request { - anthropicReq := &Request{ - Model: req.Model, - Stream: req.Stream, + out := &Request{ + Model: req.Model, + Stream: req.Stream, + TopP: req.TopP, + Temperature: req.Temperature, + StopSequences: req.Stop, + MaxTokens: req.MaxTokens, } - if req.MaxTokens > 0 { - anthropicReq.MaxTokens = req.MaxTokens - } else { - anthropicReq.MaxTokens = 4096 + if out.MaxTokens == 0 { + out.MaxTokens = 4096 + } + for _, tool := range req.Tools { + out.Tools = append(out.Tools, Tool{ + Name: tool.Function.Name, + Description: tool.Function.Description, + InputSchema: tool.Function.Parameters, + }) } - anthropicReq.TopP = req.TopP - anthropicReq.Temperature = req.Temperature - anthropicReq.StopSequences = req.Stop - var systemParts []string + var system []string for _, msg := range req.Messages { switch msg.Role { case openai.RoleSystem: - systemParts = append(systemParts, msg.Content) + system = append(system, msg.Content) case openai.RoleUser: - anthropicReq.Messages = append(anthropicReq.Messages, Message{ - Role: "user", - Content: msg.Content, - }) + out.Messages = append(out.Messages, TextMessage("user", msg.Content)) case openai.RoleAssistant: - anthropicReq.Messages = append(anthropicReq.Messages, Message{ - Role: "assistant", - Content: msg.Content, - }) + blocks := make([]ContentBlock, 0, 1+len(msg.ToolCalls)) + if msg.Content != "" { + blocks = append(blocks, ContentBlock{Type: "text", Text: msg.Content}) + } + for _, call := range msg.ToolCalls { + input := json.RawMessage(call.Function.Arguments) + if len(input) == 0 { + input = json.RawMessage(`{}`) + } + blocks = append(blocks, ContentBlock{ + Type: "tool_use", + ID: call.ID, + Name: call.Function.Name, + Input: input, + }) + } + if len(blocks) > 0 { + out.Messages = append(out.Messages, Message{Role: "assistant", Blocks: blocks}) + } case openai.RoleTool: - toolCtx := fmt.Sprintf("Tool result (%s): %s", msg.Name, msg.Content) - anthropicReq.Messages = append(anthropicReq.Messages, Message{ - Role: "user", - Content: toolCtx, - }) + block := ContentBlock{ + Type: "tool_result", + ToolUseID: msg.ToolCallID, + Content: msg.Content, + } + last := len(out.Messages) - 1 + if last >= 0 && toolResultsOnly(out.Messages[last]) { + out.Messages[last].Blocks = append(out.Messages[last].Blocks, block) + } else { + out.Messages = append(out.Messages, Message{Role: "user", Blocks: []ContentBlock{block}}) + } } } - - if len(systemParts) > 0 { - anthropicReq.System = strings.Join(systemParts, "\n\n") - } - - return anthropicReq + out.System = strings.Join(system, "\n\n") + return out } -// ToAnthropicResponse converts an OpenAI chat completion response to Anthropic format. -func NewAnthropicResponseFromChatCompletionResponse(oaiResp *openai.ChatCompletionResponse) *Response { - anthResp := &Response{ - ID: oaiResp.ID, - Type: "message", - Role: "assistant", - Model: oaiResp.Model, +func toolResultsOnly(message Message) bool { + if message.Role != "user" || len(message.Blocks) == 0 { + return false } - - if len(oaiResp.Choices) > 0 { - msg := oaiResp.Choices[0].Message - if msg.Content != "" { - anthResp.Content = append(anthResp.Content, ContentBlock{ - Type: "text", - Text: msg.Content, - }) - } - if msg.ReasoningContent != "" { - anthResp.Content = append(anthResp.Content, ContentBlock{ - Type: "thinking", - Thinking: msg.ReasoningContent, - }) + for _, block := range message.Blocks { + if block.Type != "tool_result" { + return false } - anthResp.StopReason = MapStopReasonReverse(oaiResp.Choices[0].FinishReason) } + return true +} - if oaiResp.Usage != nil && (oaiResp.Usage.PromptTokens > 0 || oaiResp.Usage.CompletionTokens > 0) { - anthResp.Usage = Usage{ - InputTokens: oaiResp.Usage.PromptTokens, - OutputTokens: oaiResp.Usage.CompletionTokens, +// NewAnthropicResponseFromChatCompletionResponse converts an OpenAI chat response +// to Anthropic's response shape. +func NewAnthropicResponseFromChatCompletionResponse(resp *openai.ChatCompletionResponse) *Response { + out := &Response{ID: resp.ID, Type: "message", Role: "assistant", Model: resp.Model} + if choice := resp.GetFirstChoice(); choice != nil { + if msg := choice.Message; msg != nil { + if msg.Content != "" { + out.Content = append(out.Content, ContentBlock{Type: "text", Text: msg.Content}) + } + if msg.ReasoningContent != "" { + out.Content = append(out.Content, ContentBlock{Type: "thinking", Thinking: msg.ReasoningContent}) + } + for _, call := range msg.ToolCalls { + input := json.RawMessage(call.Function.Arguments) + if len(input) == 0 { + input = json.RawMessage(`{}`) + } + out.Content = append(out.Content, ContentBlock{ + Type: "tool_use", ID: call.ID, Name: call.Function.Name, Input: input, + }) + } } + out.StopReason = MapStopReasonReverse(choice.FinishReason) } - - return anthResp + if resp.Usage != nil { + out.Usage = Usage{InputTokens: resp.Usage.PromptTokens, OutputTokens: resp.Usage.CompletionTokens} + } + return out } -// ToOpenAIResponse converts an Anthropic message response to OpenAI chat completion format. -func NewChatCompletionResponseFromAnthropicResponse(anthResp *Response) *openai.ChatCompletionResponse { - var content, reasoning string - for _, block := range anthResp.Content { +// NewChatCompletionResponseFromAnthropicResponse converts an Anthropic response +// into the shape consumed by the agent loop. +func NewChatCompletionResponseFromAnthropicResponse(resp *Response) *openai.ChatCompletionResponse { + message := openai.ChatCompletionMessage{Role: openai.RoleAssistant} + for _, block := range resp.Content { switch block.Type { case "text": - content += block.Text + message.Content += block.Text case "thinking": - reasoning += block.Thinking + message.ReasoningContent += block.Thinking + case "tool_use": + args := string(block.Input) + if args == "" { + args = `{}` + } + message.ToolCalls = append(message.ToolCalls, openai.ToolCall{ + Index: len(message.ToolCalls), + ID: block.ID, + Type: "function", + Function: openai.FunctionCall{ + Name: block.Name, Arguments: args, + }, + }) } } - oaiResp := openai.NewChatCompletionResponse(anthResp.ID, anthResp.Model, content, reasoning) - oaiResp.Choices[0].FinishReason = MapStopReason(anthResp.StopReason) - oaiResp.Usage = &openai.CompletionUsage{ - PromptTokens: anthResp.Usage.InputTokens, - CompletionTokens: anthResp.Usage.OutputTokens, - TotalTokens: anthResp.Usage.InputTokens + anthResp.Usage.OutputTokens, + out := openai.NewChatCompletionResponse(resp.ID, resp.Model, message.Content, message.ReasoningContent) + out.Choices[0].Message = &message + out.Choices[0].FinishReason = MapStopReason(resp.StopReason) + out.Usage = &openai.CompletionUsage{ + PromptTokens: resp.Usage.InputTokens, + CompletionTokens: resp.Usage.OutputTokens, + TotalTokens: resp.Usage.InputTokens + resp.Usage.OutputTokens, } - return oaiResp + return out } -// ToOpenAIRequest converts an Anthropic messages request to OpenAI chat completion format. +// NewChatCompletionRequestFromAnthropicRequest converts an Anthropic request to +// OpenAI chat-completion format. func NewChatCompletionRequestFromAnthropicRequest(req *Request) *openai.ChatCompletionRequest { - oaiReq := &openai.ChatCompletionRequest{ - Model: req.Model, - MaxTokens: req.MaxTokens, - Stream: req.Stream, + out := &openai.ChatCompletionRequest{ + Model: req.Model, + MaxTokens: req.MaxTokens, + Stream: req.Stream, + TopP: req.TopP, + Stop: req.StopSequences, + Temperature: req.Temperature, } - - oaiReq.TopP = req.TopP - oaiReq.Stop = req.StopSequences - oaiReq.MaxTokens = req.MaxTokens - oaiReq.Temperature = req.Temperature - if req.System != "" { - oaiReq.Messages = append(oaiReq.Messages, openai.ChatCompletionMessage{ - Role: openai.RoleSystem, - Content: req.System, + out.Messages = append(out.Messages, openai.SystemMessage(req.System)) + } + for _, tool := range req.Tools { + out.Tools = append(out.Tools, openai.ToolDef{ + Type: "function", + Function: openai.FunctionDef{ + Name: tool.Name, Description: tool.Description, Parameters: tool.InputSchema, + }, }) } - for _, msg := range req.Messages { - oaiReq.Messages = append(oaiReq.Messages, openai.ChatCompletionMessage{ - Role: msg.Role, - Content: msg.Content, - }) + switch msg.Role { + case "assistant": + converted := openai.ChatCompletionMessage{Role: openai.RoleAssistant} + for _, block := range msg.contentBlocks() { + switch block.Type { + case "text": + converted.Content += block.Text + case "thinking": + converted.ReasoningContent += block.Thinking + case "tool_use": + args := string(block.Input) + if args == "" { + args = `{}` + } + converted.ToolCalls = append(converted.ToolCalls, openai.ToolCall{ + Index: len(converted.ToolCalls), ID: block.ID, Type: "function", + Function: openai.FunctionCall{Name: block.Name, Arguments: args}, + }) + } + } + out.Messages = append(out.Messages, converted) + case "user": + var text strings.Builder + for _, block := range msg.contentBlocks() { + switch block.Type { + case "text": + text.WriteString(block.Text) + case "tool_result": + out.Messages = append(out.Messages, openai.ToolResultMessage(block.ToolUseID, "", block.Content)) + } + } + if text.Len() > 0 { + out.Messages = append(out.Messages, openai.UserMessage(text.String())) + } + } } - - return oaiReq + return out } // MapStopReason maps Anthropic stop reasons to OpenAI finish reasons. @@ -158,8 +232,7 @@ func MapStopReason(reason string) string { } func (c *Client) CreateChatCompletion(ctx context.Context, req *openai.ChatCompletionRequest) (*openai.ChatCompletionResponse, error) { - anthReq := NewAnthropicRequestFromChatCompletionRequest(req) - resp, err := c.CreateMessage(ctx, anthReq) + resp, err := c.CreateMessage(ctx, NewAnthropicRequestFromChatCompletionRequest(req)) if err != nil { return nil, err } @@ -176,72 +249,83 @@ func (c *Client) CreateChatCompletionStream(ctx context.Context, req *openai.Cha return AnthropicStreamToChatCompletionStream(stream, nil), nil } -// AnthropicStreamToChatCompletionStream converts an Anthropic message stream to OpenAI chat completion chunks. -// If onEvent is provided, each event is passed to it before conversion. +// AnthropicStreamToChatCompletionStream converts Anthropic SSE events to the +// stream shape consumed by the agent loop. func AnthropicStreamToChatCompletionStream(stream *MessageStream, onEvent func(Event)) <-chan openai.ChatCompletionResponse { ch := make(chan openai.ChatCompletionResponse) go func() { defer close(ch) var messageID, model string - var sentFirst bool var inputTokens int for event := range stream.Events { if onEvent != nil { onEvent(event) } switch event.Type { + case "error": + if event.Error != nil { + ch <- openai.ChatCompletionResponse{Error: &openai.Error{Type: event.Error.Type, Message: event.Error.Message}} + } + return case "message_start": if event.Message != nil { messageID = event.Message.ID model = event.Message.Model inputTokens = event.Message.Usage.InputTokens } + case "content_block_start": + if event.ContentBlock == nil || event.ContentBlock.Type != "tool_use" { + continue + } + index := 0 + if event.Index != nil { + index = *event.Index + } + ch <- chatDelta(messageID, model, openai.ChatCompletionMessage{ + ToolCalls: []openai.ToolCall{{ + Index: index, ID: event.ContentBlock.ID, Type: "function", + Function: openai.FunctionCall{Name: event.ContentBlock.Name}, + }}, + }) case "content_block_delta": var delta Delta if err := json.Unmarshal(event.Delta, &delta); err != nil { continue } - msg := &openai.ChatCompletionMessage{} + message := openai.ChatCompletionMessage{} switch delta.Type { case "text_delta": - msg.Content = delta.Text + message.Content = delta.Text case "thinking_delta": - msg.ReasoningContent = delta.Thinking + message.ReasoningContent = delta.Thinking + case "input_json_delta": + index := 0 + if event.Index != nil { + index = *event.Index + } + message.ToolCalls = []openai.ToolCall{{ + Index: index, + Function: openai.FunctionCall{Arguments: delta.PartialJSON}, + }} + } + if !message.IsEmpty() { + ch <- chatDelta(messageID, model, message) } - if msg.Content == "" && msg.ReasoningContent == "" { + case "message_delta": + var delta Delta + if err := json.Unmarshal(event.Delta, &delta); err != nil || delta.StopReason == "" { continue } - if !sentFirst { - msg.Role = openai.RoleAssistant - sentFirst = true + outputTokens := 0 + if event.Usage != nil { + outputTokens = event.Usage.OutputTokens } ch <- openai.ChatCompletionResponse{ - ID: messageID, - Model: model, - Object: "chat.completion.chunk", - Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: msg}}, - } - case "message_delta": - var delta Delta - if err := json.Unmarshal(event.Delta, &delta); err == nil && delta.StopReason != "" { - outputTokens := 0 - if event.Usage != nil { - outputTokens = event.Usage.OutputTokens - } - ch <- openai.ChatCompletionResponse{ - ID: messageID, - Model: model, - Object: "chat.completion.chunk", - Choices: []openai.ChatCompletionChoice{{ - Index: 0, - FinishReason: MapStopReason(delta.StopReason), - }}, - Usage: &openai.CompletionUsage{ - PromptTokens: inputTokens, - CompletionTokens: outputTokens, - TotalTokens: inputTokens + outputTokens, - }, - } + ID: messageID, Model: model, Object: "chat.completion.chunk", + Choices: []openai.ChatCompletionChoice{{Index: 0, FinishReason: MapStopReason(delta.StopReason)}}, + Usage: &openai.CompletionUsage{ + PromptTokens: inputTokens, CompletionTokens: outputTokens, TotalTokens: inputTokens + outputTokens, + }, } case "message_stop": return @@ -251,6 +335,13 @@ func AnthropicStreamToChatCompletionStream(stream *MessageStream, onEvent func(E return ch } +func chatDelta(id, model string, message openai.ChatCompletionMessage) openai.ChatCompletionResponse { + return openai.ChatCompletionResponse{ + ID: id, Model: model, Object: "chat.completion.chunk", + Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}, + } +} + // OpenAIStreamToAnthropicStream converts an OpenAI chat completion stream to Anthropic SSE events. // The onChunk callback receives each chunk before it is processed (useful for assembling the final response). // Returns the final assembled ChatCompletionResponse after the stream completes. @@ -330,9 +421,7 @@ func OpenAIStreamToAnthropicStream(chunks <-chan openai.ChatCompletionResponse, if !hasContent { sendContentBlockStart() } - anthStream.SendContentBlockDelta(0, Delta{ - Type: "text_delta", Text: delta.Content, - }) + anthStream.SendContentBlockDelta(0, Delta{Type: "text_delta", Text: delta.Content}) } if choice.FinishReason != "" { diff --git a/anthropic/openai_test.go b/anthropic/openai_test.go new file mode 100644 index 0000000..82c3ccb --- /dev/null +++ b/anthropic/openai_test.go @@ -0,0 +1,85 @@ +package anthropic + +import ( + "encoding/json" + "testing" + + "github.com/lsongdev/miya-agents/openai" +) + +func TestChatCompletionRequestPreservesTools(t *testing.T) { + req := &openai.ChatCompletionRequest{ + Model: "claude-test", + Tools: []openai.ToolDef{{ + Type: "function", + Function: openai.FunctionDef{ + Name: "read_file", Description: "read a file", + Parameters: map[string]any{"type": "object"}, + }, + }}, + Messages: []openai.ChatCompletionMessage{ + openai.UserMessage("read README.md"), + { + Role: openai.RoleAssistant, + ToolCalls: []openai.ToolCall{{ + ID: "call_1", Type: "function", + Function: openai.FunctionCall{Name: "read_file", Arguments: `{"path":"README.md"}`}, + }}, + }, + openai.ToolResultMessage("call_1", "read_file", "hello"), + }, + } + + got := NewAnthropicRequestFromChatCompletionRequest(req) + if len(got.Tools) != 1 || got.Tools[0].Name != "read_file" { + t.Fatalf("tools = %#v", got.Tools) + } + if len(got.Messages) != 3 { + t.Fatalf("messages = %#v", got.Messages) + } + use := got.Messages[1].Blocks[0] + if use.Type != "tool_use" || use.ID != "call_1" || use.Name != "read_file" || string(use.Input) != `{"path":"README.md"}` { + t.Fatalf("tool use = %#v", use) + } + result := got.Messages[2].Blocks[0] + if result.Type != "tool_result" || result.ToolUseID != "call_1" || result.Content != "hello" { + t.Fatalf("tool result = %#v", result) + } +} + +func TestAnthropicToolStreamBuildsToolCall(t *testing.T) { + events := make(chan Event, 5) + stream := &MessageStream{Events: events} + index := 0 + partial1, _ := json.Marshal(Delta{Type: "input_json_delta", PartialJSON: `{"path":"`}) + partial2, _ := json.Marshal(Delta{Type: "input_json_delta", PartialJSON: `README.md"}`}) + stop, _ := json.Marshal(Delta{StopReason: "tool_use"}) + events <- Event{Type: "message_start", Message: &MessageStart{ID: "msg_1", Model: "claude-test"}} + events <- Event{Type: "content_block_start", Index: &index, ContentBlock: &ContentBlock{Type: "tool_use", ID: "call_1", Name: "read_file"}} + events <- Event{Type: "content_block_delta", Index: &index, Delta: partial1} + events <- Event{Type: "content_block_delta", Index: &index, Delta: partial2} + events <- Event{Type: "message_delta", Delta: stop} + close(events) + + builder := openai.NewMessageBuilder() + finish := "" + for chunk := range AnthropicStreamToChatCompletionStream(stream, nil) { + if message := chunk.GetMessage(); message != nil { + builder.Update(*message) + } + if choice := chunk.GetFirstChoice(); choice != nil && choice.FinishReason != "" { + finish = choice.FinishReason + } + } + message := builder.Build() + if len(message.ToolCalls) != 1 { + t.Fatalf("tool calls = %#v", message.ToolCalls) + } + call := message.ToolCalls[0] + if call.ID != "call_1" || call.Function.Name != "read_file" || call.Function.Arguments != `{"path":"README.md"}` { + t.Fatalf("tool call = %#v", call) + } + if finish != "tool_calls" { + t.Fatalf("finish reason = %q", finish) + } +} diff --git a/anthropic/types.go b/anthropic/types.go index 3b74518..e725995 100644 --- a/anthropic/types.go +++ b/anthropic/types.go @@ -12,6 +12,7 @@ type Request struct { MaxTokens int `json:"max_tokens"` Messages []Message `json:"messages"` System string `json:"system,omitempty"` + Tools []Tool `json:"tools,omitempty"` Stream bool `json:"stream,omitempty"` Temperature *float64 `json:"temperature,omitempty"` TopP *float64 `json:"top_p,omitempty"` @@ -52,43 +53,63 @@ func (r *Request) UnmarshalJSON(data []byte) error { return nil } -// Message is a single message in an Anthropic request. +// Tool describes a client tool available to Claude. +type Tool struct { + Name string `json:"name"` + Description string `json:"description,omitempty"` + InputSchema map[string]any `json:"input_schema"` +} + +// Message keeps the common text form terse while exposing Blocks for tool use +// and other structured Anthropic content. Set either Content or Blocks. type Message struct { - Role string `json:"role"` - Content string `json:"content"` + Role string `json:"role"` + Content string `json:"-"` + Blocks []ContentBlock `json:"-"` +} + +func TextMessage(role, text string) Message { + return Message{Role: role, Content: text} +} + +func (m Message) MarshalJSON() ([]byte, error) { + var content any = m.Content + if m.Blocks != nil { + content = m.Blocks + } + return json.Marshal(struct { + Role string `json:"role"` + Content any `json:"content"` + }{m.Role, content}) } func (m *Message) UnmarshalJSON(data []byte) error { - type Alias Message - aux := &struct { + var wire struct { + Role string `json:"role"` Content json.RawMessage `json:"content"` - *Alias - }{Alias: (*Alias)(m)} - if err := json.Unmarshal(data, aux); err != nil { - return err } - if len(aux.Content) == 0 { - return nil + if err := json.Unmarshal(data, &wire); err != nil { + return err } - var s string - if err := json.Unmarshal(aux.Content, &s); err == nil { - m.Content = s + m.Role = wire.Role + if err := json.Unmarshal(wire.Content, &m.Content); err == nil { + m.Blocks = nil return nil } - var blocks []struct { - Type string `json:"type"` - Text string `json:"text,omitempty"` - } - if err := json.Unmarshal(aux.Content, &blocks); err != nil { + m.Content = "" + if err := json.Unmarshal(wire.Content, &m.Blocks); err != nil { return fmt.Errorf("content must be a string or array of content blocks: %v", err) } - var texts []string - for _, b := range blocks { - if b.Type == "text" { - texts = append(texts, b.Text) - } + return nil +} + +func (m Message) contentBlocks() []ContentBlock { + if m.Blocks != nil { + return m.Blocks + } + if m.Content != "" { + return []ContentBlock{{Type: "text", Text: m.Content}} } - m.Content = strings.Join(texts, "\n") return nil } @@ -104,11 +125,17 @@ type Response struct { StopSequence string `json:"stop_sequence,omitempty"` } -// ContentBlock is a block of content in an Anthropic response. +// ContentBlock is a block in an Anthropic message or response. type ContentBlock struct { - Type string `json:"type"` // "text", "thinking", "redacted_thinking", "tool_use" - Text string `json:"text"` - Thinking string `json:"thinking,omitempty"` + Type string `json:"type"` // text, thinking, redacted_thinking, tool_use, tool_result + Text string `json:"text,omitempty"` + Thinking string `json:"thinking,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Input json.RawMessage `json:"input,omitempty"` + ToolUseID string `json:"tool_use_id,omitempty"` + Content string `json:"content,omitempty"` + IsError bool `json:"is_error,omitempty"` } // Usage contains token usage from Anthropic. @@ -143,10 +170,11 @@ type MessageStart struct { Usage Usage `json:"usage"` } -// Delta represents the incremental payload inside a content_block_delta event. +// Delta represents the incremental payload inside streamed Anthropic events. type Delta struct { - Type string `json:"type"` // "text_delta", "thinking_delta", "signature_delta" - Text string `json:"text,omitempty"` - Thinking string `json:"thinking,omitempty"` - StopReason string `json:"stop_reason,omitempty"` + Type string `json:"type"` // text_delta, thinking_delta, input_json_delta, signature_delta + Text string `json:"text,omitempty"` + Thinking string `json:"thinking,omitempty"` + PartialJSON string `json:"partial_json,omitempty"` + StopReason string `json:"stop_reason,omitempty"` } diff --git a/openai/client.go b/openai/client.go index 89bb61e..38b3afa 100644 --- a/openai/client.go +++ b/openai/client.go @@ -13,12 +13,6 @@ import ( "github.com/lsongdev/miya-agents/sse" ) -type ChatClient interface { - CreateChatCompletion(ctx context.Context, request *ChatCompletionRequest) (*ChatCompletionResponse, error) - CreateChatCompletionStream(ctx context.Context, request *ChatCompletionRequest) (<-chan ChatCompletionResponse, error) - CreateEmbeddings(ctx context.Context, request *EmbeddingRequest) (*EmbeddingResponse, error) -} - type Client struct { config *Configuration client *http.Client @@ -100,14 +94,14 @@ func (client *Client) MakeRequest(ctx context.Context, path string, data any) (i return res.Body, nil } -// Model represents a model in the API format +// Model represents a model in the API format. type Model struct { ID string `json:"id"` Object string `json:"object"` OwnedBy string `json:"owned_by"` } -// Models fetches the list of available models from the API +// Models fetches the list of available models from the API. func (client *Client) Models() (models []Model, err error) { body, err := client.MakeRequest(context.Background(), "/models", nil) if err != nil { @@ -125,10 +119,10 @@ func (client *Client) Models() (models []Model, err error) { } func (resp *ChatCompletionResponse) GetFirstChoice() *ChatCompletionChoice { - for _, choice := range resp.Choices { - return &choice + if len(resp.Choices) == 0 { + return nil } - return nil + return &resp.Choices[0] } func (resp *ChatCompletionResponse) GetMessage() *ChatCompletionMessage { @@ -204,7 +198,7 @@ func (m *ChatCompletionMessage) UnmarshalJSON(data []byte) error { } func (m *ChatCompletionMessage) IsEmpty() bool { - return m.Role == "" && m.Content == "" && m.ReasoningContent == "" && (len(m.ToolCalls) == 0) + return m.Role == "" && m.Content == "" && m.ReasoningContent == "" && len(m.ToolCalls) == 0 } func (m *ChatCompletionMessage) HasToolCall() bool { @@ -242,6 +236,8 @@ func (c *Client) CreateEmbeddings(ctx context.Context, request *EmbeddingRequest if err != nil { return nil, err } + defer body.Close() + data, err := io.ReadAll(body) if err != nil { return nil, err @@ -263,6 +259,8 @@ func (c *Client) CreateChatCompletion(ctx context.Context, request *ChatCompleti if err != nil { return nil, err } + defer body.Close() + data, err := io.ReadAll(body) if err != nil { return nil, err @@ -278,23 +276,13 @@ func (c *Client) CreateChatCompletion(ctx context.Context, request *ChatCompleti return &resp, err } -// ChatCompletionStream represents a streaming response channel -type ChatCompletionStream struct { - Error chan error - Response chan ChatCompletionResponse -} - -// Close closes the stream channels -func (stream *ChatCompletionStream) Close() { - close(stream.Error) - close(stream.Response) -} - -// CreateChatCompletionStream creates a streaming chat completion +// CreateChatCompletionStream sends a streaming chat completion request. func (c *Client) CreateChatCompletionStream(ctx context.Context, request *ChatCompletionRequest) (<-chan ChatCompletionResponse, error) { resp := make(chan ChatCompletionResponse) + streamRequest := *request + streamRequest.Stream = true - payload, err := json.Marshal(request) + payload, err := json.Marshal(&streamRequest) if err != nil { return nil, fmt.Errorf("json error: %v", err) } diff --git a/openai/client_refactor_test.go b/openai/client_refactor_test.go new file mode 100644 index 0000000..d5b386d --- /dev/null +++ b/openai/client_refactor_test.go @@ -0,0 +1,49 @@ +package openai + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" +) + +func TestGetFirstChoiceReturnsSliceElement(t *testing.T) { + resp := ChatCompletionResponse{Choices: []ChatCompletionChoice{{FinishReason: "stop"}}} + choice := resp.GetFirstChoice() + if choice == nil { + t.Fatal("missing first choice") + } + choice.FinishReason = "length" + if resp.Choices[0].FinishReason != "length" { + t.Fatal("GetFirstChoice returned a copy") + } +} + +func TestCreateChatCompletionStreamForcesStreamWithoutMutation(t *testing.T) { + seenStream := make(chan bool, 1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req ChatCompletionRequest + _ = json.NewDecoder(r.Body).Decode(&req) + seenStream <- req.Stream + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: [DONE]\n\n") + })) + defer server.Close() + + client, _ := NewClient(&Configuration{API: server.URL}) + request := &ChatCompletionRequest{Model: "gpt-test"} + stream, err := client.CreateChatCompletionStream(context.Background(), request) + if err != nil { + t.Fatal(err) + } + for range stream { + } + if !<-seenStream { + t.Fatal("stream request did not set stream=true") + } + if request.Stream { + t.Fatal("CreateChatCompletionStream mutated its request") + } +}