package llm import ( "bytes" "context" "encoding/json" "errors" "fmt" "io" "net/http" "net/url" "strings" "time" "gitea.maximumdirect.net/eric/scriptorium/internal/domain" ) var ( ErrInvalidConfig = errors.New("invalid llm client configuration") ErrInvalidRequest = errors.New("invalid generate request") ErrRequestFailed = errors.New("llm request failed") ErrUnexpectedStatus = errors.New("llm returned non-success status") ErrMalformedResponse = errors.New("malformed llm response") ) type OpenAICompatibleConfig struct { BaseURL string APIKey string Model string Timeout time.Duration HTTPClient *http.Client } type OpenAICompatibleClient struct { baseURL string apiKey string defaultModel string timeout time.Duration httpClient *http.Client } func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) { if strings.TrimSpace(cfg.BaseURL) == "" { return nil, fmt.Errorf("%w: base URL is required", ErrInvalidConfig) } if _, err := url.ParseRequestURI(cfg.BaseURL); err != nil { return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err) } timeout := cfg.Timeout if timeout <= 0 { timeout = 10 * time.Minute } var client *http.Client if cfg.HTTPClient != nil { client = cfg.HTTPClient if client.Timeout == 0 { client.Timeout = timeout } } else { client = &http.Client{Timeout: timeout} } return &OpenAICompatibleClient{ baseURL: strings.TrimRight(cfg.BaseURL, "/"), apiKey: cfg.APIKey, defaultModel: cfg.Model, timeout: timeout, httpClient: client, }, nil } func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) { if req.Target.TimeoutSeconds < 0 { return nil, fmt.Errorf("%w: timeout_seconds must be greater than or equal to 0", ErrInvalidRequest) } model := strings.TrimSpace(req.Target.Model) if model == "" { model = strings.TrimSpace(c.defaultModel) } if model == "" { return nil, fmt.Errorf("%w: model is required", ErrInvalidRequest) } endpoint := strings.TrimSpace(req.Target.Endpoint) if endpoint == "" { endpoint = c.baseURL } endpoint = strings.TrimRight(endpoint, "/") + "/chat/completions" wireReq := openAIChatRequest{ Model: model, } wireReq.Messages = make([]openAIChatMessage, 0, len(req.Prompt.Messages)) for _, msg := range req.Prompt.Messages { wireReq.Messages = append(wireReq.Messages, openAIChatMessage{ Role: msg.Role, Content: msg.Content, }) } if req.Target.Temperature != 0 { wireReq.Temperature = &req.Target.Temperature } if req.Target.MaxTokens != 0 { wireReq.MaxTokens = &req.Target.MaxTokens } if req.Target.TopP != 0 { wireReq.TopP = &req.Target.TopP } payload, err := json.Marshal(wireReq) if err != nil { return nil, fmt.Errorf("%w: failed to encode request: %v", ErrRequestFailed, err) } httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(payload)) if err != nil { return nil, fmt.Errorf("%w: failed to create request: %v", ErrRequestFailed, err) } httpReq.Header.Set("Content-Type", "application/json") if strings.TrimSpace(c.apiKey) != "" { httpReq.Header.Set("Authorization", "Bearer "+c.apiKey) } effectiveTimeout := c.timeout if req.Target.TimeoutSeconds > 0 { effectiveTimeout = time.Duration(req.Target.TimeoutSeconds) * time.Second } httpClient := c.httpClient if httpClient == nil { httpClient = &http.Client{Timeout: effectiveTimeout} } else if httpClient.Timeout != effectiveTimeout { cloned := *httpClient cloned.Timeout = effectiveTimeout httpClient = &cloned } httpResp, err := httpClient.Do(httpReq) if err != nil { return nil, fmt.Errorf("%w: %v", ErrRequestFailed, err) } defer httpResp.Body.Close() if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 { body, _ := io.ReadAll(io.LimitReader(httpResp.Body, 4096)) return nil, fmt.Errorf("%w: status=%d body=%q", ErrUnexpectedStatus, httpResp.StatusCode, strings.TrimSpace(string(body))) } var wireResp openAIChatResponse if err := json.NewDecoder(httpResp.Body).Decode(&wireResp); err != nil { return nil, fmt.Errorf("%w: failed to decode response: %v", ErrMalformedResponse, err) } if len(wireResp.Choices) == 0 { return nil, fmt.Errorf("%w: no choices returned", ErrMalformedResponse) } content := wireResp.Choices[0].Message.Content if content == "" { return nil, fmt.Errorf("%w: first choice has empty message content", ErrMalformedResponse) } return &domain.GenerateResponse{ Content: content, Usage: domain.TokenUsage{ PromptTokens: wireResp.Usage.PromptTokens, CompletionTokens: wireResp.Usage.CompletionTokens, TotalTokens: wireResp.Usage.TotalTokens, }, }, nil } type openAIChatRequest struct { Model string `json:"model"` Messages []openAIChatMessage `json:"messages"` Temperature *float64 `json:"temperature,omitempty"` MaxTokens *int `json:"max_tokens,omitempty"` TopP *float64 `json:"top_p,omitempty"` } type openAIChatMessage struct { Role string `json:"role"` Content string `json:"content"` } type openAIChatResponse struct { Choices []struct { Message openAIChatMessage `json:"message"` } `json:"choices"` Usage struct { PromptTokens int `json:"prompt_tokens"` CompletionTokens int `json:"completion_tokens"` TotalTokens int `json:"total_tokens"` } `json:"usage"` }