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, ToolCallID: m.ToolCallID, Name: m.Name, } if len(m.ToolCalls) > 0 { calls := make([]openaiToolCall, len(m.ToolCalls)) for j, tc := range m.ToolCalls { calls[j] = openaiToolCall{ ID: tc.ID, Type: "function", Function: openaiFunction{ Name: tc.Name, Arguments: string(tc.Arguments), }, } } messages[i].ToolCalls = calls } } 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"` ToolCallID string `json:"tool_call_id,omitempty"` Name string `json:"name,omitempty"` ToolCalls []openaiToolCall `json:"tool_calls,omitempty"` } 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"` }