Files
chorus/internal/core/provider/openai.go
T

421 lines
13 KiB
Go

package provider
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"mime/multipart"
"net/http"
"net/url"
"path"
"strings"
"git.ilapage.cn/OPC/chorus/internal/core/model"
)
const defaultMaxResponseBytes int64 = 20 << 20
var (
ErrInvalidConfig = errors.New("provider configuration is invalid")
ErrInvalidRequest = errors.New("provider request is invalid")
)
// OpenAIConfig applies to both OpenAI-compatible and Gemini protocol shapes.
// The name is retained for MVP-0 callers; APIType controls the fixed endpoint.
type OpenAIConfig struct {
BaseURL string
AuthType AuthType
APIKey string
ExtraBody json.RawMessage
MaxResponseBytes int64
}
type OpenAI struct {
http HTTPClient
baseURL *url.URL
authType AuthType
apiKey string
extraBody map[string]any
maxResponseBytes int64
}
func NewOpenAI(httpClient HTTPClient, config OpenAIConfig) (*OpenAI, error) {
if httpClient == nil || strings.TrimSpace(config.BaseURL) == "" {
return nil, ErrInvalidConfig
}
baseURL, err := url.Parse(config.BaseURL)
if err != nil || (baseURL.Scheme != "http" && baseURL.Scheme != "https") || baseURL.Hostname() == "" || baseURL.User != nil || baseURL.RawQuery != "" || baseURL.Fragment != "" {
return nil, ErrInvalidConfig
}
authType := config.AuthType
if authType == "" {
if config.APIKey == "" {
authType = AuthNone
} else {
authType = AuthBearer
}
}
if !authType.Valid() || (authType != AuthNone && strings.TrimSpace(config.APIKey) == "") {
return nil, ErrInvalidConfig
}
extraBody, err := parseExtraBody(config.ExtraBody)
if err != nil {
return nil, err
}
if config.MaxResponseBytes <= 0 {
config.MaxResponseBytes = defaultMaxResponseBytes
}
return &OpenAI{http: httpClient, baseURL: baseURL, authType: authType, apiKey: config.APIKey, extraBody: extraBody, maxResponseBytes: config.MaxResponseBytes}, nil
}
func (c *OpenAI) Generate(ctx context.Context, request Request) ([]Output, error) {
if strings.TrimSpace(request.ModelID) == "" || strings.TrimSpace(request.RenderedPrompt) == "" {
return nil, ErrInvalidRequest
}
switch request.APIType {
case model.APIChat:
if request.Kind != model.KindText || len(request.Inputs) != 0 {
return nil, ErrInvalidRequest
}
return c.chat(ctx, request)
case model.APIImages:
if request.Kind != model.KindImage || len(request.Inputs) != 0 {
return nil, ErrInvalidRequest
}
return c.images(ctx, request)
case model.APIImagesEdits:
if request.Kind != model.KindImage || !validImageInputs(request.Inputs, true) {
return nil, ErrInvalidRequest
}
return c.imagesEdits(ctx, request)
case model.APIGemini:
if request.Kind != model.KindText && request.Kind != model.KindImage {
return nil, ErrInvalidRequest
}
if len(request.Inputs) > 0 && !validImageInputs(request.Inputs, true) {
return nil, ErrInvalidRequest
}
return c.gemini(ctx, request)
default:
return nil, ErrInvalidRequest
}
}
func (c *OpenAI) chat(ctx context.Context, request Request) ([]Output, error) {
body := map[string]any{
"model": request.ModelID,
"messages": []map[string]string{{"role": "user", "content": request.RenderedPrompt}},
}
mergeExtra(body, c.extraBody)
encoded, err := json.Marshal(body)
if err != nil {
return nil, ErrInvalidRequest
}
response, err := c.do(ctx, c.openAIURL("/chat/completions"), "application/json", encoded)
if err != nil {
return nil, err
}
var decoded struct {
Choices []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
} `json:"choices"`
}
if json.Unmarshal(response.Body, &decoded) != nil || len(decoded.Choices) == 0 || decoded.Choices[0].Message.Content == "" {
return nil, providerError(CodeUnknown, FailureOther)
}
return []Output{{Kind: request.Kind, Text: decoded.Choices[0].Message.Content, ContentType: "text/plain; charset=utf-8"}}, nil
}
func (c *OpenAI) images(ctx context.Context, request Request) ([]Output, error) {
body := map[string]any{
"model": request.ModelID,
"prompt": request.RenderedPrompt,
"n": 1,
"response_format": "b64_json",
}
mergeExtra(body, c.extraBody)
encoded, err := json.Marshal(body)
if err != nil {
return nil, ErrInvalidRequest
}
response, err := c.do(ctx, c.openAIURL("/images/generations"), "application/json", encoded)
if err != nil {
return nil, err
}
return c.decodeOpenAIImages(ctx, request.Kind, response, true)
}
func (c *OpenAI) imagesEdits(ctx context.Context, request Request) ([]Output, error) {
var body bytes.Buffer
writer := multipart.NewWriter(&body)
if err := writer.WriteField("model", request.ModelID); err != nil {
return nil, ErrInvalidRequest
}
if err := writer.WriteField("prompt", request.RenderedPrompt); err != nil {
return nil, ErrInvalidRequest
}
for key, value := range c.extraBody {
encoded, err := scalarString(value)
if err != nil || writer.WriteField(key, encoded) != nil {
return nil, ErrInvalidRequest
}
}
for index, input := range request.Inputs {
part, err := writer.CreateFormFile("image[]", fmt.Sprintf("image-%d%s", index+1, imageExtension(input.MIMEType)))
if err != nil {
return nil, ErrInvalidRequest
}
if _, err := part.Write(input.Content); err != nil {
return nil, ErrInvalidRequest
}
}
if err := writer.Close(); err != nil {
return nil, ErrInvalidRequest
}
response, err := c.do(ctx, c.openAIURL("/images/edits"), writer.FormDataContentType(), body.Bytes())
if err != nil {
return nil, err
}
return c.decodeOpenAIImages(ctx, request.Kind, response, false)
}
func (c *OpenAI) gemini(ctx context.Context, request Request) ([]Output, error) {
parts := make([]map[string]any, 0, len(request.Inputs)+1)
parts = append(parts, map[string]any{"text": request.RenderedPrompt})
for _, input := range request.Inputs {
parts = append(parts, map[string]any{"inlineData": map[string]string{
"mimeType": input.MIMEType,
"data": base64.StdEncoding.EncodeToString(input.Content),
}})
}
body := map[string]any{"contents": []any{map[string]any{"role": "user", "parts": parts}}}
mergeExtra(body, c.extraBody)
encoded, err := json.Marshal(body)
if err != nil {
return nil, ErrInvalidRequest
}
response, err := c.do(ctx, c.geminiURL(request.ModelID), "application/json", encoded)
if err != nil {
return nil, err
}
var decoded struct {
Candidates []struct {
Content struct {
Parts []struct {
Text string `json:"text"`
InlineData *struct {
MIMEType string `json:"mimeType"`
Data string `json:"data"`
} `json:"inlineData"`
} `json:"parts"`
} `json:"content"`
} `json:"candidates"`
}
if json.Unmarshal(response.Body, &decoded) != nil || len(decoded.Candidates) == 0 {
return nil, providerError(CodeUnknown, FailureOther)
}
for _, candidate := range decoded.Candidates {
for _, part := range candidate.Content.Parts {
if request.Kind == model.KindText && part.Text != "" {
return []Output{{Kind: request.Kind, Text: part.Text, ContentType: "text/plain; charset=utf-8"}}, nil
}
if request.Kind == model.KindImage && part.InlineData != nil {
if !strings.HasPrefix(strings.ToLower(part.InlineData.MIMEType), "image/") {
return nil, providerError(CodeUnknown, FailureOther)
}
content, decodeErr := base64.StdEncoding.DecodeString(part.InlineData.Data)
if decodeErr != nil || int64(len(content)) > c.maxResponseBytes {
return nil, providerError(CodeUnknown, FailureOther)
}
return []Output{{Kind: request.Kind, Content: content, ContentType: part.InlineData.MIMEType}}, nil
}
}
}
return nil, providerError(CodeUnknown, FailureOther)
}
func (c *OpenAI) decodeOpenAIImages(ctx context.Context, kind model.GenerationKind, response HTTPResponse, requireOne bool) ([]Output, error) {
var decoded struct {
Data []struct {
B64JSON string `json:"b64_json"`
URL string `json:"url"`
} `json:"data"`
}
if json.Unmarshal(response.Body, &decoded) != nil || len(decoded.Data) == 0 || (requireOne && len(decoded.Data) != 1) {
return nil, providerError(CodeUnknown, FailureOther)
}
outputs := make([]Output, 0, len(decoded.Data))
for _, item := range decoded.Data {
var content []byte
contentType := "image/png"
var err error
switch {
case item.B64JSON != "":
content, err = base64.StdEncoding.DecodeString(item.B64JSON)
if err != nil || int64(len(content)) > c.maxResponseBytes {
return nil, providerError(CodeUnknown, FailureOther)
}
case item.URL != "":
fetched, fetchErr := c.http.Fetch(ctx, item.URL, c.maxResponseBytes)
if fetchErr != nil {
return nil, classifyNetworkError(fetchErr)
}
if fetched.StatusCode < 200 || fetched.StatusCode >= 300 {
return nil, fromHTTPStatus(fetched.StatusCode, nil)
}
content, contentType = fetched.Body, strings.TrimSpace(strings.Split(fetched.ContentType, ";")[0])
if !strings.HasPrefix(strings.ToLower(contentType), "image/") || int64(len(content)) > c.maxResponseBytes {
return nil, providerError(CodeUnknown, FailureOther)
}
default:
return nil, providerError(CodeUnknown, FailureOther)
}
outputs = append(outputs, Output{Kind: kind, Content: content, ContentType: contentType})
}
return outputs, nil
}
func (c *OpenAI) do(ctx context.Context, target string, contentType string, body []byte) (HTTPResponse, error) {
header := http.Header{"Content-Type": []string{contentType}, "Accept": []string{"application/json"}}
switch c.authType {
case AuthBearer:
header.Set("Authorization", "Bearer "+c.apiKey)
case AuthGoogleAPIKey:
header.Set("x-goog-api-key", c.apiKey)
}
response, err := c.http.Do(ctx, HTTPRequest{Method: http.MethodPost, URL: target, Header: header, Body: body, MaxBytes: c.maxResponseBytes})
if err != nil {
return HTTPResponse{}, classifyNetworkError(err)
}
if response.StatusCode < 200 || response.StatusCode >= 300 {
return HTTPResponse{}, fromHTTPStatus(response.StatusCode, response.Body)
}
return response, nil
}
func (c *OpenAI) openAIURL(endpoint string) string {
target := *c.baseURL
target.Path = path.Join(strings.TrimSuffix(c.baseURL.Path, "/"), endpoint)
target.RawPath = ""
return target.String()
}
func (c *OpenAI) geminiURL(modelID string) string {
target := *c.baseURL
target.Path = path.Join(strings.TrimSuffix(c.baseURL.Path, "/"), "models", modelID+":generateContent")
target.RawPath = path.Join(strings.TrimSuffix(c.baseURL.EscapedPath(), "/"), "models", url.PathEscape(modelID)+":generateContent")
return target.String()
}
func validImageInputs(inputs []Input, requirePrimary bool) bool {
if len(inputs) == 0 {
return !requirePrimary
}
primary := 0
positions := make(map[uint32]struct{}, len(inputs))
for _, input := range inputs {
if len(input.Content) == 0 || !strings.HasPrefix(strings.ToLower(input.MIMEType), "image/") {
return false
}
if _, exists := positions[input.Position]; exists {
return false
}
positions[input.Position] = struct{}{}
if input.Role == model.RolePrimary {
primary++
} else if input.Role != model.RoleReference {
return false
}
}
return !requirePrimary || primary == 1
}
func parseExtraBody(raw json.RawMessage) (map[string]any, error) {
result := map[string]any{}
if len(raw) == 0 || string(raw) == "null" {
return result, nil
}
if err := json.Unmarshal(raw, &result); err != nil {
return nil, ErrInvalidConfig
}
allowed := map[string]bool{"temperature": true, "max_tokens": true, "size": true, "quality": true}
for key, value := range result {
if !allowed[key] {
return nil, ErrInvalidConfig
}
switch value.(type) {
case string, float64, bool:
default:
return nil, ErrInvalidConfig
}
}
return result, nil
}
func mergeExtra(target, extra map[string]any) {
for key, value := range extra {
target[key] = value
}
}
func scalarString(value any) (string, error) {
switch value := value.(type) {
case string:
return value, nil
case float64:
return fmt.Sprintf("%v", value), nil
case bool:
return fmt.Sprintf("%t", value), nil
default:
return "", ErrInvalidConfig
}
}
func imageExtension(mimeType string) string {
if mimeType == "image/jpeg" {
return ".jpg"
}
return ".png"
}
func fromHTTPStatus(status int, body []byte) *Error {
if status == http.StatusBadRequest && policyRejected(body) {
return providerError(CodePolicyRejected, FailurePolicyRejected)
}
switch status {
case http.StatusTooManyRequests:
return providerError(CodeRateLimited, FailureRateLimited)
case http.StatusBadRequest:
return providerError(CodeBadRequest, FailureBadRequest)
case http.StatusUnauthorized:
return providerError(CodeUnauthorized, FailureUnauthorized)
}
if status >= 500 {
return providerError(CodeServerError, FailureServer)
}
return providerError(CodeUnknown, FailureOther)
}
func policyRejected(body []byte) bool {
lower := strings.ToLower(string(body))
return strings.Contains(lower, "content_policy") || strings.Contains(lower, "safety")
}
func classifyNetworkError(err error) *Error {
if errors.Is(err, context.DeadlineExceeded) {
return providerError(CodeTimeout, FailureTimeout)
}
return providerError(CodeConnection, FailureConnection)
}
func providerError(code ErrorCode, class FailureClass) *Error {
return &Error{Code: code, Class: class, Message: "upstream request failed"}
}
var _ Client = (*OpenAI)(nil)