diff --git a/pkg/llm/mock/mock.go b/pkg/llm/mock/mock.go new file mode 100644 index 0000000..f34f625 --- /dev/null +++ b/pkg/llm/mock/mock.go @@ -0,0 +1,121 @@ +package mock + +import ( + "context" + "iter" + "strings" + + "github.com/VictorVargas/rony-llm-agent/pkg/llm" +) + +// MockLLMClient is a deterministic implementation of llm.LLMClient for testing. +type MockLLMClient struct { + GenerateFunc func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) + StreamFunc func(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] + NameFunc func() string + CapabilitiesFunc func() llm.ProviderCapabilities +} + +func (m *MockLLMClient) Generate(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) { + if m.GenerateFunc != nil { + return m.GenerateFunc(ctx, req) + } + return llm.CompletionResponse{ + Content: "default response", + StopReason: llm.StopReasonEndTurn, + }, nil +} + +func (m *MockLLMClient) Stream(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] { + if m.StreamFunc != nil { + return m.StreamFunc(ctx, req) + } + return func(yield func(llm.StreamChunk, error) bool) { + yield(llm.StreamChunk{ + Delta: "default response", + }, nil) + } +} + +func (m *MockLLMClient) Name() string { + if m.NameFunc != nil { + return m.NameFunc() + } + return "mock" +} + +func (m *MockLLMClient) Capabilities() llm.ProviderCapabilities { + if m.CapabilitiesFunc != nil { + return m.CapabilitiesFunc() + } + return llm.ProviderCapabilities{ + SupportsTools: true, + SupportsVision: false, + SupportsJSON: true, + MaxContextWindow: 128000, + } +} + +// New returns a MockLLMClient with default behaviors. +func New() *MockLLMClient { + return &MockLLMClient{} +} + +// NewWithGenerate returns a MockLLMClient that returns the given response. +func NewWithGenerate(resp llm.CompletionResponse) *MockLLMClient { + return &MockLLMClient{ + GenerateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) { + return resp, nil + }, + } +} + +// NewWithStream returns a MockLLMClient that streams the given chunks. +func NewWithStream(chunks []llm.StreamChunk) *MockLLMClient { + return &MockLLMClient{ + StreamFunc: func(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] { + return func(yield func(llm.StreamChunk, error) bool) { + for _, c := range chunks { + if !yield(c, nil) { + return + } + } + } + }, + } +} + +// NewWithMatch returns a MockLLMClient that matches input messages against patterns. +func NewWithMatch(responses []MatchResponse) *MockLLMClient { + return &MockLLMClient{ + GenerateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) { + for _, r := range responses { + if r.Match == "*" || matchesAny(r.Match, req.Messages) { + return llm.CompletionResponse{ + Content: r.Response, + StopReason: llm.StopReasonEndTurn, + }, nil + } + } + return llm.CompletionResponse{ + Content: "no match", + StopReason: llm.StopReasonEndTurn, + }, nil + }, + } +} + +// MatchResponse pairs a pattern with a response string. +type MatchResponse struct { + Match string // exact string or "*" for wildcard + Response string +} + +func matchesAny(pattern string, msgs []llm.Message) bool { + for _, m := range msgs { + if m.Role == llm.RoleUser && strings.Contains(m.Content, pattern) { + return true + } + } + return false +} diff --git a/pkg/llm/providers/llamacpp/client.go b/pkg/llm/providers/llamacpp/client.go new file mode 100644 index 0000000..314803e --- /dev/null +++ b/pkg/llm/providers/llamacpp/client.go @@ -0,0 +1,322 @@ +package llamacpp + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "iter" + "strings" + + "github.com/VictorVargas/rony-llm-agent/pkg/llm" +) + +// Config holds the settings needed to create a llama.cpp client. +type Config struct { + BaseURL string // defaults to http://localhost:8080/v1 + Model string + Timeout int // request timeout in seconds (0 = default) + TopK int // top-k sampling (0 = default) + TopP float32 + Temperature float32 +} + +// Client implements llm.LLMClient for llama.cpp. +type Client struct { + baseURL string + http *http.Client +} + +// New returns a new llama.cpp client. +func New(cfg Config) (*Client, error) { + baseURL := cfg.BaseURL + if baseURL == "" { + baseURL = "http://localhost:8080/v1" + } + + return &Client{ + baseURL: baseURL, + http: http.DefaultClient, + }, nil +} + +func (c *Client) Generate(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) { + endpoint := c.baseURL + "/chat/completions" + + payload, err := c.buildRequest(req, false) + if err != nil { + return llm.CompletionResponse{}, fmt.Errorf("building request: %w", err) + } + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, payload) + if err != nil { + return llm.CompletionResponse{}, fmt.Errorf("creating request: %w", err) + } + httpReq.Header.Set("Content-Type", "application/json") + + resp, err := c.http.Do(httpReq) + if err != nil { + return llm.CompletionResponse{}, fmt.Errorf("request failed: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + body, _ := io.ReadAll(resp.Body) + return llm.CompletionResponse{}, fmt.Errorf("API error %d: %s", resp.StatusCode, string(body)) + } + + var apiResp llamaChatResponse + if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil { + return llm.CompletionResponse{}, fmt.Errorf("decoding response: %w", err) + } + + return c.toResponse(apiResp), nil +} + +func (c *Client) Stream(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] { + return func(yield func(llm.StreamChunk, error) bool) { + endpoint := c.baseURL + "/chat/completions" + + payload, err := c.buildRequest(req, true) + if err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("building request: %w", err)) + return + } + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, payload) + if err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("creating request: %w", err)) + return + } + httpReq.Header.Set("Content-Type", "application/json") + + resp, err := c.http.Do(httpReq) + if err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("request failed: %w", err)) + return + } + defer resp.Body.Close() + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + body, _ := io.ReadAll(resp.Body) + yield(llm.StreamChunk{}, fmt.Errorf("API error %d: %s", resp.StatusCode, string(body))) + return + } + + scanner := bufio.NewScanner(resp.Body) + for scanner.Scan() { + line := scanner.Text() + if !strings.HasPrefix(line, "data: ") { + continue + } + data := strings.TrimPrefix(line, "data: ") + if data == "[DONE]" { + return + } + + var event llamaStreamEvent + if err := json.Unmarshal([]byte(data), &event); err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("decoding event: %w", err)) + return + } + + for _, choice := range event.Choices { + chunk := llm.StreamChunk{ + Delta: choice.Delta.Content, + } + if choice.FinishReason != "" { + chunk.FinishReason = choice.FinishReason + } + if !yield(chunk, nil) { + return + } + } + } + if err := scanner.Err(); err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("stream error: %w", err)) + } + } +} + +func (c *Client) Name() string { + return "llama.cpp" +} + +func (c *Client) Capabilities() llm.ProviderCapabilities { + return llm.ProviderCapabilities{ + SupportsTools: true, + SupportsVision: false, + SupportsJSON: true, + MaxContextWindow: 32768, + } +} + +// buildRequest converts an llm.CompletionRequest to the llama.cpp API format. +func (c *Client) buildRequest(req llm.CompletionRequest, stream bool) (io.Reader, error) { + messages := make([]llamaMessage, len(req.Messages)) + for i, m := range req.Messages { + messages[i] = llamaMessage{ + Role: string(m.Role), + Content: m.Content, + } + } + + var tools []llamaTool + for _, t := range req.Tools { + var tool llamaTool + if err := json.Unmarshal(t, &tool); err != nil { + return nil, fmt.Errorf("parsing tool: %w", err) + } + tools = append(tools, tool) + } + + openReq := llamaChatRequest{ + Model: req.Model, + Messages: messages, + Stream: stream, + } + if len(tools) > 0 { + openReq.Tools = tools + } + if req.ToolChoice != nil { + openReq.ToolChoice = req.ToolChoice + } + if req.Temperature != nil { + tmp := *req.Temperature + openReq.Temperature = tmp + } + if req.MaxTokens != nil { + tmp := *req.MaxTokens + openReq.MaxTokens = tmp + } + if len(req.Stop) > 0 { + openReq.Stop = req.Stop + } + + data, err := json.Marshal(openReq) + if err != nil { + return nil, fmt.Errorf("marshaling request: %w", err) + } + return strings.NewReader(string(data)), nil +} + +// toResponse converts a llama.cpp API response to our CompletionResponse. +func (c *Client) toResponse(resp llamaChatResponse) llm.CompletionResponse { + choice := resp.Choices[0] + result := llm.CompletionResponse{ + ID: resp.ID, + Model: resp.Model, + Content: choice.Message.Content, + StopReason: choice.FinishReason, + } + + for _, tc := range choice.Message.ToolCalls { + result.ToolCalls = append(result.ToolCalls, llm.ToolCall{ + ID: tc.ID, + Name: tc.Function.Name, + Arguments: json.RawMessage(tc.Function.Arguments), + }) + } + + result.Usage = llm.TokenUsage{ + InputTokens: resp.Usage.PromptTokens, + OutputTokens: resp.Usage.CompletionTokens, + TotalTokens: resp.Usage.TotalTokens, + } + + return result +} + +// llama.cpp API types + +type llamaChatRequest struct { + Model string `json:"model"` + Messages []llamaMessage `json:"messages"` + Tools []llamaTool `json:"tools,omitempty"` + ToolChoice interface{} `json:"tool_choice,omitempty"` + Temperature float32 `json:"temperature,omitempty"` + MaxTokens int `json:"max_tokens,omitempty"` + TopK int `json:"top_k,omitempty"` + TopP float32 `json:"top_p,omitempty"` + Stop []string `json:"stop,omitempty"` + Stream bool `json:"stream"` +} + +type llamaMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type llamaTool struct { + Type string `json:"type"` + Function json.RawMessage `json:"function"` +} + +type llamaChatResponse struct { + ID string `json:"id"` + Model string `json:"model"` + Choices []llamaChoice `json:"choices"` + Usage llamaUsage `json:"usage"` +} + +type llamaChoice struct { + Index int `json:"index"` + Message llamaMessageResult `json:"message"` + FinishReason string `json:"finish_reason"` +} + +type llamaMessageResult struct { + Role string `json:"role"` + Content string `json:"content"` + ToolCalls []llamaToolCall `json:"tool_calls"` +} + +type llamaToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + Function llamaFunction `json:"function"` +} + +type llamaFunction struct { + Name string `json:"name"` + Arguments string `json:"arguments"` +} + +type llamaUsage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` +} + +// Stream event types + +type llamaStreamEvent struct { + ID string `json:"id"` + Choices []llamaStreamChoice `json:"choices"` +} + +type llamaStreamChoice struct { + Index int `json:"index"` + Delta llamaStreamDelta `json:"delta"` + FinishReason string `json:"finish_reason"` +} + +type llamaStreamDelta struct { + Content string `json:"content"` + Role string `json:"role"` + ToolCalls []llamaStreamToolCall `json:"tool_calls"` +} + +type llamaStreamToolCall struct { + Index int `json:"index"` + ID string `json:"id"` + Type string `json:"type"` + Function llamaStreamFunction `json:"function"` +} + +type llamaStreamFunction struct { + Name string `json:"name"` + Arguments string `json:"arguments"` +} diff --git a/pkg/llm/providers/llamacpp/client_test.go b/pkg/llm/providers/llamacpp/client_test.go new file mode 100644 index 0000000..7b98377 --- /dev/null +++ b/pkg/llm/providers/llamacpp/client_test.go @@ -0,0 +1,232 @@ +package llamacpp + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/VictorVargas/rony-llm-agent/pkg/llm" +) + +func TestClient_Name(t *testing.T) { + client := &Client{} + if client.Name() != "llama.cpp" { + t.Errorf("expected 'llama.cpp', got %q", client.Name()) + } +} + +func TestClient_Capabilities(t *testing.T) { + client := &Client{} + caps := client.Capabilities() + if !caps.SupportsTools { + t.Error("expected SupportsTools to be true") + } + if caps.SupportsVision { + t.Error("expected SupportsVision to be false") + } +} + +func TestClient_Generate(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(llamaChatResponse{ + ID: "llama-123", + Model: "llama3", + Choices: []llamaChoice{ + { + Index: 0, + FinishReason: "stop", + Message: llamaMessageResult{ + Role: "assistant", + Content: "Hello from llama.cpp!", + }, + }, + }, + Usage: llamaUsage{ + PromptTokens: 10, + CompletionTokens: 5, + TotalTokens: 15, + }, + }) + })) + defer server.Close() + + client, err := New(Config{ + BaseURL: server.URL + "/v1", + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + resp, err := client.Generate(context.Background(), llm.CompletionRequest{ + Model: "llama3", + Messages: []llm.Message{ + {Role: llm.RoleUser, Content: "Hi"}, + }, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if resp.Content != "Hello from llama.cpp!" { + t.Errorf("expected 'Hello from llama.cpp!', got %q", resp.Content) + } + if resp.ID != "llama-123" { + t.Errorf("expected 'llama-123', got %q", resp.ID) + } + if resp.Usage.InputTokens != 10 { + t.Errorf("expected 10 input tokens, got %d", resp.Usage.InputTokens) + } +} + +func TestClient_Generate_Error(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + w.Write([]byte("server error")) + })) + defer server.Close() + + client, err := New(Config{ + BaseURL: server.URL + "/v1", + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + _, err = client.Generate(context.Background(), llm.CompletionRequest{}) + if err == nil { + t.Fatal("expected error, got nil") + } +} + +func TestClient_Generate_ToolCall(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(llamaChatResponse{ + ID: "llama-tool-1", + Model: "llama3", + Choices: []llamaChoice{ + { + Index: 0, + FinishReason: "tool_calls", + Message: llamaMessageResult{ + Role: "assistant", + Content: "", + ToolCalls: []llamaToolCall{ + { + ID: "call-1", + Type: "function", + Function: llamaFunction{ + Name: "calculate", + Arguments: `{"a":1,"b":2}`, + }, + }, + }, + }, + }, + }, + Usage: llamaUsage{ + PromptTokens: 20, + CompletionTokens: 10, + TotalTokens: 30, + }, + }) + })) + defer server.Close() + + client, err := New(Config{ + BaseURL: server.URL + "/v1", + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + resp, err := client.Generate(context.Background(), llm.CompletionRequest{ + Model: "llama3", + Messages: []llm.Message{ + {Role: llm.RoleUser, Content: "What is 1 + 2?"}, + }, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(resp.ToolCalls) != 1 { + t.Fatalf("expected 1 tool call, got %d", len(resp.ToolCalls)) + } + if resp.ToolCalls[0].Name != "calculate" { + t.Errorf("expected 'calculate', got %q", resp.ToolCalls[0].Name) + } +} + +func TestClient_Stream(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + w.Write([]byte("data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n")) + w.Write([]byte("data: {\"choices\":[{\"delta\":{\"content\":\" world\"},\"finish_reason\":null}]}\n")) + w.Write([]byte("data: {\"choices\":[{\"delta\":{\"content\":\"!\"},\"finish_reason\":\"stop\"}]}\n")) + w.Write([]byte("data: [DONE]\n")) + })) + defer server.Close() + + client, err := New(Config{ + BaseURL: server.URL + "/v1", + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + var chunks []llm.StreamChunk + stream := client.Stream(context.Background(), llm.CompletionRequest{}) + for chunk, err := range stream { + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + chunks = append(chunks, chunk) + } + + if len(chunks) != 3 { + t.Errorf("expected 3 chunks, got %d", len(chunks)) + } + if chunks[0].Delta != "Hello" { + t.Errorf("expected 'Hello', got %q", chunks[0].Delta) + } + if chunks[1].Delta != " world" { + t.Errorf("expected ' world', got %q", chunks[1].Delta) + } + if chunks[2].Delta != "!" { + t.Errorf("expected '!', got %q", chunks[2].Delta) + } +} + +func TestClient_Stream_FinishReason(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + w.Write([]byte("data: {\"choices\":[{\"delta\":{\"content\":\"done\"},\"finish_reason\":\"stop\"}]}\n")) + w.Write([]byte("data: [DONE]\n")) + })) + defer server.Close() + + client, err := New(Config{ + BaseURL: server.URL + "/v1", + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + var chunks []llm.StreamChunk + stream := client.Stream(context.Background(), llm.CompletionRequest{}) + for chunk, err := range stream { + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + chunks = append(chunks, chunk) + } + + if len(chunks) != 1 { + t.Fatalf("expected 1 chunk, got %d", len(chunks)) + } + if chunks[0].FinishReason != "stop" { + t.Errorf("expected 'stop' finish reason, got %q", chunks[0].FinishReason) + } +} diff --git a/pkg/llm/providers/openai/client.go b/pkg/llm/providers/openai/client.go new file mode 100644 index 0000000..d3967cf --- /dev/null +++ b/pkg/llm/providers/openai/client.go @@ -0,0 +1,327 @@ +package openai + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "iter" + "strings" + + "github.com/VictorVargas/rony-llm-agent/pkg/llm" +) + +// Config holds the settings needed to create an OpenAI client. +type Config struct { + APIKey string + Model string + BaseURL string // defaults to https://api.openai.com/v1 +} + +// Client implements llm.LLMClient for OpenAI. +type Client struct { + apiKey string + baseURL string + http *http.Client +} + +// New returns a new OpenAI client. +func New(cfg Config) (*Client, error) { + if cfg.APIKey == "" { + return nil, fmt.Errorf("openai: API key is required") + } + + baseURL := cfg.BaseURL + if baseURL == "" { + baseURL = "https://api.openai.com/v1" + } + + return &Client{ + apiKey: cfg.APIKey, + baseURL: baseURL, + http: http.DefaultClient, + }, nil +} + +func (c *Client) Generate(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) { + endpoint := c.baseURL + "/chat/completions" + + payload, err := c.buildRequest(req) + if err != nil { + return llm.CompletionResponse{}, fmt.Errorf("building request: %w", err) + } + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, payload) + if err != nil { + return llm.CompletionResponse{}, fmt.Errorf("creating request: %w", err) + } + httpReq.Header.Set("Authorization", "Bearer "+c.apiKey) + httpReq.Header.Set("Content-Type", "application/json") + + resp, err := c.http.Do(httpReq) + if err != nil { + return llm.CompletionResponse{}, fmt.Errorf("request failed: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + body, _ := io.ReadAll(resp.Body) + return llm.CompletionResponse{}, fmt.Errorf("API error %d: %s", resp.StatusCode, string(body)) + } + + var apiResp openaiChatResponse + if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil { + return llm.CompletionResponse{}, fmt.Errorf("decoding response: %w", err) + } + + return c.toResponse(apiResp), nil +} + +func (c *Client) Stream(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] { + return func(yield func(llm.StreamChunk, error) bool) { + endpoint := c.baseURL + "/chat/completions" + + payload, err := c.buildRequest(req) + if err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("building request: %w", err)) + return + } + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, payload) + if err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("creating request: %w", err)) + return + } + httpReq.Header.Set("Authorization", "Bearer "+c.apiKey) + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Accept", "text/event-stream") + + resp, err := c.http.Do(httpReq) + if err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("request failed: %w", err)) + return + } + defer resp.Body.Close() + + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + body, _ := io.ReadAll(resp.Body) + yield(llm.StreamChunk{}, fmt.Errorf("API error %d: %s", resp.StatusCode, string(body))) + return + } + + scanner := bufio.NewScanner(resp.Body) + for scanner.Scan() { + line := scanner.Text() + if !strings.HasPrefix(line, "data: ") { + continue + } + data := strings.TrimPrefix(line, "data: ") + if data == "[DONE]" { + return + } + + var event openaiStreamEvent + if err := json.Unmarshal([]byte(data), &event); err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("decoding event: %w", err)) + return + } + + for _, choice := range event.Choices { + chunk := llm.StreamChunk{ + Delta: choice.Delta.Content, + } + if choice.FinishReason != "" { + chunk.FinishReason = choice.FinishReason + } + if !yield(chunk, nil) { + return + } + } + } + if err := scanner.Err(); err != nil { + yield(llm.StreamChunk{}, fmt.Errorf("stream error: %w", err)) + } + } +} + +func (c *Client) Name() string { + return "openai" +} + +func (c *Client) Capabilities() llm.ProviderCapabilities { + return llm.ProviderCapabilities{ + SupportsTools: true, + SupportsVision: true, + SupportsJSON: true, + MaxContextWindow: 128000, + } +} + +// buildRequest converts an llm.CompletionRequest to the OpenAI API format. +func (c *Client) buildRequest(req llm.CompletionRequest) (io.Reader, error) { + // Convert messages to OpenAI format + messages := make([]openaiMessage, len(req.Messages)) + for i, m := range req.Messages { + messages[i] = openaiMessage{ + Role: string(m.Role), + Content: m.Content, + } + } + + tools := make([]openaiTool, len(req.Tools)) + for i, t := range req.Tools { + var tool openaiTool + if err := json.Unmarshal(t, &tool); err != nil { + return nil, fmt.Errorf("parsing tool %d: %w", i, err) + } + tools[i] = tool + } + + openaiReq := openaiChatRequest{ + Model: req.Model, + Messages: messages, + Stream: false, + } + if len(tools) > 0 { + openaiReq.Tools = tools + } + if req.ToolChoice != nil { + openaiReq.ToolChoice = req.ToolChoice + } + if req.Temperature != nil { + tmp := *req.Temperature + openaiReq.Temperature = &tmp + } + if req.MaxTokens != nil { + max := *req.MaxTokens + openaiReq.MaxTokens = &max + } + if len(req.Stop) > 0 { + openaiReq.Stop = req.Stop + } + + data, err := json.Marshal(openaiReq) + if err != nil { + return nil, fmt.Errorf("marshaling request: %w", err) + } + return strings.NewReader(string(data)), nil +} + +// toResponse converts an OpenAI API response to our CompletionResponse. +func (c *Client) toResponse(resp openaiChatResponse) llm.CompletionResponse { + choice := resp.Choices[0] + result := llm.CompletionResponse{ + ID: resp.ID, + Model: resp.Model, + Content: choice.Message.Content, + StopReason: choice.FinishReason, + } + + for _, tc := range choice.Message.ToolCalls { + result.ToolCalls = append(result.ToolCalls, llm.ToolCall{ + ID: tc.ID, + Name: tc.Function.Name, + Arguments: json.RawMessage(tc.Function.Arguments), + }) + } + + result.Usage = llm.TokenUsage{ + InputTokens: resp.Usage.PromptTokens, + OutputTokens: resp.Usage.CompletionTokens, + TotalTokens: resp.Usage.TotalTokens, + } + + return result +} + +// OpenAI API types + +type openaiChatRequest struct { + Model string `json:"model"` + Messages []openaiMessage `json:"messages"` + Tools []openaiTool `json:"tools,omitempty"` + ToolChoice interface{} `json:"tool_choice,omitempty"` + Temperature *float32 `json:"temperature,omitempty"` + MaxTokens *int `json:"max_tokens,omitempty"` + Stop []string `json:"stop,omitempty"` + Stream bool `json:"stream"` +} + +type openaiMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type openaiTool struct { + Type string `json:"type"` + Function json.RawMessage `json:"function"` +} + +type openaiChatResponse struct { + ID string `json:"id"` + Model string `json:"model"` + Choices []openaiChoice `json:"choices"` + Usage openaiUsage `json:"usage"` +} + +type openaiChoice struct { + Index int `json:"index"` + Message openaiMessageResult `json:"message"` + FinishReason string `json:"finish_reason"` +} + +type openaiMessageResult struct { + Role string `json:"role"` + Content string `json:"content"` + ToolCalls []openaiToolCall `json:"tool_calls"` +} + +type openaiToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + Function openaiFunction `json:"function"` +} + +type openaiFunction struct { + Name string `json:"name"` + Arguments string `json:"arguments"` +} + +type openaiUsage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` +} + +// Stream event types + +type openaiStreamEvent struct { + ID string `json:"id"` + Choices []openaiStreamChoice `json:"choices"` +} + +type openaiStreamChoice struct { + Index int `json:"index"` + Delta openaiStreamDelta `json:"delta"` + FinishReason string `json:"finish_reason"` +} + +type openaiStreamDelta struct { + Content string `json:"content"` + Role string `json:"role"` + ToolCalls []openaiStreamToolCall `json:"tool_calls"` +} + +type openaiStreamToolCall struct { + Index int `json:"index"` + ID string `json:"id"` + Type string `json:"type"` + Function openaiStreamFunction `json:"function"` +} + +type openaiStreamFunction struct { + Name string `json:"name"` + Arguments string `json:"arguments"` +} diff --git a/pkg/llm/providers/openai/client_test.go b/pkg/llm/providers/openai/client_test.go new file mode 100644 index 0000000..dd838f6 --- /dev/null +++ b/pkg/llm/providers/openai/client_test.go @@ -0,0 +1,141 @@ +package openai + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/VictorVargas/rony-llm-agent/pkg/llm" +) + +func TestClient_Name(t *testing.T) { + client := &Client{} + if client.Name() != "openai" { + t.Errorf("expected 'openai', got %q", client.Name()) + } +} + +func TestClient_Capabilities(t *testing.T) { + client := &Client{} + caps := client.Capabilities() + if !caps.SupportsTools { + t.Error("expected SupportsTools to be true") + } + if !caps.SupportsVision { + t.Error("expected SupportsVision to be true") + } +} + +func TestClient_Generate(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + json.NewEncoder(w).Encode(openaiChatResponse{ + ID: "test-123", + Model: "gpt-4", + Choices: []openaiChoice{ + { + Index: 0, + FinishReason: "stop", + Message: openaiMessageResult{ + Role: "assistant", + Content: "Hello!", + }, + }, + }, + Usage: openaiUsage{ + PromptTokens: 10, + CompletionTokens: 5, + TotalTokens: 15, + }, + }) + })) + defer server.Close() + + client, err := New(Config{ + APIKey: "test-key", + BaseURL: server.URL, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + resp, err := client.Generate(context.Background(), llm.CompletionRequest{ + Model: "gpt-4", + Messages: []llm.Message{ + {Role: llm.RoleUser, Content: "Hi"}, + }, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if resp.Content != "Hello!" { + t.Errorf("expected 'Hello!', got %q", resp.Content) + } + if resp.ID != "test-123" { + t.Errorf("expected 'test-123', got %q", resp.ID) + } + if resp.Usage.InputTokens != 10 { + t.Errorf("expected 10 input tokens, got %d", resp.Usage.InputTokens) + } +} + +func TestClient_Generate_Error(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + w.Write([]byte("invalid api key")) + })) + defer server.Close() + + client, err := New(Config{ + APIKey: "bad-key", + BaseURL: server.URL, + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + _, err = client.Generate(context.Background(), llm.CompletionRequest{}) + if err == nil { + t.Fatal("expected error, got nil") + } +} + +func TestClient_Stream(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + w.Write([]byte("data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n")) + w.Write([]byte("data: {\"choices\":[{\"delta\":{\"content\":\" world\"},\"finish_reason\":null}]}\n")) + w.Write([]byte("data: [DONE]\n")) + })) + defer server.Close() + + client, err := New(Config{ + APIKey: "test-key", + BaseURL: server.URL + "/v1", + }) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + var chunks []llm.StreamChunk + stream := client.Stream(context.Background(), llm.CompletionRequest{}) + for chunk, err := range stream { + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + chunks = append(chunks, chunk) + } + + if len(chunks) != 2 { + t.Errorf("expected 2 chunks, got %d", len(chunks)) + } + if chunks[0].Delta != "Hello" { + t.Errorf("expected 'Hello', got %q", chunks[0].Delta) + } + if chunks[1].Delta != " world" { + t.Errorf("expected ' world', got %q", chunks[1].Delta) + } +} diff --git a/pkg/llm/types.go b/pkg/llm/types.go new file mode 100644 index 0000000..f7d9fa8 --- /dev/null +++ b/pkg/llm/types.go @@ -0,0 +1,120 @@ +package llm + +import ( + "context" + "encoding/json" + "iter" +) + +// Role represents the role of a message participant. +type Role string + +const ( + RoleSystem Role = "system" + RoleUser Role = "user" + RoleAssistant Role = "assistant" + RoleTool Role = "tool" +) + +// Message is a single message in a conversation. +type Message struct { + Role Role `json:"role"` + Content string `json:"content"` + ToolCallID string `json:"tool_call_id,omitempty"` + Name string `json:"name,omitempty"` +} + +// ProviderCapabilities describes what a model supports. +type ProviderCapabilities struct { + SupportsTools bool `json:"supports_tools"` + SupportsVision bool `json:"supports_vision"` + SupportsJSON bool `json:"supports_json"` + MaxContextWindow int `json:"max_context_window"` +} + +// TokenUsage tracks token consumption for a response. +type TokenUsage struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + TotalTokens int `json:"total_tokens"` +} + +// ToolChoice controls how the model chooses which tools to call. +type ToolChoice string + +const ( + ToolChoiceAuto ToolChoice = "auto" + ToolChoiceNone ToolChoice = "none" + ToolChoiceRequired ToolChoice = "required" +) + +// ToolRef is used when ToolChoice must invoke a specific tool by name. +type ToolRef struct { + Name string `json:"name"` +} + +func (t *ToolRef) MarshalJSON() ([]byte, error) { + type Alias ToolRef + return json.Marshal(&struct { + *Alias + Type string `json:"type"` + }{ + Alias: (*Alias)(t), + Type: "function", + }) +} + +// CompletionRequest is sent to an LLM provider. +type CompletionRequest struct { + Model string `json:"model"` + Messages []Message `json:"messages"` + Tools []json.RawMessage `json:"tools,omitempty"` + ToolChoice interface{} `json:"tool_choice,omitempty"` // ToolChoice, ToolRef, or null + Temperature *float32 `json:"temperature,omitempty"` + MaxTokens *int `json:"max_tokens,omitempty"` + Stop []string `json:"stop,omitempty"` + Metadata map[string]string `json:"metadata,omitempty"` +} + +// CompletionResponse is returned from an LLM provider. +type CompletionResponse struct { + ID string `json:"id"` + Model string `json:"model"` + Content string `json:"content"` + ToolCalls []ToolCall `json:"tool_calls,omitempty"` + StopReason string `json:"stop_reason"` + Usage TokenUsage `json:"usage,omitempty"` +} + +// StopReason values. +const ( + StopReasonEndTurn = "end_turn" + StopReasonToolUse = "tool_use" + StopReasonMaxTokens = "max_tokens" + StopReasonStopSeq = "stop_sequence" +) + +// ToolCall represents a function invocation requested by the model. +type ToolCall struct { + ID string `json:"id"` + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` + Thought string `json:"thought,omitempty"` +} + +// StreamChunk is emitted by the iterator returned from Stream(). +type StreamChunk struct { + Delta string `json:"delta"` + ToolCalls []ToolCall `json:"tool_calls,omitempty"` + FinishReason string `json:"finish_reason,omitempty"` + Usage TokenUsage `json:"usage,omitempty"` // only present on final chunk +} + +// LLMClient is the interface that every provider implements. +// Products must never call a provider SDK directly; they go through this interface. +type LLMClient interface { + Generate(ctx context.Context, req CompletionRequest) (CompletionResponse, error) + Stream(ctx context.Context, req CompletionRequest) iter.Seq2[StreamChunk, error] + Name() string + Capabilities() ProviderCapabilities +}