package service import ( "bytes" "context" "encoding/json" "fmt" "io" "math" "net/http" "strings" "time" "infogenie-backend/internal/database" "infogenie-backend/internal/model" ) type ChatMessage struct { Role string `json:"role"` Content string `json:"content"` } type chatRequest struct { Model string `json:"model"` Messages []ChatMessage `json:"messages"` Temperature float64 `json:"temperature"` MaxTokens int `json:"max_tokens"` } type chatResponse struct { Choices []struct { Message struct { Content string `json:"content"` } `json:"message"` } `json:"choices"` } // loadAIConfig 从数据库读取AI配置 func loadAIConfig(provider string) (apiKey, apiBase, defaultModel string, models []string, ok bool) { if database.DB == nil { return "", "", "", nil, false } var config model.AIConfig if err := database.DB.Where("provider = ? AND is_enabled = ?", provider, true).First(&config).Error; err != nil { return "", "", "", nil, false } // 解析models JSON var modelList []string if config.Models != "" { if err := json.Unmarshal([]byte(config.Models), &modelList); err != nil { // 如果解析失败,返回空的模型列表 modelList = []string{} } } return config.APIKey, config.APIBase, config.DefaultModel, modelList, true } // loadRuntimeDeepSeek 读取管理员在后台配置的 DeepSeek 兼容接口(OpenAI 格式),优先于 ai_config.json func loadRuntimeDeepSeek() (apiBase, apiKey, defModel string, ok bool) { if database.DB == nil { return "", "", "", false } var row model.SiteAIRuntime if err := database.DB.First(&row, 1).Error; err != nil { return "", "", "", false } base := strings.TrimSpace(row.APIBase) key := strings.TrimSpace(row.APIKey) dm := strings.TrimSpace(row.DefaultModel) if base != "" && key != "" { return base, key, dm, true } return "", "", "", false } // openDeepSeekChatURL 解析 DeepSeek 兼容 /chat/completions 的完整 URL、密钥与最终落库模型名 func openDeepSeekChatURL(model string) (fullURL, apiKey, resolvedModel string, err error) { if base, key, defModel, ok := loadRuntimeDeepSeek(); ok { if model == "" { model = defModel } if model == "" { model = "deepseek-chat" } return strings.TrimSuffix(base, "/") + "/chat/completions", key, model, nil } apiKey, apiBase, defaultModel, models, ok := loadAIConfig("deepseek") if !ok { return "", "", "", fmt.Errorf("DeepSeek配置未设置,请在管理员后台配置API Key和Base URL") } if model == "" { model = defaultModel } if model == "" { model = "deepseek-chat" } if len(models) > 0 { allowed := false for _, m := range models { if m == model { allowed = true break } } if !allowed { model = models[0] } } return strings.TrimSuffix(apiBase, "/") + "/chat/completions", apiKey, model, nil } func CallDeepSeek(messages []ChatMessage, model string, maxRetries int) (string, error) { urlStr, key, m, err := openDeepSeekChatURL(model) if err != nil { return "", err } return callOpenAICompatible(urlStr, key, m, messages, maxRetries, 90*time.Second) } // openKimiChatURL 解析 Kimi /v1/chat/completions func openKimiChatURL(model string) (fullURL, apiKey, resolvedModel string, err error) { apiKey, apiBase, defaultModel, models, ok := loadAIConfig("kimi") if !ok { return "", "", "", fmt.Errorf("Kimi配置未设置,请在管理员后台配置API Key和Base URL") } if model == "" { model = defaultModel } if model == "" { model = "kimi-k2-0905-preview" } if len(models) > 0 { allowed := false for _, m := range models { if m == model { allowed = true break } } if !allowed { model = models[0] } } return strings.TrimSuffix(apiBase, "/") + "/v1/chat/completions", apiKey, model, nil } func CallKimi(messages []ChatMessage, model string) (string, error) { urlStr, key, m, err := openKimiChatURL(model) if err != nil { return "", err } return callOpenAICompatible(urlStr, key, m, messages, 1, 30*time.Second) } func callOpenAICompatible(url, apiKey, model string, messages []ChatMessage, maxRetries int, timeout time.Duration) (string, error) { reqBody := chatRequest{ Model: model, Messages: messages, Temperature: 0.7, MaxTokens: 2000, } bodyBytes, err := json.Marshal(reqBody) if err != nil { return "", fmt.Errorf("序列化请求失败: %w", err) } client := &http.Client{Timeout: timeout} var lastErr error for attempt := 0; attempt < maxRetries; attempt++ { req, _ := http.NewRequest("POST", url, bytes.NewReader(bodyBytes)) req.Header.Set("Authorization", "Bearer "+apiKey) req.Header.Set("Content-Type", "application/json") resp, err := client.Do(req) if err != nil { lastErr = err if attempt < maxRetries-1 { backoff := time.Duration(math.Pow(2, float64(attempt))) * time.Second time.Sleep(backoff) continue } return "", fmt.Errorf("API调用异常(已重试%d次): %w", maxRetries, err) } respBody, _ := io.ReadAll(resp.Body) resp.Body.Close() if resp.StatusCode == 200 { var result chatResponse if err := json.Unmarshal(respBody, &result); err != nil { return "", fmt.Errorf("解析响应失败: %w", err) } if len(result.Choices) == 0 { return "", fmt.Errorf("AI未返回有效内容") } return result.Choices[0].Message.Content, nil } lastErr = fmt.Errorf("API调用失败: %d - %s", resp.StatusCode, string(respBody)) if attempt < maxRetries-1 { backoff := time.Duration(math.Pow(2, float64(attempt))) * time.Second time.Sleep(backoff) } } return "", lastErr } func CallAI(provider, model string, messages []ChatMessage) (string, error) { switch provider { case "deepseek": return CallDeepSeek(messages, model, 3) case "kimi": return CallKimi(messages, model) default: return "", fmt.Errorf("不支持的AI提供商: %s,目前支持的提供商: deepseek, kimi", provider) } } // OpenAIChatStream 向上游发起 stream:true 的请求;返回的 ReadCloser 需由调用方 Close。statusCode 非 200 时 body 已读完并关闭,rc 为 nil。 func OpenAIChatStream(ctx context.Context, provider, model string, messages []ChatMessage) (rc io.ReadCloser, statusCode int, err error) { var urlStr, apiKey, m string switch provider { case "deepseek": urlStr, apiKey, m, err = openDeepSeekChatURL(model) case "kimi": urlStr, apiKey, m, err = openKimiChatURL(model) default: return nil, 0, fmt.Errorf("不支持的AI提供商: %s", provider) } if err != nil { return nil, 0, err } streamBody := map[string]interface{}{ "model": m, "messages": messages, "temperature": 0.7, "max_tokens": 2000, "stream": true, } bodyBytes, jerr := json.Marshal(streamBody) if jerr != nil { return nil, 0, fmt.Errorf("序列化请求失败: %w", jerr) } req, rerr := http.NewRequestWithContext(ctx, "POST", urlStr, bytes.NewReader(bodyBytes)) if rerr != nil { return nil, 0, rerr } req.Header.Set("Authorization", "Bearer "+apiKey) req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "text/event-stream") client := &http.Client{} resp, derr := client.Do(req) if derr != nil { return nil, 0, derr } if resp.StatusCode != http.StatusOK { b, _ := io.ReadAll(resp.Body) resp.Body.Close() return nil, resp.StatusCode, fmt.Errorf("%s", strings.TrimSpace(string(b))) } return resp.Body, http.StatusOK, nil }