diff --git a/AGENTS.md b/AGENTS.md index 76691f9..aa41aba 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -64,6 +64,8 @@ 评论聚合:独立页面通过“我的作品评论”“竞品作品评论”两个 TAB 严格区分来源,各 TAB 保留筛选与页码。列表展示评论发布时间、评论者、内容、所属账号名称及对应作品的小封面;所属账号列与筛选下拉均不展示 UID,缺名仅显示前端“未命名账号”,评论者显示不变;作品使用本地封面,点击在新标签页打开抖音作品页面,缺图明确提示,不展示作品标题、不使用远程图片兜底。支持所属账号筛选,以及最近 1/6/12 小时、1/3/5/7 天筛选,默认最近 1 天。按评论发布时间计算范围,最新评论在前,同时间按评论 ID 倒序;未记录发布时间、未来时间及未采集评论不参与。筛选先作用于全部已采集评论再分页,条件或每页数量变化回到第一页;只读展示,不新增采集机制。 +AI 服务:仅支持 OpenAI Compatible,BASEURL、APIKEY、MODEL 在“系统设置 → AI 服务”保存;BASEURL 是接口根地址,不补默认地址,不包含 `/models` 或 `/chat/completions`。填写地址与密钥后通过 `GET /models` 自动获取模型,MODEL 仅用可搜索下拉选择,支持手动刷新;获取失败、模型为空或已保存模型不可用必须明确提示,不预设模型或手填兜底。实际分析每次使用数据库当前配置,不使用百炼专用客户端或 `BAILIAN_*` 环境变量,不需要额外审批勾选。仅修改 AI 配置或保存未变化的采集配置不得重置采集窗口、游标、进度、完成时间或指标计划,也不因采集运行中而拒绝;实际修改采集配置仍保留运行中冲突检查。原有 MODEL 在数据库更新时保留,首次使用需补齐地址与密钥。 + 前端框架:Umi Max 4.7 + React 19; 组件库:antd 6.6.5 + @ant-design/pro-components 3.x(beta 线)+ @ant-design/icons;仅使用 antd/pro 默认组件原样实现,禁止自定义封装与样式魔改;组件不满足业务时改交互逻辑适配组件; diff --git a/compose.yaml b/compose.yaml index 8fa4d86..91f780d 100644 --- a/compose.yaml +++ b/compose.yaml @@ -7,8 +7,6 @@ services: CONTROL_PLANE_PASSWORD: ${CONTROL_PLANE_PASSWORD:?required} CREATORHUB_CREDENTIAL_MASTER_KEY: >- ${CREATORHUB_CREDENTIAL_MASTER_KEY:?required} - BAILIAN_API_KEY: ${BAILIAN_API_KEY:-} - BAILIAN_BASE_URL: ${BAILIAN_BASE_URL:-} CREATOR_MEDIA_DIR: /var/lib/creatorhub/materials CREATOR_COVER_DIR: /var/lib/creatorhub/covers CREATOR_TRANSCRIPTION_BIN: ${CREATOR_TRANSCRIPTION_BIN:-} diff --git a/internal/controlplane/api/ai_models_test.go b/internal/controlplane/api/ai_models_test.go new file mode 100644 index 0000000..9c5b25f --- /dev/null +++ b/internal/controlplane/api/ai_models_test.go @@ -0,0 +1,159 @@ +package api + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "os" + "strings" + "testing" + "time" + + "git.ipao.vip/rogee/creator-hub/internal/creator" + "github.com/gofiber/fiber/v3" +) + +func TestAIModelsDiscoveryBeforeSavingSettings(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet || r.URL.Path != "/custom/v1/models" || r.Header.Get("Authorization") != "Bearer draft-key" { + t.Errorf("unexpected discovery request: %s %s", r.Method, r.URL.Path) + } + _, _ = w.Write([]byte(`{"data":[{"id":"model-b"},{"id":"model-a"}]}`)) + })) + defer server.Close() + app := fiber.New() + registerCreatorWithServices(app, nil, nil, nil, nil) + for _, tc := range []struct { + body string + status int + contains string + }{ + {`{"base_url":"` + server.URL + `/custom/v1","api_key":"draft-key"}`, 200, `"model-a"`}, + {`{`, 400, ""}, + {`{"base_url":"","api_key":"key"}`, 400, ""}, + {`{"base_url":"https://example.com/v1"}`, 400, ""}, + } { + request := httptest.NewRequest(http.MethodPost, "/api/creator/settings/models", strings.NewReader(tc.body)) + request.Header.Set("Content-Type", "application/json") + response, err := app.Test(request) + if err != nil { + t.Fatal(err) + } + body, err := io.ReadAll(response.Body) + response.Body.Close() + if err != nil { + t.Fatal(err) + } + if response.StatusCode != tc.status || !strings.Contains(string(body), tc.contains) { + t.Fatalf("response=%d %s", response.StatusCode, body) + } + } +} + +func TestAISettingsRoutesSaveReloadAndUseConfiguredAnalyzer(t *testing.T) { + databaseURL := os.Getenv("CREATORHUB_POSTGRES_TEST_URL") + if databaseURL == "" { + t.Skip("set CREATORHUB_POSTGRES_TEST_URL to run PostgreSQL integration tests") + } + store, phaseA, ctx := openCreatorIntegrationStoreForAPITest(t, databaseURL) + var requests int + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests++ + if r.URL.Path != "/v1/chat/completions" || r.Header.Get("Authorization") != "Bearer saved-api-key" { + t.Errorf("wrong configured connection: %s", r.URL.Path) + } + var payload struct { + Model string `json:"model"` + } + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + t.Error(err) + } + if payload.Model != "selected-model" { + t.Errorf("model=%s", payload.Model) + } + _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"match\":true,\"reason\":\"符合规则\"}"}}]}`)) + })) + defer server.Close() + app := fiber.New() + registerCreatorWithServices(app, store, phaseA, nil, nil) + settings, err := store.GetSettings(ctx) + if err != nil { + t.Fatal(err) + } + input := creator.SettingsUpdate{LookbackDays: settings.LookbackDays, NewWorkIntervalSeconds: settings.NewWorkIntervalSeconds, MetricInitialIntervalSeconds: settings.MetricInitialIntervalSeconds, MetricMultiplier: settings.MetricMultiplier, MetricMaxIntervalSeconds: settings.MetricMaxIntervalSeconds, MetricAgeSeconds: settings.MetricAgeSeconds, AIBaseURL: server.URL + "/v1", AIAPIKey: "saved-api-key", AIModel: "selected-model"} + encoded, err := json.Marshal(input) + if err != nil { + t.Fatal(err) + } + if status, body := apiRequest(t, app, http.MethodPut, "/api/creator/settings", string(encoded)); status != 200 { + t.Fatalf("save=%d %s", status, body) + } + status, body := apiRequest(t, app, http.MethodGet, "/api/creator/settings", "") + if status != 200 { + t.Fatalf("reload=%d %s", status, body) + } + var saved creator.Settings + if err := json.Unmarshal([]byte(body), &saved); err != nil || saved.AIBaseURL != input.AIBaseURL || saved.AIAPIKey != input.AIAPIKey || saved.AIModel != input.AIModel { + t.Fatalf("reload=%+v err=%v", saved, err) + } + owner, err := store.CreateCompetitor(ctx, creator.CompetitorInput{Platform: "douyin", PlatformAccountKey: "configured-ai-owner", UniqueID: "ai-owner", Nickname: "analysis", HomepageURL: "https://www.douyin.com/user/configured-ai-owner"}) + if err != nil { + t.Fatal(err) + } + now := time.Now().UTC() + work, _, err := store.UpsertWork(ctx, creator.WorkInput{Platform: "douyin", WorkKey: "configured-ai-work", SourceType: creator.SourceCompetitor, SourceID: owner.ID, Title: "work", PublishedAt: &now, PublishedAtStatus: "verified"}, now) + if err != nil { + t.Fatal(err) + } + comment, _, err := store.SaveComment(ctx, creator.CommentInput{Platform: "douyin", WorkID: work.ID, CommentKey: "configured-ai-comment", AuthorUID: "author", Content: "comment", CommentType: "top_level", PublishedAt: &now}) + if err != nil { + t.Fatal(err) + } + rule, err := store.CreateRule(ctx, creator.LeadRuleInput{Name: "AI analysis rule", Enabled: true, SourceType: creator.SourceCompetitor, Topic: "work", IncludeKeywords: []string{"comment"}, AIRequirement: "rule"}) + if err != nil { + t.Fatal(err) + } + status, body = apiRequest(t, app, http.MethodPost, fmt.Sprintf("/api/creator/comments/%s/analyze", comment.ID), fmt.Sprintf(`{"rule_id":%q}`, rule.ID)) + if status != 200 || requests != 2 || !strings.Contains(body, "符合规则") { + t.Fatalf("analysis=%d %s requests=%d", status, body, requests) + } +} + +func apiRequest(t *testing.T, app *fiber.App, method, path, body string) (int, string) { + t.Helper() + request := httptest.NewRequest(method, path, strings.NewReader(body)) + request.Header.Set("Content-Type", "application/json") + response, err := app.Test(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + result, err := io.ReadAll(response.Body) + if err != nil { + t.Fatal(err) + } + return response.StatusCode, string(result) +} + +func TestAIModelsDiscoverySurfacesProviderFailure(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(401) + _, _ = w.Write([]byte(`{"error":{"message":"invalid key"}}`)) + })) + defer server.Close() + app := fiber.New() + registerCreatorWithServices(app, nil, nil, nil, nil) + request := httptest.NewRequest(http.MethodPost, "/api/creator/settings/models", strings.NewReader(`{"base_url":"`+server.URL+`","api_key":"wrong"}`)) + request.Header.Set("Content-Type", "application/json") + response, err := app.Test(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + body, _ := io.ReadAll(response.Body) + if response.StatusCode != 502 || !strings.Contains(string(body), "401") || !strings.Contains(string(body), "invalid key") { + t.Fatalf("response=%d %s", response.StatusCode, body) + } +} diff --git a/internal/controlplane/api/creator.go b/internal/controlplane/api/creator.go index 2980f89..acfbfc4 100644 --- a/internal/controlplane/api/creator.go +++ b/internal/controlplane/api/creator.go @@ -59,6 +59,9 @@ func registerCreator(app *fiber.App, store *creator.Store, phaseAStore *accountd } func registerCreatorWithServices(app *fiber.App, store *creator.Store, phaseAStore *accountdomain.Store, hubStore *hub.Store, analyzer creator.ThemeAnalyzer) { + if analyzer == nil && store != nil { + analyzer = &creator.ConfiguredOpenAI{Store: store} + } registerEnvironmentLoginRoutes(app, store, hubStore) registerAccountEventRoutes(app, store) registerPrivateMessageRoutes(app, store, gatewayPrivateMessageSender(hubStore, store)) @@ -75,6 +78,21 @@ func registerCreatorWithServices(app *fiber.App, store *creator.Store, phaseASto app.Post(path, func(c fiber.Ctx) error { return creatorError(c, creator.ErrConflict) }) } + app.Post("/api/creator/settings/models", func(c fiber.Ctx) error { + var connection creator.AIConnection + if err := decodeCreator(c, &connection); err != nil { + return creatorError(c, err) + } + models, err := creator.ListOpenAIModels(c.Context(), connection, nil) + if err != nil { + if errors.Is(err, creator.ErrInvalid) { + return creatorError(c, err) + } + return c.Status(fiber.StatusBadGateway).JSON(fiber.Map{"message": err.Error()}) + } + return c.JSON(fiber.Map{"models": models}) + }) + app.Get("/api/creator/settings", func(c fiber.Ctx) error { settings, err := store.GetSettings(c.Context()) if err != nil { diff --git a/internal/controlplane/app/app.go b/internal/controlplane/app/app.go index 6de768b..c7fd7c2 100644 --- a/internal/controlplane/app/app.go +++ b/internal/controlplane/app/app.go @@ -35,7 +35,6 @@ import ( type config struct { listenAddr, webDir, databaseURL, credentialStoreDir string username, password string - aiAPIKey, aiBaseURL string credentialMasterKey []byte logLevel logrus.Level } @@ -157,8 +156,6 @@ func loadConfig() (config, error) { _ = v.BindEnv("log_level", "LOG_LEVEL") _ = v.BindEnv("username", "CONTROL_PLANE_USERNAME") _ = v.BindEnv("password", "CONTROL_PLANE_PASSWORD") - _ = v.BindEnv("ai_api_key", "BAILIAN_API_KEY") - _ = v.BindEnv("ai_base_url", "BAILIAN_BASE_URL") level, err := logrus.ParseLevel(v.GetString("log_level")) if err != nil { @@ -171,8 +168,6 @@ func loadConfig() (config, error) { credentialStoreDir: strings.TrimSpace(v.GetString("credential_store_dir")), username: strings.TrimSpace(v.GetString("username")), password: v.GetString("password"), - aiAPIKey: strings.TrimSpace(v.GetString("ai_api_key")), - aiBaseURL: strings.TrimSpace(v.GetString("ai_base_url")), logLevel: level, } if cfg.listenAddr == "" { @@ -202,12 +197,6 @@ func loadConfig() (config, error) { (databaseURL.Scheme != "postgres" && databaseURL.Scheme != "postgresql") { return config{}, errors.New("DATABASE_URL must be a postgres URL with a host") } - if cfg.aiBaseURL != "" { - aiURL, parseErr := url.Parse(cfg.aiBaseURL) - if parseErr != nil || aiURL.Host == "" || (aiURL.Scheme != "http" && aiURL.Scheme != "https") || aiURL.User != nil { - return config{}, errors.New("BAILIAN_BASE_URL must be an HTTP(S) URL without credentials") - } - } return cfg, nil } diff --git a/internal/creator/ai_settings_test.go b/internal/creator/ai_settings_test.go new file mode 100644 index 0000000..d3f743f --- /dev/null +++ b/internal/creator/ai_settings_test.go @@ -0,0 +1,115 @@ +package creator + +import ( + "errors" + "reflect" + "testing" + "time" +) + +func TestAISettingsValidation(t *testing.T) { + if _, err := validateAISettings(SettingsUpdate{}); err != nil { + t.Fatalf("unconfigured settings: %v", err) + } + for _, input := range []SettingsUpdate{ + {AIBaseURL: "https://example.com/v1"}, + {AIAPIKey: "key"}, + {AIModel: "model"}, + {AIBaseURL: "https://example.com/v1", AIAPIKey: "key"}, + {AIBaseURL: "bad-url", AIAPIKey: "key", AIModel: "model"}, + } { + if _, err := validateAISettings(input); !errors.Is(err, ErrInvalid) { + t.Fatalf("incomplete settings accepted: %+v err=%v", input, err) + } + } + input, err := validateAISettings(SettingsUpdate{AIBaseURL: " https://example.com/custom/v1/ ", AIAPIKey: " key ", AIModel: " model "}) + if err != nil || input.AIBaseURL != "https://example.com/custom/v1" || input.AIAPIKey != "key" || input.AIModel != "model" { + t.Fatalf("normalization=%+v err=%v", input, err) + } +} + +func TestAISettingsDoNotResetActiveCollectionOrMetrics(t *testing.T) { + store, phaseAStore, ctx := openCreatorIntegrationStore(t) + accountID := createIntegrationAccount(t, ctx, phaseAStore, "ai-settings-owner") + now := time.Now().UTC().Truncate(time.Second) + start, end, err := NewCollectionWindow(now, 30) + if err != nil { + t.Fatal(err) + } + lease, err := store.beginCheckpoint(ctx, SourceOwned, accountID, "works", start, end) + if err != nil { + t.Fatal(err) + } + if lease == "" { + t.Fatal("checkpoint lease missing") + } + if _, err := store.db.ExecContext(ctx, `UPDATE creator_collection_checkpoint SET cursor='page-4',last_error='keep this',last_completed_at=$1 WHERE source_type=$2 AND source_id=$3`, now.Add(-time.Hour), SourceOwned, accountID); err != nil { + t.Fatal(err) + } + before, err := store.checkpoint(ctx, SourceOwned, accountID, "works") + if err != nil { + t.Fatal(err) + } + work, _, err := store.UpsertWork(ctx, WorkInput{Platform: "douyin", WorkKey: "ai-settings-work", SourceType: SourceOwned, SourceID: accountID, Title: "work", PublishedAt: &now, PublishedAtStatus: "verified"}, now) + if err != nil { + t.Fatal(err) + } + var nextBefore, nextAfter *time.Time + if err := store.db.QueryRowContext(ctx, `SELECT metric_plan_next_at FROM creator_work WHERE work_id=$1`, work.ID).Scan(&nextBefore); err != nil { + t.Fatal(err) + } + settings, err := store.GetSettings(ctx) + if err != nil { + t.Fatal(err) + } + input := SettingsUpdate{LookbackDays: settings.LookbackDays, NewWorkIntervalSeconds: settings.NewWorkIntervalSeconds, MetricInitialIntervalSeconds: settings.MetricInitialIntervalSeconds, MetricMultiplier: settings.MetricMultiplier, MetricMaxIntervalSeconds: settings.MetricMaxIntervalSeconds, MetricAgeSeconds: settings.MetricAgeSeconds, AIBaseURL: "https://example.com/v1", AIAPIKey: "key", AIModel: "model"} + if _, err := store.UpdateSettings(ctx, input); err != nil { + t.Fatalf("AI-only save must work during collection: %v", err) + } + after, err := store.checkpoint(ctx, SourceOwned, accountID, "works") + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(before, after) { + t.Fatalf("checkpoint changed: before=%+v after=%+v", before, after) + } + if err := store.db.QueryRowContext(ctx, `SELECT metric_plan_next_at FROM creator_work WHERE work_id=$1`, work.ID).Scan(&nextAfter); err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(nextBefore, nextAfter) { + t.Fatalf("metric schedule changed: %v -> %v", nextBefore, nextAfter) + } + if _, err := store.UpdateSettings(ctx, input); err != nil { + t.Fatalf("unchanged save: %v", err) + } + input.AIBaseURL, input.AIAPIKey, input.AIModel = "", "", "" + if _, err := store.UpdateSettings(ctx, input); err != nil { + t.Fatalf("clear AI: %v", err) + } + saved, err := store.GetSettings(ctx) + if err != nil || saved.AIBaseURL != "" || saved.AIAPIKey != "" || saved.AIModel != "" { + t.Fatalf("clear was not persisted: %+v err=%v", saved, err) + } + input.LookbackDays++ + if _, err := store.UpdateSettings(ctx, input); !errors.Is(err, ErrConflict) { + t.Fatalf("collection change while running: %v", err) + } + if err := store.finishCheckpoint(ctx, SourceOwned, accountID, "works", lease, "succeeded", "", nil); err != nil { + t.Fatal(err) + } + if _, err := store.UpdateSettings(ctx, input); err != nil { + t.Fatalf("collection change after completion: %v", err) + } + if err := store.db.QueryRowContext(ctx, `SELECT metric_plan_next_at FROM creator_work WHERE work_id=$1`, work.ID).Scan(&nextAfter); err != nil || nextAfter == nil { + t.Fatalf("metric plan not recalculated: %v err=%v", nextAfter, err) + } + input.AIBaseURL = "not-a-url" + if _, err := store.UpdateSettings(ctx, input); !errors.Is(err, ErrInvalid) { + t.Fatalf("invalid URL: %v", err) + } + input.AIBaseURL = "" + input.LookbackDays = 0 + if _, err := store.UpdateSettings(ctx, input); !errors.Is(err, ErrInvalid) { + t.Fatalf("invalid collection setting: %v", err) + } +} diff --git a/internal/creator/bailian.go b/internal/creator/bailian.go deleted file mode 100644 index 80f94ca..0000000 --- a/internal/creator/bailian.go +++ /dev/null @@ -1,186 +0,0 @@ -package creator - -import ( - "bytes" - "context" - "encoding/json" - "fmt" - "io" - "net/http" - "strings" - "time" -) - -const defaultBailianBaseURL = "https://dashscope.aliyuncs.com/compatible-mode/v1" - -type BailianClient struct { - BaseURL string - APIKey string - Model string - HTTPClient *http.Client -} - -func NewBailianClient(baseURL, apiKey, model string, client *http.Client) (*BailianClient, error) { - baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/") - if baseURL == "" { - baseURL = defaultBailianBaseURL - } - if strings.TrimSpace(apiKey) == "" || strings.TrimSpace(model) == "" { - return nil, fmt.Errorf("%w: BAILIAN_API_KEY and AI model are required", ErrUnavailable) - } - if client == nil { - client = &http.Client{Timeout: 60 * time.Second} - } - return &BailianClient{BaseURL: baseURL, APIKey: apiKey, Model: model, HTTPClient: client}, nil -} - -type bailianChatRequest struct { - Model string `json:"model"` - Messages []bailianChatMessage `json:"messages"` -} - -type bailianChatMessage struct { - Role string `json:"role"` - Content string `json:"content"` -} - -type bailianChatResponse struct { - Choices []struct { - Message bailianChatMessage `json:"message"` - } `json:"choices"` -} - -func (c *BailianClient) chat(ctx context.Context, instruction, input string) (string, error) { - if c == nil || c.HTTPClient == nil || strings.TrimSpace(c.APIKey) == "" || strings.TrimSpace(c.Model) == "" { - return "", fmt.Errorf("%w: BAILIAN client is not configured", ErrUnavailable) - } - body, err := json.Marshal(bailianChatRequest{ - Model: c.Model, - Messages: []bailianChatMessage{ - {Role: "system", Content: instruction}, - {Role: "user", Content: input}, - }, - }) - if err != nil { - return "", err - } - endpoint := c.BaseURL + "/chat/completions" - if strings.HasSuffix(c.BaseURL, "/chat/completions") { - endpoint = c.BaseURL - } - request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) - if err != nil { - return "", err - } - request.Header.Set("Authorization", "Bearer "+c.APIKey) - request.Header.Set("Content-Type", "application/json") - response, err := c.HTTPClient.Do(request) - if err != nil { - return "", err - } - defer response.Body.Close() - responseBody, err := io.ReadAll(io.LimitReader(response.Body, 2<<20)) - if err != nil { - return "", err - } - if response.StatusCode < 200 || response.StatusCode >= 300 { - return "", fmt.Errorf("bailian request failed with HTTP %d: %s", response.StatusCode, strings.TrimSpace(string(responseBody))) - } - var result bailianChatResponse - if err := json.Unmarshal(responseBody, &result); err != nil { - return "", fmt.Errorf("decode bailian response: %w", err) - } - if len(result.Choices) == 0 || strings.TrimSpace(result.Choices[0].Message.Content) == "" { - return "", fmt.Errorf("bailian response contained no message") - } - return strings.TrimSpace(result.Choices[0].Message.Content), nil -} - -func (c *BailianClient) Generate(ctx context.Context, instruction, input string) (string, error) { - return c.chat(ctx, instruction, input) -} - -func (c *BailianClient) MatchTheme(ctx context.Context, title, body, topic string) (bool, string, error) { - content, err := c.chat(ctx, - "判断作品是否符合给定主题。只返回 JSON,不要 Markdown 或额外文字,格式必须是 {\"match\":true或false,\"reason\":\"简短原因\"}。", - fmt.Sprintf("主题:%s\n标题:%s\n正文:%s", topic, title, body)) - if err != nil { - return false, "", err - } - var result struct { - Match *bool `json:"match"` - Reason string `json:"reason"` - } - if err := json.Unmarshal([]byte(content), &result); err != nil || result.Match == nil { - if err != nil { - return false, "", fmt.Errorf("decode bailian theme result: %w", err) - } - return false, "", fmt.Errorf("decode bailian theme result: match is required") - } - return *result.Match, strings.TrimSpace(result.Reason), nil -} - -func (c *BailianClient) MatchLead(ctx context.Context, work, comment, requirement string) (bool, string, error) { - content, err := c.chat(ctx, - "判断评论是否是有效业务线索。只返回 JSON,不要 Markdown 或额外文字,格式必须是 {\"match\":true或false,\"reason\":\"简短原因\"}。", - fmt.Sprintf("判定要求:%s\n作品:%s\n评论:%s", requirement, work, comment)) - if err != nil { - return false, "", err - } - var result struct { - Match *bool `json:"match"` - Reason string `json:"reason"` - } - if err := json.Unmarshal([]byte(content), &result); err != nil || result.Match == nil { - if err != nil { - return false, "", fmt.Errorf("decode bailian lead result: %w", err) - } - return false, "", fmt.Errorf("decode bailian lead result: match is required") - } - return *result.Match, strings.TrimSpace(result.Reason), nil -} - -type ConfiguredBailian struct { - Store *Store - APIKey string - BaseURL string - HTTPClient *http.Client -} - -func (b *ConfiguredBailian) client(ctx context.Context) (*BailianClient, error) { - if b == nil || b.Store == nil { - return nil, fmt.Errorf("%w: BAILIAN client is not configured", ErrUnavailable) - } - settings, err := b.Store.GetSettings(ctx) - if err != nil { - return nil, err - } - if !settings.AIConfigured || settings.AIProvider != "bailian" { - return nil, fmt.Errorf("%w: BAILIAN is not enabled in creator settings", ErrUnavailable) - } - return NewBailianClient(b.BaseURL, b.APIKey, settings.AIModel, b.HTTPClient) -} - -func (b *ConfiguredBailian) Generate(ctx context.Context, instruction, input string) (string, error) { - client, err := b.client(ctx) - if err != nil { - return "", err - } - return client.Generate(ctx, instruction, input) -} - -func (b *ConfiguredBailian) MatchTheme(ctx context.Context, title, body, topic string) (bool, string, error) { - client, err := b.client(ctx) - if err != nil { - return false, "", err - } - return client.MatchTheme(ctx, title, body, topic) -} - -func (b *ConfiguredBailian) MatchLead(ctx context.Context, work, comment, requirement string) (bool, string, error) { - client, err := b.client(ctx) - if err != nil { - return false, "", err - } - return client.MatchLead(ctx, work, comment, requirement) -} diff --git a/internal/creator/bailian_test.go b/internal/creator/bailian_test.go deleted file mode 100644 index e176c25..0000000 --- a/internal/creator/bailian_test.go +++ /dev/null @@ -1,95 +0,0 @@ -package creator - -import ( - "context" - "encoding/json" - "net/http" - "net/http/httptest" - "testing" -) - -func TestBailianClientGeneratesAndParsesDecisions(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPost || r.URL.Path != "/v1/chat/completions" { - t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) - } - if r.Header.Get("Authorization") != "Bearer test-key" { - t.Fatal("missing authorization") - } - var request bailianChatRequest - if err := json.NewDecoder(r.Body).Decode(&request); err != nil { - t.Fatal(err) - } - if request.Model != "qwen-test" || len(request.Messages) != 2 { - t.Fatalf("unexpected request body: %#v", request) - } - _ = json.NewEncoder(w).Encode(map[string]any{"choices": []any{map[string]any{ - "message": map[string]string{"role": "assistant", "content": `{"match":true,"reason":"相关"}`}, - }}}) - })) - defer server.Close() - - client, err := NewBailianClient(server.URL+"/v1", "test-key", "qwen-test", server.Client()) - if err != nil { - t.Fatal(err) - } - text, err := client.Generate(context.Background(), "要求", "输入") - if err != nil || text == "" { - t.Fatalf("Generate() = %q, %v", text, err) - } - match, reason, err := client.MatchLead(context.Background(), "作品", "评论", "要求") - if err != nil || !match || reason != "相关" { - t.Fatalf("MatchLead() = %v, %q, %v", match, reason, err) - } - match, reason, err = client.MatchTheme(context.Background(), "标题", "正文", "主题") - if err != nil || !match || reason != "相关" { - t.Fatalf("MatchTheme() = %v, %q, %v", match, reason, err) - } -} - -func TestBailianClientRejectsHTTPAndMalformedResponses(t *testing.T) { - responses := []string{"not-json", `{"choices":[]}`} - for _, body := range responses { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusBadGateway) - _, _ = w.Write([]byte(body)) - })) - client, err := NewBailianClient(server.URL, "key", "model", server.Client()) - if err != nil { - server.Close() - t.Fatal(err) - } - if _, err := client.Generate(context.Background(), "instruction", "input"); err == nil { - server.Close() - t.Fatal("expected HTTP failure") - } - server.Close() - } - var nilClient *BailianClient - if _, err := nilClient.Generate(context.Background(), "instruction", "input"); err == nil { - t.Fatal("nil client must fail closed") - } -} - -func TestBailianClientRejectsIncompleteConfiguration(t *testing.T) { - if _, err := NewBailianClient("", "", "qwen-test", nil); err == nil { - t.Fatal("expected missing key to fail") - } - if _, err := NewBailianClient("", "test-key", "", nil); err == nil { - t.Fatal("expected missing model to fail") - } -} - -func TestBailianClientRejectsMissingDecisionField(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{}"}}]}`)) - })) - defer server.Close() - client, err := NewBailianClient(server.URL, "test-key", "qwen-test", server.Client()) - if err != nil { - t.Fatal(err) - } - if _, _, err := client.MatchLead(context.Background(), "work", "comment", "requirement"); err == nil { - t.Fatal("expected missing match decision to fail closed") - } -} diff --git a/internal/creator/coverage_unit_test.go b/internal/creator/coverage_unit_test.go index 43076d0..060eb77 100644 --- a/internal/creator/coverage_unit_test.go +++ b/internal/creator/coverage_unit_test.go @@ -64,8 +64,8 @@ func TestCreatorPaginationGuards(t *testing.T) { } } -func TestConfiguredBailianWithoutStore(t *testing.T) { - var client *ConfiguredBailian +func TestConfiguredOpenAIWithoutStore(t *testing.T) { + var client *ConfiguredOpenAI if _, err := client.Generate(context.Background(), "instruction", "input"); !errors.Is(err, ErrUnavailable) { t.Fatalf("nil Generate: %v", err) } diff --git a/internal/creator/integration_test.go b/internal/creator/integration_test.go index 23ed0d5..00ac820 100644 --- a/internal/creator/integration_test.go +++ b/internal/creator/integration_test.go @@ -745,7 +745,7 @@ func TestCreatorPostgresSettingsResetCollectionCheckpoint(t *testing.T) { t.Fatal(err) } settings.LookbackDays = 3 - if _, err := store.UpdateSettings(ctx, SettingsUpdate{LookbackDays: settings.LookbackDays, NewWorkIntervalSeconds: settings.NewWorkIntervalSeconds, MetricInitialIntervalSeconds: settings.MetricInitialIntervalSeconds, MetricMultiplier: settings.MetricMultiplier, MetricMaxIntervalSeconds: settings.MetricMaxIntervalSeconds, MetricAgeSeconds: settings.MetricAgeSeconds, AIProvider: settings.AIProvider, AIModel: settings.AIModel, AIConfigured: settings.AIConfigured}); err != nil { + if _, err := store.UpdateSettings(ctx, SettingsUpdate{LookbackDays: settings.LookbackDays, NewWorkIntervalSeconds: settings.NewWorkIntervalSeconds, MetricInitialIntervalSeconds: settings.MetricInitialIntervalSeconds, MetricMultiplier: settings.MetricMultiplier, MetricMaxIntervalSeconds: settings.MetricMaxIntervalSeconds, MetricAgeSeconds: settings.MetricAgeSeconds, AIBaseURL: settings.AIBaseURL, AIAPIKey: settings.AIAPIKey, AIModel: settings.AIModel}); err != nil { t.Fatalf("update settings: %v", err) } checkpoint, err := store.checkpoint(ctx, SourceOwned, accountID, "works") @@ -760,7 +760,7 @@ func TestCreatorPostgresSettingsResetCollectionCheckpoint(t *testing.T) { t.Fatalf("begin running checkpoint: %v", err) } settings.LookbackDays = 4 - if _, err := store.UpdateSettings(ctx, SettingsUpdate{LookbackDays: settings.LookbackDays, NewWorkIntervalSeconds: settings.NewWorkIntervalSeconds, MetricInitialIntervalSeconds: settings.MetricInitialIntervalSeconds, MetricMultiplier: settings.MetricMultiplier, MetricMaxIntervalSeconds: settings.MetricMaxIntervalSeconds, MetricAgeSeconds: settings.MetricAgeSeconds, AIProvider: settings.AIProvider, AIModel: settings.AIModel, AIConfigured: settings.AIConfigured}); !errors.Is(err, ErrConflict) { + if _, err := store.UpdateSettings(ctx, SettingsUpdate{LookbackDays: settings.LookbackDays, NewWorkIntervalSeconds: settings.NewWorkIntervalSeconds, MetricInitialIntervalSeconds: settings.MetricInitialIntervalSeconds, MetricMultiplier: settings.MetricMultiplier, MetricMaxIntervalSeconds: settings.MetricMaxIntervalSeconds, MetricAgeSeconds: settings.MetricAgeSeconds, AIBaseURL: settings.AIBaseURL, AIAPIKey: settings.AIAPIKey, AIModel: settings.AIModel}); !errors.Is(err, ErrConflict) { t.Fatalf("settings changed during running checkpoint: err=%v", err) } if err := store.finishCheckpoint(ctx, SourceOwned, accountID, "works", lease, "succeeded", "", nil); err != nil { diff --git a/internal/creator/models.go b/internal/creator/models.go index 79c7c1b..8bba015 100644 --- a/internal/creator/models.go +++ b/internal/creator/models.go @@ -33,9 +33,9 @@ type Settings struct { MetricMultiplier float64 `json:"metric_multiplier"` MetricMaxIntervalSeconds int64 `json:"metric_max_interval_seconds"` MetricAgeSeconds int64 `json:"metric_age_seconds"` - AIProvider string `json:"ai_provider"` + AIBaseURL string `json:"ai_base_url"` + AIAPIKey string `json:"ai_api_key"` AIModel string `json:"ai_model"` - AIConfigured bool `json:"ai_configured"` UpdatedAt time.Time `json:"updated_at"` } @@ -46,9 +46,9 @@ type SettingsUpdate struct { MetricMultiplier float64 `json:"metric_multiplier"` MetricMaxIntervalSeconds int64 `json:"metric_max_interval_seconds"` MetricAgeSeconds int64 `json:"metric_age_seconds"` - AIProvider string `json:"ai_provider"` + AIBaseURL string `json:"ai_base_url"` + AIAPIKey string `json:"ai_api_key"` AIModel string `json:"ai_model"` - AIConfigured bool `json:"ai_configured"` } type AccountProfile struct { diff --git a/internal/creator/openai.go b/internal/creator/openai.go new file mode 100644 index 0000000..4e57bef --- /dev/null +++ b/internal/creator/openai.go @@ -0,0 +1,247 @@ +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) +} diff --git a/internal/creator/openai_test.go b/internal/creator/openai_test.go new file mode 100644 index 0000000..30f9e41 --- /dev/null +++ b/internal/creator/openai_test.go @@ -0,0 +1,177 @@ +package creator + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "testing" +) + +func TestOpenAIModelsUsesConfiguredRootAndKey(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet || r.URL.Path != "/custom/v1/models" || r.Header.Get("Authorization") != "Bearer test-key" { + t.Errorf("unexpected request: %s %s, authorization=%q", r.Method, r.URL.Path, r.Header.Get("Authorization")) + } + _, _ = w.Write([]byte(`{"data":[{"id":"z-model"},{"id":"a-model"},{"id":"z-model"}]}`)) + })) + defer server.Close() + got, err := ListOpenAIModels(context.Background(), AIConnection{BaseURL: " " + server.URL + "/custom/v1/ ", APIKey: " test-key "}, server.Client()) + if err != nil || !reflect.DeepEqual(got, []string{"a-model", "z-model"}) { + t.Fatalf("models=%v err=%v", got, err) + } +} + +func TestOpenAIModelsRejectsUnavailableAndNonCompatibleResponses(t *testing.T) { + for _, tc := range []struct { + name, body string + status int + }{ + {"unauthorized", `{"error":{"message":"invalid key"}}`, 401}, + {"not found", "missing", 404}, + {"invalid JSON", "not-json", 200}, + {"wrong protocol", `{"models":[{"name":"model"}]}`, 200}, + {"empty", `{"data":[]}`, 200}, + {"blank id", `{"data":[{"id":" "}]}`, 200}, + {"oversized", strings.Repeat("x", (2<<20)+1), 200}, + } { + t.Run(tc.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(tc.status) + _, _ = w.Write([]byte(tc.body)) + })) + defer server.Close() + if _, err := ListOpenAIModels(context.Background(), AIConnection{BaseURL: server.URL, APIKey: "key"}, server.Client()); err == nil { + t.Fatal("invalid models response accepted") + } + }) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := ListOpenAIModels(ctx, AIConnection{BaseURL: "http://localhost:1/v1", APIKey: "key"}, nil); err == nil { + t.Fatal("canceled request accepted") + } +} + +func TestOpenAIConnectionValidation(t *testing.T) { + for _, base := range []string{"", "not-a-url", "ftp://example.com/v1", "https://example.com/v1?key=x", "https://example.com/v1#models", "https://user@example.com/v1", "https://example.com/v1/models", "https://example.com/v1/chat/completions"} { + if _, err := ListOpenAIModels(context.Background(), AIConnection{BaseURL: base, APIKey: "key"}, nil); !errors.Is(err, ErrInvalid) { + t.Fatalf("base=%q err=%v", base, err) + } + } + if _, err := ListOpenAIModels(context.Background(), AIConnection{BaseURL: "https://example.com/v1"}, nil); !errors.Is(err, ErrInvalid) { + t.Fatalf("missing key: %v", err) + } + if _, err := NewOpenAIClient("", "key", "model", nil); !errors.Is(err, ErrUnavailable) { + t.Fatalf("empty URL must not fall back: %v", err) + } + if _, err := NewOpenAIClient("https://example.com/v1", "key", "", nil); !errors.Is(err, ErrUnavailable) { + t.Fatalf("missing model: %v", err) + } + client, err := NewOpenAIClient(" https://example.com/custom/v1/ ", " key ", " model ", nil) + if err != nil || client.BaseURL != "https://example.com/custom/v1" || client.APIKey != "key" || client.Model != "model" || client.HTTPClient.Timeout == 0 { + t.Fatalf("client=%+v err=%v", client, err) + } +} + +func TestOpenAIClientGeneratesAndParsesDecisions(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != "/v1/chat/completions" || r.Header.Get("Authorization") != "Bearer test-key" { + t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path) + } + var request struct { + Model string `json:"model"` + Messages []struct{ Role, Content string } `json:"messages"` + } + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Error(err) + } + if request.Model != "model-test" || len(request.Messages) != 2 || request.Messages[0].Role != "system" || request.Messages[1].Role != "user" { + t.Errorf("request=%+v", request) + } + _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"match\":true,\"reason\":\"相关\"}"}}]}`)) + })) + defer server.Close() + client, err := NewOpenAIClient(server.URL+"/v1", "test-key", "model-test", server.Client()) + if err != nil { + t.Fatal(err) + } + if text, err := client.Generate(context.Background(), "要求", "输入"); err != nil || text == "" { + t.Fatalf("text=%q err=%v", text, err) + } + if match, reason, err := client.MatchTheme(context.Background(), "标题", "正文", "主题"); err != nil || !match || reason != "相关" { + t.Fatalf("theme=%v/%q err=%v", match, reason, err) + } + if match, reason, err := client.MatchLead(context.Background(), "作品", "评论", "要求"); err != nil || !match || reason != "相关" { + t.Fatalf("lead=%v/%q err=%v", match, reason, err) + } +} + +func TestOpenAIChatAndDecisionsRejectInvalidResponses(t *testing.T) { + for _, body := range []string{"bad-json", `{"choices":[]}`, `{"choices":[{"message":{"content":" "}}]}`, `{"choices":[{"message":{"content":"{}"}}]}`, `{"choices":[{"message":{"content":"bad-json"}}]}`} { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte(body)) })) + client, err := NewOpenAIClient(server.URL, "key", "model", server.Client()) + if err != nil { + t.Fatal(err) + } + if _, _, err := client.MatchTheme(context.Background(), "title", "body", "topic"); err == nil { + t.Fatal("invalid theme accepted") + } + if _, _, err := client.MatchLead(context.Background(), "work", "comment", "rule"); err == nil { + t.Fatal("invalid lead accepted") + } + server.Close() + } + var client *OpenAIClient + if _, err := client.Generate(context.Background(), "", ""); !errors.Is(err, ErrUnavailable) { + t.Fatalf("nil client: %v", err) + } +} + +func TestConfiguredOpenAIUsesSavedSettingsAndChangesImmediately(t *testing.T) { + store, _, ctx := openCreatorIntegrationStore(t) + var expectedKey, expectedModel string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "Bearer "+expectedKey { + t.Errorf("authorization does not use saved key") + } + var request map[string]any + _ = json.NewDecoder(r.Body).Decode(&request) + if request["model"] != expectedModel { + t.Errorf("model=%v want=%s", request["model"], expectedModel) + } + _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{\"match\":false,\"reason\":\"无关\"}"}}]}`)) + })) + defer server.Close() + client := &ConfiguredOpenAI{Store: store, HTTPClient: server.Client()} + if _, err := client.Generate(ctx, "", ""); !errors.Is(err, ErrUnavailable) { + t.Fatalf("unconfigured: %v", err) + } + settings, err := store.GetSettings(ctx) + if err != nil { + t.Fatal(err) + } + for _, model := range []string{"first", "second"} { + expectedKey, expectedModel = "key-"+model, model + input := SettingsUpdate{LookbackDays: settings.LookbackDays, NewWorkIntervalSeconds: settings.NewWorkIntervalSeconds, MetricInitialIntervalSeconds: settings.MetricInitialIntervalSeconds, MetricMultiplier: settings.MetricMultiplier, MetricMaxIntervalSeconds: settings.MetricMaxIntervalSeconds, MetricAgeSeconds: settings.MetricAgeSeconds, AIBaseURL: server.URL + "/v1", AIAPIKey: expectedKey, AIModel: expectedModel} + saved, err := store.UpdateSettings(ctx, input) + if err != nil { + t.Fatal(err) + } + if saved.AIBaseURL != input.AIBaseURL || saved.AIAPIKey != expectedKey || saved.AIModel != expectedModel { + t.Fatalf("settings=%+v", saved) + } + if _, err := client.Generate(ctx, "", ""); err != nil { + t.Fatal(err) + } + if match, reason, err := client.MatchLead(ctx, "", "", ""); err != nil || match || reason != "无关" { + t.Fatalf("lead=%v/%q err=%v", match, reason, err) + } + if match, reason, err := client.MatchTheme(ctx, "", "", ""); err != nil || match || reason != "无关" { + t.Fatalf("theme=%v/%q err=%v", match, reason, err) + } + } +} diff --git a/internal/creator/settings.go b/internal/creator/settings.go index 71a0cf9..7dc74c8 100644 --- a/internal/creator/settings.go +++ b/internal/creator/settings.go @@ -5,115 +5,141 @@ import ( "fmt" "strings" "time" + + "github.com/sirupsen/logrus" ) func scanSettings(scanner interface{ Scan(...any) error }) (Settings, error) { var result Settings if err := scanner.Scan(&result.LookbackDays, &result.NewWorkIntervalSeconds, &result.MetricInitialIntervalSeconds, &result.MetricMultiplier, &result.MetricMaxIntervalSeconds, &result.MetricAgeSeconds, - &result.AIProvider, &result.AIModel, &result.AIConfigured, &result.UpdatedAt); err != nil { + &result.AIBaseURL, &result.AIAPIKey, &result.AIModel, &result.UpdatedAt); err != nil { return Settings{}, err } result.UpdatedAt = result.UpdatedAt.UTC() return result, nil } -const settingsSelect = `SELECT lookback_days, new_work_interval_seconds, metric_initial_interval_seconds, - metric_multiplier, metric_max_interval_seconds, metric_age_seconds, ai_provider, ai_model, - ai_configured, updated_at - FROM creator_settings` +const settingsColumns = `lookback_days, new_work_interval_seconds, metric_initial_interval_seconds, + metric_multiplier, metric_max_interval_seconds, metric_age_seconds, ai_base_url, ai_api_key, ai_model, updated_at` +const settingsSelect = `SELECT ` + settingsColumns + ` FROM creator_settings` func (s *Store) GetSettings(ctx context.Context) (Settings, error) { result, err := scanSettings(s.db.QueryRowContext(ctx, settingsSelect)) return result, rowError(err) } -func (s *Store) UpdateSettings(ctx context.Context, input SettingsUpdate) (Settings, error) { - input.AIProvider = strings.TrimSpace(input.AIProvider) +func validateAISettings(input SettingsUpdate) (SettingsUpdate, error) { + input.AIBaseURL = strings.TrimSpace(input.AIBaseURL) + input.AIAPIKey = strings.TrimSpace(input.AIAPIKey) input.AIModel = strings.TrimSpace(input.AIModel) + if input.AIBaseURL == "" && input.AIAPIKey == "" && input.AIModel == "" { + return input, nil + } + connection, err := validateAIConnection(AIConnection{BaseURL: input.AIBaseURL, APIKey: input.AIAPIKey}) + if err != nil { + return SettingsUpdate{}, err + } + if input.AIModel == "" { + return SettingsUpdate{}, fmt.Errorf("%w: 请从模型列表选择 MODEL", ErrInvalid) + } + input.AIBaseURL = connection.BaseURL + return input, nil +} + +func (s *Store) UpdateSettings(ctx context.Context, input SettingsUpdate) (Settings, error) { if err := ValidateSettings(input); err != nil { return Settings{}, err } - if input.AIConfigured && (input.AIProvider == "" || input.AIModel == "") { - return Settings{}, ErrInvalid + input, err := validateAISettings(input) + if err != nil { + return Settings{}, err } tx, err := s.db.BeginTx(ctx, nil) if err != nil { return Settings{}, fmt.Errorf("begin creator settings update: %w", err) } defer tx.Rollback() - if _, err := tx.ExecContext(ctx, ` - UPDATE creator_settings SET lookback_days=$1, new_work_interval_seconds=$2, - metric_initial_interval_seconds=$3, metric_multiplier=$4, metric_max_interval_seconds=$5, - metric_age_seconds=$6, ai_provider=$7, ai_model=$8, ai_configured=$9, - updated_at=now()`, input.LookbackDays, input.NewWorkIntervalSeconds, + current, err := scanSettings(tx.QueryRowContext(ctx, settingsSelect+` FOR UPDATE`)) + if err != nil { + return Settings{}, rowError(err) + } + collectionChanged := current.LookbackDays != input.LookbackDays || current.NewWorkIntervalSeconds != input.NewWorkIntervalSeconds || + current.MetricInitialIntervalSeconds != input.MetricInitialIntervalSeconds || current.MetricMultiplier != input.MetricMultiplier || + current.MetricMaxIntervalSeconds != input.MetricMaxIntervalSeconds || current.MetricAgeSeconds != input.MetricAgeSeconds + value, err := scanSettings(tx.QueryRowContext(ctx, ` + UPDATE creator_settings SET lookback_days=$1, new_work_interval_seconds=$2, + metric_initial_interval_seconds=$3, metric_multiplier=$4, metric_max_interval_seconds=$5, + metric_age_seconds=$6, ai_base_url=$7, ai_api_key=$8, ai_model=$9, + updated_at=now() RETURNING `+settingsColumns, input.LookbackDays, input.NewWorkIntervalSeconds, input.MetricInitialIntervalSeconds, input.MetricMultiplier, input.MetricMaxIntervalSeconds, - input.MetricAgeSeconds, input.AIProvider, input.AIModel, input.AIConfigured); err != nil { - return Settings{}, databaseError(err) - } - var runningCheckpoints int - if err := tx.QueryRowContext(ctx, `SELECT count(*) FROM creator_collection_checkpoint WHERE status='running'`).Scan(&runningCheckpoints); err != nil { - return Settings{}, databaseError(err) - } - if runningCheckpoints > 0 { - return Settings{}, ErrConflict - } - rows, err := tx.QueryContext(ctx, ` - SELECT work_id, published_at FROM creator_work - WHERE published_at IS NOT NULL - FOR UPDATE`) + input.MetricAgeSeconds, input.AIBaseURL, input.AIAPIKey, input.AIModel)) if err != nil { return Settings{}, databaseError(err) } - type metricSchedule struct { - workID string - publishedAt *time.Time - } - schedules := make([]metricSchedule, 0) - for rows.Next() { - var item metricSchedule - if err := rows.Scan(&item.workID, &item.publishedAt); err != nil { + if collectionChanged { + var runningCheckpoints int + if err := tx.QueryRowContext(ctx, `SELECT count(*) FROM creator_collection_checkpoint WHERE status='running'`).Scan(&runningCheckpoints); err != nil { + return Settings{}, databaseError(err) + } + if runningCheckpoints > 0 { + return Settings{}, ErrConflict + } + rows, err := tx.QueryContext(ctx, `SELECT work_id, published_at FROM creator_work WHERE published_at IS NOT NULL FOR UPDATE`) + if err != nil { + return Settings{}, databaseError(err) + } + type metricSchedule struct { + workID string + publishedAt *time.Time + } + schedules := make([]metricSchedule, 0) + for rows.Next() { + var item metricSchedule + if err := rows.Scan(&item.workID, &item.publishedAt); err != nil { + _ = rows.Close() + return Settings{}, err + } + schedules = append(schedules, item) + } + if err := rows.Err(); err != nil { _ = rows.Close() return Settings{}, err } - schedules = append(schedules, item) - } - if err := rows.Err(); err != nil { - _ = rows.Close() - return Settings{}, err - } - if err := rows.Close(); err != nil { - return Settings{}, err - } - now := time.Now().UTC() - start, end, err := NewCollectionWindow(now, input.LookbackDays) - if err != nil { - return Settings{}, err - } - if _, err := tx.ExecContext(ctx, ` - UPDATE creator_collection_checkpoint - SET window_start=$1, window_end=$2, cursor='', lease_token='', lease_until=NULL, - status=CASE WHEN status='blocked' THEN 'blocked' ELSE 'idle' END, - last_error=CASE WHEN status='blocked' THEN last_error ELSE '' END - WHERE status <> 'running'`, start, end); err != nil { - return Settings{}, databaseError(err) - } - for _, schedule := range schedules { - nextAt, reason := NextMetricAtValue(schedule.publishedAt, now, input) - stopped := nextAt.IsZero() + if err := rows.Close(); err != nil { + return Settings{}, err + } + now := time.Now().UTC() + start, end, err := NewCollectionWindow(now, input.LookbackDays) + if err != nil { + return Settings{}, err + } if _, err := tx.ExecContext(ctx, ` - UPDATE creator_work SET next_metric_at=$2, metric_stop_reason=$3, - metric_plan_next_at=$2, metric_plan_interval_seconds=$4, metric_plan_multiplier=$5, - metric_plan_max_interval_seconds=$6, metric_plan_stopped=$7, updated_at=now() - WHERE work_id=$1`, schedule.workID, nullableArg(nextAt), reason, input.MetricInitialIntervalSeconds, - input.MetricMultiplier, input.MetricMaxIntervalSeconds, stopped); err != nil { + UPDATE creator_collection_checkpoint + SET window_start=$1, window_end=$2, cursor='', lease_token='', lease_until=NULL, + status=CASE WHEN status='blocked' THEN 'blocked' ELSE 'idle' END, + last_error=CASE WHEN status='blocked' THEN last_error ELSE '' END + WHERE status <> 'running'`, start, end); err != nil { return Settings{}, databaseError(err) } + for _, schedule := range schedules { + nextAt, reason := NextMetricAtValue(schedule.publishedAt, now, input) + stopped := nextAt.IsZero() + if _, err := tx.ExecContext(ctx, ` + UPDATE creator_work SET next_metric_at=$2, metric_stop_reason=$3, + metric_plan_next_at=$2, metric_plan_interval_seconds=$4, metric_plan_multiplier=$5, + metric_plan_max_interval_seconds=$6, metric_plan_stopped=$7, updated_at=now() + WHERE work_id=$1`, schedule.workID, nullableArg(nextAt), reason, input.MetricInitialIntervalSeconds, + input.MetricMultiplier, input.MetricMaxIntervalSeconds, stopped); err != nil { + return Settings{}, databaseError(err) + } + } } if err := tx.Commit(); err != nil { return Settings{}, fmt.Errorf("commit creator settings update: %w", err) } - return s.GetSettings(ctx) + logrus.WithFields(logrus.Fields{"collection_changed": collectionChanged, "ai_base_url": value.AIBaseURL, "ai_model": value.AIModel}).Info("Creator settings saved") + return value, nil } func NextMetricAtValue(publishedAt *time.Time, now time.Time, input SettingsUpdate) (time.Time, string) { diff --git a/internal/environment/migrations/1057_openai_settings.sql b/internal/environment/migrations/1057_openai_settings.sql new file mode 100644 index 0000000..b618d49 --- /dev/null +++ b/internal/environment/migrations/1057_openai_settings.sql @@ -0,0 +1,5 @@ +ALTER TABLE creator_settings + ADD COLUMN ai_base_url TEXT NOT NULL DEFAULT '', + ADD COLUMN ai_api_key TEXT NOT NULL DEFAULT '', + DROP COLUMN ai_provider, + DROP COLUMN ai_configured; diff --git a/internal/environment/openai_schema_upgrade_test.go b/internal/environment/openai_schema_upgrade_test.go new file mode 100644 index 0000000..8fa18dd --- /dev/null +++ b/internal/environment/openai_schema_upgrade_test.go @@ -0,0 +1,86 @@ +package environment + +import ( + "context" + "database/sql" + "fmt" + "net/url" + "os" + "testing" + "time" +) + +func TestOpenAISettingsSchemaFreshUpgradeAndRestart(t *testing.T) { + databaseURL := os.Getenv("CREATORHUB_POSTGRES_TEST_URL") + if databaseURL == "" { + t.Skip("set CREATORHUB_POSTGRES_TEST_URL to run migration tests") + } + ctx := context.Background() + admin, err := sql.Open("pgx", databaseURL) + if err != nil { + t.Fatal(err) + } + defer admin.Close() + schema := fmt.Sprintf("openai_upgrade_%d", time.Now().UnixNano()) + if _, err := admin.ExecContext(ctx, "CREATE SCHEMA "+schema); err != nil { + t.Fatal(err) + } + defer func() { + if _, err := admin.ExecContext(ctx, "DROP SCHEMA "+schema+" CASCADE"); err != nil { + t.Errorf("cleanup: %v", err) + } + }() + parsed, err := url.Parse(databaseURL) + if err != nil { + t.Fatal(err) + } + query := parsed.Query() + query.Set("search_path", schema) + parsed.RawQuery = query.Encode() + store, err := Open(ctx, parsed.String()) + if err != nil { + t.Fatal(err) + } + defer store.Close() + assertColumns := func() { + t.Helper() + var columns int + if err := store.db.QueryRowContext(ctx, `SELECT count(*) FROM information_schema.columns WHERE table_schema=current_schema() AND table_name='creator_settings' AND column_name IN ('ai_base_url','ai_api_key','ai_model')`).Scan(&columns); err != nil || columns != 3 { + t.Fatalf("AI columns=%d err=%v", columns, err) + } + if err := store.db.QueryRowContext(ctx, `SELECT count(*) FROM information_schema.columns WHERE table_schema=current_schema() AND table_name='creator_settings' AND column_name IN ('ai_provider','ai_configured')`).Scan(&columns); err != nil || columns != 0 { + t.Fatalf("legacy columns=%d err=%v", columns, err) + } + } + assertColumns() + // Recreate the previous schema only in this isolated database, preserving its settings row. + if _, err := store.db.ExecContext(ctx, `ALTER TABLE creator_settings DROP COLUMN ai_base_url,DROP COLUMN ai_api_key,ADD COLUMN ai_provider text NOT NULL DEFAULT '',ADD COLUMN ai_configured boolean NOT NULL DEFAULT false; + UPDATE creator_settings SET lookback_days=12,ai_model='legacy-model',ai_provider='bailian',ai_configured=true; + DELETE FROM schema_migration WHERE version=1057`); err != nil { + t.Fatal(err) + } + if err := store.migrate(ctx); err != nil { + t.Fatalf("upgrade: %v", err) + } + assertColumns() + var days int + var model, base, key string + if err := store.db.QueryRowContext(ctx, `SELECT lookback_days,ai_model,ai_base_url,ai_api_key FROM creator_settings`).Scan(&days, &model, &base, &key); err != nil || days != 12 || model != "legacy-model" || base != "" || key != "" { + t.Fatalf("preserved settings=%d/%s/%s/%s err=%v", days, model, base, key, err) + } + if _, err := store.db.ExecContext(ctx, `UPDATE creator_settings SET ai_base_url='https://example.com/v1',ai_api_key='saved-key',ai_model='chosen-model'`); err != nil { + t.Fatal(err) + } + for i := 0; i < 2; i++ { + if err := store.migrate(ctx); err != nil { + t.Fatalf("repeat startup: %v", err) + } + } + if err := store.db.QueryRowContext(ctx, `SELECT lookback_days,ai_model,ai_base_url,ai_api_key FROM creator_settings`).Scan(&days, &model, &base, &key); err != nil || days != 12 || model != "chosen-model" || base != "https://example.com/v1" || key != "saved-key" { + t.Fatalf("restart changed settings: %d/%s/%s/%s err=%v", days, model, base, key, err) + } + var applied int + if err := store.db.QueryRowContext(ctx, `SELECT count(*) FROM schema_migration WHERE version=1057`).Scan(&applied); err != nil || applied != 1 { + t.Fatalf("applied=%d err=%v", applied, err) + } +} diff --git a/internal/environment/store.go b/internal/environment/store.go index 4b30711..ebd4a47 100644 --- a/internal/environment/store.go +++ b/internal/environment/store.go @@ -209,6 +209,9 @@ var migration1055 string //go:embed migrations/1056_event_details.sql var migration1056 string +//go:embed migrations/1057_openai_settings.sql +var migration1057 string + var ( ErrConflict = errors.New("resource conflicts with existing state") ErrInvalid = errors.New("invalid environment input") @@ -351,7 +354,7 @@ func (s *Store) migrate(ctx context.Context) error { {1029, migration1029}, {1030, migration1030}, {1031, migration1031}, {1032, migration1032}, {1033, migration1033}, {1034, migration1034}, {1035, migration1035}, {1036, migration1036}, {1037, migration1037}, {1038, migration1038}, {1039, migration1039}, {1040, migration1040}, {1041, migration1041}, {1042, migration1042}, {1043, migration1043}, {1044, migration1044}, {1045, migration1045}, {1046, migration1046}, - {43, migration043}, {44, migration044}, {1047, migration1047}, {1048, migration1048}, {1049, migration1049}, {1050, migration1050}, {1051, migration1051}, {1052, migration1052}, {1053, migration1053}, {1054, migration1054}, {1055, migration1055}, {1056, migration1056}} { + {43, migration043}, {44, migration044}, {1047, migration1047}, {1048, migration1048}, {1049, migration1049}, {1050, migration1050}, {1051, migration1051}, {1052, migration1052}, {1053, migration1053}, {1054, migration1054}, {1055, migration1055}, {1056, migration1056}, {1057, migration1057}} { var applied bool if err := tx.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM schema_migration WHERE version = $1)`, migration.version).Scan(&applied); err != nil { return errors.New("read environment schema migration state") diff --git a/web/src/pages/creator/settings/index.tsx b/web/src/pages/creator/settings/index.tsx index 30ea84f..213eacb 100644 --- a/web/src/pages/creator/settings/index.tsx +++ b/web/src/pages/creator/settings/index.tsx @@ -1,8 +1,7 @@ -// 采集设置:语义对齐 web.archived CreatorSettingsPage.jsx(提交 d97cade)。 -// 四类配置用 antd Tabs 组织;仅用 antd 默认组件。 import { useCallback, useEffect, useState } from 'react'; -import { Alert, App, Button, Card, Checkbox, Form, Input, InputNumber, Tabs, Typography } from 'antd'; -import { creatorGet, creatorUpdate } from '@/services/api'; +import { ReloadOutlined } from '@ant-design/icons'; +import { Alert, App, Button, Card, Flex, Form, Input, InputNumber, Select, Tabs, Typography } from 'antd'; +import { creatorAction, creatorGet, creatorUpdate } from '@/services/api'; import { conflictMessage } from '@/utils/helpers'; const numericFields = new Set([ @@ -13,26 +12,35 @@ const numericFields = new Set([ 'metric_max_interval_seconds', 'metric_age_seconds', ]); - const editableFields = [ - 'lookback_days', - 'new_work_interval_seconds', - 'metric_initial_interval_seconds', - 'metric_multiplier', - 'metric_max_interval_seconds', - 'metric_age_seconds', - 'ai_provider', + ...numericFields, + 'ai_base_url', + 'ai_api_key', 'ai_model', - 'ai_configured', ]; - const editableSettings = (value: any) => Object.fromEntries(editableFields.map((field) => [field, value[field]])); +type ModelList = { + baseURL: string; + apiKey: string; + options: string[]; + loading: boolean; + error: string | null; +}; + export default function Page() { const [form] = Form.useForm(); const [pending, setPending] = useState(true); const [error, setError] = useState(null); const [busy, setBusy] = useState(false); + const [refreshToken, setRefreshToken] = useState(0); + const [models, setModels] = useState({ baseURL: '', apiKey: '', options: [], loading: false, error: null }); + const baseURL = (Form.useWatch('ai_base_url', { form, preserve: true }) ?? '').trim().replace(/\/+$/, ''); + const apiKey = (Form.useWatch('ai_api_key', { form, preserve: true }) ?? '').trim(); + const connectionMatches = models.baseURL === baseURL && models.apiKey === apiKey; + const modelOptions = connectionMatches ? models.options : []; + const loadingModels = Boolean(baseURL && apiKey) && (!connectionMatches || models.loading); + const modelError = connectionMatches ? models.error : null; const { message: messageApi } = App.useApp(); const load = useCallback(async () => { @@ -52,11 +60,56 @@ export default function Page() { load(); }, [load]); + useEffect(() => { + if (pending) return; + if (!baseURL || !apiKey) { + setModels({ baseURL, apiKey, options: [], loading: false, error: null }); + return; + } + let active = true; + setModels({ baseURL, apiKey, options: [], loading: true, error: null }); + const timer = setTimeout(async () => { + try { + const result = await creatorAction('/creator/settings/models', { base_url: baseURL, api_key: apiKey }); + if (!active) return; + if (!Array.isArray(result.models) || result.models.length === 0) { + form.setFieldValue('ai_model', undefined); + throw new Error('服务未返回可选模型,请检查 BASEURL 和 APIKEY。'); + } + const selected = form.getFieldValue('ai_model'); + const missingSelected = selected && !result.models.includes(selected); + if (missingSelected) form.setFieldValue('ai_model', undefined); + setModels({ + baseURL, apiKey, options: result.models, loading: false, + error: missingSelected ? '已保存的 MODEL 不在当前模型列表中,请重新选择。' : null, + }); + } catch (modelLoadError) { + if (!active) return; + setModels({ baseURL, apiKey, options: [], loading: false, error: conflictMessage(modelLoadError, '模型列表获取失败') }); + } + }, 500); + return () => { active = false; clearTimeout(timer); }; + }, [baseURL, apiKey, pending, refreshToken, form]); + + function connectionChanged(changed: Record) { + if ('ai_base_url' in changed || 'ai_api_key' in changed) { + form.setFieldValue('ai_model', undefined); + setRefreshToken((value) => value + 1); + } + if ('ai_model' in changed && modelOptions.includes(String(changed.ai_model))) { + setModels((value) => ({ ...value, error: null })); + } + } + async function save(values: any) { + const hasAI = [values.ai_base_url, values.ai_api_key, values.ai_model].some((value) => String(value ?? '').trim()); + if (hasAI && (!connectionMatches || loadingModels || !baseURL || !apiKey || !modelOptions.includes(values.ai_model))) { + messageApi.error('请填写 BASEURL 和 APIKEY,并从成功获取的模型列表中选择 MODEL。'); + return; + } setBusy(true); try { - // 数值字段清空时归零,对齐归档版 Number('') 的提交语义;服务侧字段为必填数值。 - const payload = Object.fromEntries(editableFields.map((field) => [field, numericFields.has(field) ? values[field] ?? 0 : values[field]])); + const payload = Object.fromEntries(editableFields.map((field) => [field, numericFields.has(field) ? values[field] ?? 0 : String(values[field] ?? '').trim()])); const result = await creatorUpdate('/creator/settings', payload); form.setFieldsValue(editableSettings(result)); messageApi.success('采集与 AI 配置已保存。'); @@ -72,14 +125,13 @@ export default function Page() { return (
- 按配置类别管理 UTC 采集窗口、指标采集和已批准的服务。 + 管理采集策略、指标采集和 AI 服务。仅修改 AI 服务不会改变采集进度。 -
+ @@ -89,10 +141,10 @@ export default function Page() {
- ) }, + ), + }, { - key: 'metrics', - label: '指标采集', + key: 'metrics', label: '指标采集', forceRender: true, children: (
@@ -108,30 +160,40 @@ export default function Page() {
- ) }, + ), + }, { - key: 'ai', - label: 'AI 服务', + key: 'ai', label: 'AI 服务', forceRender: true, children: (
- - + 仅支持 OpenAI Compatible 格式;模型列表与实际调用使用同一地址和密钥。 + + - - + + - - 已完成 AI 配置审批 + MODEL} + > +