327 lines
8.4 KiB
Go
327 lines
8.4 KiB
Go
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"`
|
|
}
|