package llm import ( "bytes" "context" "encoding/json" "errors" "fmt" "io" "net/http" "net/url" "os" "strings" "time" "gitea.maximumdirect.net/eric/promptkit/internal/defaults" "gitea.maximumdirect.net/eric/promptkit/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 requestFailedError struct { cause error } func (e *requestFailedError) Error() string { return ErrRequestFailed.Error() } func (e *requestFailedError) Unwrap() []error { return []error{ErrRequestFailed, e.cause} } type OpenAICompatibleConfig struct { BaseURL string Model string Timeout time.Duration HTTPClient *http.Client } type OpenAICompatibleClient struct { baseURL string defaultModel string httpClient *http.Client } func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleClient, error) { baseURL := "" if strings.TrimSpace(cfg.BaseURL) != "" { var err error baseURL, err = domain.NormalizeOpenAICompatibleBaseEndpoint(cfg.BaseURL) if err != nil { return nil, fmt.Errorf("%w: invalid base URL: %v", ErrInvalidConfig, err) } } timeout := cfg.Timeout if timeout <= 0 { timeout = defaults.LLMRequestTimeoutDefault } var client *http.Client if cfg.HTTPClient != nil { cloned := *cfg.HTTPClient if cloned.Timeout <= 0 { cloned.Timeout = timeout } client = &cloned } else { client = &http.Client{Timeout: timeout} } return &OpenAICompatibleClient{ baseURL: baseURL, defaultModel: cfg.Model, httpClient: client, }, nil } func (c *OpenAICompatibleClient) Generate(ctx context.Context, req domain.GenerateRequest) (*domain.GenerateResponse, error) { if err := domain.ValidateExecutionTargetSettings(req.Target); err != nil { return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err) } selectedEndpoint := req.Target.Endpoint if strings.TrimSpace(selectedEndpoint) == "" { selectedEndpoint = c.baseURL } endpoint, err := domain.NormalizeOpenAICompatibleBaseEndpoint(selectedEndpoint) if err != nil { return nil, fmt.Errorf("%w: invalid endpoint: %v", ErrInvalidRequest, err) } endpoint, err = url.JoinPath(endpoint, defaults.OpenAIChatCompletionsPath) if err != nil { return nil, fmt.Errorf("%w: invalid endpoint path: %v", ErrInvalidRequest, err) } wireReq, err := openAIChatRequestFromGenerateRequest(req, c.defaultModel) if err != nil { return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err) } wirePayload, err := openAIChatRequestPayload(wireReq) if err != nil { return nil, fmt.Errorf("%w: %v", ErrInvalidRequest, err) } payload, err := json.Marshal(wirePayload) if err != nil { return nil, fmt.Errorf("%w: failed to encode request: %v", ErrRequestFailed, err) } requestContext := ctx if req.Target.TimeoutSeconds > 0 { var cancel context.CancelFunc requestContext, cancel = context.WithTimeout( ctx, time.Duration(req.Target.TimeoutSeconds)*time.Second, ) defer cancel() } httpReq, err := http.NewRequestWithContext(requestContext, 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 apiKey := strings.TrimSpace(req.Target.APIKey); apiKey != "" { httpReq.Header.Set("Authorization", "Bearer "+apiKey) } else if envName := strings.TrimSpace(req.Target.APIKeyEnv); envName != "" { apiKey := strings.TrimSpace(os.Getenv(envName)) if apiKey == "" { return nil, fmt.Errorf("%w: api key environment variable %q is not set", ErrInvalidRequest, envName) } httpReq.Header.Set("Authorization", "Bearer "+apiKey) } httpClient := c.httpClient if httpClient == nil { httpClient = &http.Client{Timeout: defaults.LLMRequestTimeoutDefault} } httpResp, err := httpClient.Do(httpReq) if err != nil { return nil, &requestFailedError{cause: err} } defer httpResp.Body.Close() if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 { _, _ = io.Copy(io.Discard, io.LimitReader(httpResp.Body, 4096)) return nil, fmt.Errorf("%w: status=%d", ErrUnexpectedStatus, httpResp.StatusCode) } 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, CachedTokens: wireResp.Usage.PromptTokensDetails.CachedTokens, CacheWriteTokens: wireResp.Usage.CacheWriteTokens, }, }, nil } func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultModel string) (openAIChatRequest, error) { model := strings.TrimSpace(req.Target.Model) if model == "" { model = strings.TrimSpace(defaultModel) } if model == "" { return openAIChatRequest{}, errors.New("model is required") } wireReq := openAIChatRequest{ Model: model, } sessionID, err := domain.NormalizeSessionID(req.Prompt.SessionID) if err != nil { return openAIChatRequest{}, err } wireReq.SessionID = sessionID wireReq.Messages = make([]openAIChatRequestMessage, 0, len(req.Prompt.Messages)) for _, msg := range req.Prompt.Messages { wireReq.Messages = append(wireReq.Messages, openAIChatRequestMessageFromRenderedMessage(msg)) } if req.Target.Temperature != 0 || req.TargetPresence.Temperature { wireReq.Temperature = &req.Target.Temperature } if req.Target.MaxTokens != 0 || req.TargetPresence.MaxTokens { wireReq.MaxTokens = &req.Target.MaxTokens } if req.Target.TopP != 0 || req.TargetPresence.TopP { wireReq.TopP = &req.Target.TopP } if strings.TrimSpace(req.Target.ServiceTier) != "" { wireReq.ServiceTier = req.Target.ServiceTier } if strings.TrimSpace(req.Target.ReasoningEffort) != "" { wireReq.ReasoningEffort = req.Target.ReasoningEffort } if len(req.Target.ExtraParams) > 0 { wireReq.ExtraParams = req.Target.ExtraParams } if req.StructuredOutput != nil { responseFormat, err := toOpenAIResponseFormat(req.StructuredOutput) if err != nil { return openAIChatRequest{}, err } wireReq.ResponseFormat = responseFormat } return wireReq, nil } type openAIChatRequest struct { Model string `json:"model"` SessionID string `json:"session_id,omitempty"` Messages []openAIChatRequestMessage `json:"messages"` Temperature *float64 `json:"temperature,omitempty"` MaxTokens *int `json:"max_tokens,omitempty"` TopP *float64 `json:"top_p,omitempty"` ServiceTier string `json:"service_tier,omitempty"` ReasoningEffort string `json:"reasoning_effort,omitempty"` ResponseFormat *openAIResponseFormat `json:"response_format,omitempty"` ExtraParams map[string]any `json:"-"` } func openAIChatRequestPayload(req openAIChatRequest) (map[string]any, error) { out := map[string]any{ "model": req.Model, "messages": req.Messages, } if req.SessionID != "" { out["session_id"] = req.SessionID } if req.Temperature != nil { out["temperature"] = *req.Temperature } if req.MaxTokens != nil { out["max_tokens"] = *req.MaxTokens } if req.TopP != nil { out["top_p"] = *req.TopP } if req.ServiceTier != "" { out["service_tier"] = req.ServiceTier } if req.ReasoningEffort != "" { out["reasoning_effort"] = req.ReasoningEffort } if req.ResponseFormat != nil { out["response_format"] = req.ResponseFormat } for key, value := range req.ExtraParams { if key == "" { return nil, errors.New("extra_params key must not be empty") } if IsReservedOpenAIChatRequestField(key) { return nil, fmt.Errorf("extra_params key %q collides with reserved request field", key) } if _, err := json.Marshal(value); err != nil { return nil, fmt.Errorf("extra_params.%s must be JSON-serializable: %w", key, err) } out[key] = value } return out, nil } // IsReservedOpenAIChatRequestField reports whether name is owned by the // standard OpenAI-compatible chat request rather than extra parameters. func IsReservedOpenAIChatRequestField(name string) bool { switch name { case "model", "session_id", "messages", "temperature", "max_tokens", "top_p", "service_tier", "reasoning_effort", "response_format": return true default: return false } } type openAIChatRequestMessage struct { Role string `json:"role"` Content any `json:"content"` } type openAIChatTextContentBlock struct { Type string `json:"type"` Text string `json:"text"` CacheControl *openAICacheControl `json:"cache_control,omitempty"` } type openAICacheControl struct { Type string `json:"type"` TTL string `json:"ttl,omitempty"` } type openAIChatResponseMessage struct { Role string `json:"role"` Content string `json:"content"` } type openAIChatResponse struct { Choices []struct { Message openAIChatResponseMessage `json:"message"` } `json:"choices"` Usage struct { PromptTokens int `json:"prompt_tokens"` CompletionTokens int `json:"completion_tokens"` TotalTokens int `json:"total_tokens"` PromptTokensDetails struct { CachedTokens int `json:"cached_tokens"` } `json:"prompt_tokens_details"` CacheWriteTokens int `json:"cache_write_tokens"` } `json:"usage"` } type openAIResponseFormat struct { Type string `json:"type"` JSONSchema *openAIJSONSchemaEnvelope `json:"json_schema,omitempty"` } type openAIJSONSchemaEnvelope struct { Name string `json:"name"` Strict bool `json:"strict"` Schema any `json:"schema"` } func openAIChatRequestMessageFromRenderedMessage(msg domain.RenderedMessage) openAIChatRequestMessage { wireMsg := openAIChatRequestMessage{ Role: msg.Role, Content: msg.Content, } if msg.CacheControl == nil { return wireMsg } wireMsg.Content = []openAIChatTextContentBlock{ { Type: "text", Text: msg.Content, CacheControl: &openAICacheControl{ Type: string(msg.CacheControl.Type), TTL: msg.CacheControl.TTL, }, }, } return wireMsg } func toOpenAIResponseFormat(spec *domain.StructuredOutputSpec) (*openAIResponseFormat, error) { if spec == nil { return nil, nil } switch spec.Type { case domain.StructuredOutputJSONSchema: if spec.JSONSchema == nil { return nil, errors.New("json_schema structured output requires schema payload") } if strings.TrimSpace(spec.JSONSchema.Name) == "" { return nil, errors.New("json_schema structured output requires non-empty schema name") } if spec.JSONSchema.Schema == nil { return nil, errors.New("json_schema structured output requires schema document") } return &openAIResponseFormat{ Type: "json_schema", JSONSchema: &openAIJSONSchemaEnvelope{ Name: spec.JSONSchema.Name, Strict: spec.JSONSchema.Strict, Schema: spec.JSONSchema.Schema, }, }, nil default: return nil, fmt.Errorf("unsupported structured output type %q", spec.Type) } }