248 lines
8.8 KiB
Go
248 lines
8.8 KiB
Go
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)
|
||
}
|