421 lines
13 KiB
Go
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)
|