Files
creator-hub/internal/creator/openai.go
T
rogee f345e368df
douyin-release-gate / verify (push) Failing after 19m29s
feat: configure OpenAI-compatible AI services and discover models
2026-10-07 17:17:33 +08:00

248 lines
8.8 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package creator
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"sort"
"strings"
"time"
"github.com/sirupsen/logrus"
)
// AIConnection is the unsaved connection used to discover available models.
type AIConnection struct {
BaseURL string `json:"base_url"`
APIKey string `json:"api_key"`
}
func validateAIConnection(connection AIConnection) (AIConnection, error) {
connection.BaseURL = strings.TrimRight(strings.TrimSpace(connection.BaseURL), "/")
connection.APIKey = strings.TrimSpace(connection.APIKey)
parsed, err := url.Parse(connection.BaseURL)
if err != nil || parsed.Hostname() == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.RawQuery != "" || parsed.Fragment != "" || parsed.User != nil || strings.HasSuffix(parsed.Path, "/models") || strings.HasSuffix(parsed.Path, "/chat/completions") {
return AIConnection{}, fmt.Errorf("%w: BASEURL 必须为 HTTP(S) 接口根地址,不包含 /models、/chat/completions、查询参数或片段", ErrInvalid)
}
if connection.APIKey == "" {
return AIConnection{}, fmt.Errorf("%w: 请填写 APIKEY", ErrInvalid)
}
return connection, nil
}
type OpenAIClient struct {
BaseURL string
APIKey string
Model string
HTTPClient *http.Client
}
func NewOpenAIClient(baseURL, apiKey, model string, client *http.Client) (*OpenAIClient, error) {
connection, err := validateAIConnection(AIConnection{BaseURL: baseURL, APIKey: apiKey})
if err != nil {
return nil, fmt.Errorf("%w: %v", ErrUnavailable, err)
}
model = strings.TrimSpace(model)
if model == "" {
return nil, fmt.Errorf("%w: 请在 AI 服务中选择 MODEL", ErrUnavailable)
}
return &OpenAIClient{BaseURL: connection.BaseURL, APIKey: connection.APIKey, Model: model, HTTPClient: openAIHTTPClient(client)}, nil
}
func openAIHTTPClient(client *http.Client) *http.Client {
if client != nil {
return client
}
return &http.Client{Timeout: 60 * time.Second}
}
// ListOpenAIModels never substitutes a static list or another provider endpoint.
func ListOpenAIModels(ctx context.Context, connection AIConnection, client *http.Client) (models []string, resultErr error) {
connection, err := validateAIConnection(connection)
if err != nil {
return nil, err
}
defer func() {
if resultErr != nil {
logrus.WithField("base_url", connection.BaseURL).WithError(resultErr).Error("OpenAI model discovery failed")
}
}()
api := &OpenAIClient{BaseURL: connection.BaseURL, APIKey: connection.APIKey, HTTPClient: openAIHTTPClient(client)}
body, err := api.request(ctx, http.MethodGet, "/models", nil)
if err != nil {
return nil, err
}
var response struct {
Data []struct {
ID string `json:"id"`
} `json:"data"`
}
if err := json.Unmarshal(body, &response); err != nil {
return nil, fmt.Errorf("OpenAI Compatible 模型列表格式错误: %w", err)
}
if len(response.Data) == 0 {
return nil, fmt.Errorf("OpenAI Compatible 服务未返回可选模型(data 为空或缺失)")
}
models = make([]string, 0, len(response.Data))
seen := make(map[string]bool)
for _, model := range response.Data {
if strings.TrimSpace(model.ID) == "" {
return nil, fmt.Errorf("OpenAI Compatible 模型列表包含缺失的模型 ID")
}
if !seen[model.ID] {
models = append(models, model.ID)
seen[model.ID] = true
}
}
sort.Strings(models)
logrus.WithFields(logrus.Fields{"base_url": connection.BaseURL, "model_count": len(models)}).Info("OpenAI model discovery completed")
return models, nil
}
func (client *OpenAIClient) request(ctx context.Context, method, path string, payload []byte) (body []byte, resultErr error) {
started := time.Now()
defer func() {
entry := logrus.WithFields(logrus.Fields{"base_url": client.BaseURL, "operation": path, "model": client.Model, "elapsed_ms": time.Since(started).Milliseconds()})
if resultErr != nil {
entry.WithError(resultErr).Error("OpenAI request failed")
} else {
entry.Info("OpenAI HTTP request completed")
}
}()
request, err := http.NewRequestWithContext(ctx, method, client.BaseURL+path, bytes.NewReader(payload))
if err != nil {
return nil, fmt.Errorf("OpenAI request: %w", err)
}
request.Header.Set("Authorization", "Bearer "+client.APIKey)
request.Header.Set("Accept", "application/json")
if payload != nil {
request.Header.Set("Content-Type", "application/json")
}
response, err := client.HTTPClient.Do(request)
if err != nil {
return nil, fmt.Errorf("OpenAI %s: %w", path, err)
}
defer response.Body.Close()
const maxResponseBytes = 2 << 20
body, err = io.ReadAll(io.LimitReader(response.Body, maxResponseBytes+1))
if err != nil {
return nil, fmt.Errorf("OpenAI %s 读取响应失败: %w", path, err)
}
if len(body) > maxResponseBytes {
return nil, fmt.Errorf("OpenAI %s 响应超过 2 MB", path)
}
if response.StatusCode < 200 || response.StatusCode >= 300 {
return nil, fmt.Errorf("OpenAI %s HTTP %d: %s", path, response.StatusCode, strings.TrimSpace(string(body)))
}
return body, nil
}
type chatMessage struct {
Role string `json:"role"`
Content string `json:"content"`
}
func (client *OpenAIClient) chat(ctx context.Context, system, user string) (string, error) {
if client == nil || client.APIKey == "" || client.Model == "" || client.BaseURL == "" {
return "", ErrUnavailable
}
payload, err := json.Marshal(struct {
Model string `json:"model"`
Messages []chatMessage `json:"messages"`
}{client.Model, []chatMessage{{"system", system}, {"user", user}}})
if err != nil {
return "", err
}
body, err := client.request(ctx, http.MethodPost, "/chat/completions", payload)
if err != nil {
return "", err
}
var parsed struct {
Choices []struct {
Message chatMessage `json:"message"`
} `json:"choices"`
}
if err := json.Unmarshal(body, &parsed); err != nil {
return "", fmt.Errorf("OpenAI response: %w", err)
}
if len(parsed.Choices) == 0 || strings.TrimSpace(parsed.Choices[0].Message.Content) == "" {
return "", fmt.Errorf("OpenAI response contains no text")
}
return strings.TrimSpace(parsed.Choices[0].Message.Content), nil
}
func (client *OpenAIClient) MatchTheme(ctx context.Context, title, body, topic string) (bool, string, error) {
text, err := client.chat(ctx, `判断作品是否匹配主题。只返回 JSON:{"match":true或false,"reason":"简短理由"}。不得补充事实。`, fmt.Sprintf("主题:%s\n标题:%s\n正文:%s", topic, title, body))
if err != nil {
return false, "", err
}
return parseDecision(text)
}
func (client *OpenAIClient) MatchLead(ctx context.Context, workTitle, comment, rule string) (bool, string, error) {
text, err := client.chat(ctx, `判断评论是否符合线索规则。只返回 JSON:{"match":true或false,"reason":"简短理由"}。不得猜测用户身份或联系方式。`, fmt.Sprintf("规则:%s\n作品:%s\n评论:%s", rule, workTitle, comment))
if err != nil {
return false, "", err
}
return parseDecision(text)
}
func (client *OpenAIClient) Generate(ctx context.Context, prompt, input string) (string, error) {
return client.chat(ctx, "按用户要求生成文本,不执行任何平台操作。", fmt.Sprintf("要求:%s\n输入:%s", prompt, input))
}
func parseDecision(text string) (bool, string, error) {
text = strings.TrimSpace(text)
var decision struct {
Match *bool `json:"match"`
Reason string `json:"reason"`
}
if err := json.Unmarshal([]byte(text), &decision); err != nil {
return false, "", fmt.Errorf("AI decision JSON: %w", err)
}
if decision.Match == nil {
return false, "", fmt.Errorf("AI decision omitted match")
}
return *decision.Match, strings.TrimSpace(decision.Reason), nil
}
// ConfiguredOpenAI reads persisted settings per operation, so changes take effect
// without restarting the service or changing an environment variable.
type ConfiguredOpenAI struct {
Store *Store
HTTPClient *http.Client
}
func (configured *ConfiguredOpenAI) client(ctx context.Context) (*OpenAIClient, error) {
if configured == nil || configured.Store == nil {
return nil, ErrUnavailable
}
settings, err := configured.Store.GetSettings(ctx)
if err != nil {
return nil, err
}
return NewOpenAIClient(settings.AIBaseURL, settings.AIAPIKey, settings.AIModel, configured.HTTPClient)
}
func (configured *ConfiguredOpenAI) MatchTheme(ctx context.Context, title, body, topic string) (bool, string, error) {
client, err := configured.client(ctx)
if err != nil {
return false, "", err
}
return client.MatchTheme(ctx, title, body, topic)
}
func (configured *ConfiguredOpenAI) MatchLead(ctx context.Context, title, comment, rule string) (bool, string, error) {
client, err := configured.client(ctx)
if err != nil {
return false, "", err
}
return client.MatchLead(ctx, title, comment, rule)
}
func (configured *ConfiguredOpenAI) Generate(ctx context.Context, prompt, input string) (string, error) {
client, err := configured.client(ctx)
if err != nil {
return "", err
}
return client.Generate(ctx, prompt, input)
}