feat: configure OpenAI-compatible AI services and discover models
douyin-release-gate / verify (push) Failing after 19m29s

This commit is contained in:
2026-10-07 17:17:33 +08:00
parent 51091cf0c5
commit f345e368df
19 changed files with 1135 additions and 408 deletions
+2
View File
@@ -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 默认组件原样实现,禁止自定义封装与样式魔改;组件不满足业务时改交互逻辑适配组件;
-2
View File
@@ -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:-}
+159
View File
@@ -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)
}
}
+18
View File
@@ -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 {
-11
View File
@@ -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
}
+115
View File
@@ -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)
}
}
-186
View File
@@ -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)
}
-95
View File
@@ -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")
}
}
+2 -2
View File
@@ -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)
}
+2 -2
View File
@@ -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 {
+4 -4
View File
@@ -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 {
+247
View File
@@ -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)
}
+177
View File
@@ -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)
}
}
}
+93 -67
View File
@@ -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) {
@@ -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;
@@ -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)
}
}
+4 -1
View File
@@ -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")
+100 -38
View File
@@ -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<any>(null);
const [busy, setBusy] = useState(false);
const [refreshToken, setRefreshToken] = useState(0);
const [models, setModels] = useState<ModelList>({ 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<string, unknown>) {
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 (
<div>
<Typography.Paragraph type="secondary" style={{ marginBottom: 16 }}>按配置类别管理 UTC 采集窗口、指标采集和已批准的服务。</Typography.Paragraph>
<Typography.Paragraph type="secondary" style={{ marginBottom: 16 }}>管理采集策略、指标采集和 AI 服务。仅修改 AI 服务不会改变采集进度。</Typography.Paragraph>
<Card variant="outlined">
<Form form={form} layout="vertical" onFinish={save} initialValues={{}}>
<Form form={form} layout="vertical" onFinish={save} onValuesChange={connectionChanged} disabled={busy}>
<Tabs
items={[
{
key: 'collection',
label: '采集策略',
key: 'collection', label: '采集策略', forceRender: true,
children: (
<div style={{ maxWidth: 640 }}>
<Form.Item name="lookback_days" label="竞品回溯天数" extra="首次回溯默认 30 天;不会猜测缺失发布时间。">
@@ -89,10 +141,10 @@ export default function Page() {
<InputNumber min={1} style={{ width: 200 }} />
</Form.Item>
</div>
) },
),
},
{
key: 'metrics',
label: '指标采集',
key: 'metrics', label: '指标采集', forceRender: true,
children: (
<div style={{ maxWidth: 640 }}>
<Form.Item name="metric_initial_interval_seconds" label="指标首次间隔(秒)">
@@ -108,30 +160,40 @@ export default function Page() {
<InputNumber min={1} style={{ width: 200 }} />
</Form.Item>
</div>
) },
),
},
{
key: 'ai',
label: 'AI 服务',
key: 'ai', label: 'AI 服务', forceRender: true,
children: (
<div style={{ maxWidth: 640 }}>
<Form.Item name="ai_provider" label="AI Provider">
<Input />
<Typography.Paragraph type="secondary">仅支持 OpenAI Compatible 格式;模型列表与实际调用使用同一地址和密钥。</Typography.Paragraph>
<Form.Item name="ai_base_url" label="BASEURL" extra="填写接口根地址,通常以 /v1 结尾,不包含 /models 或 /chat/completions。">
<Input placeholder="https://api.openai.com/v1" autoComplete="off" />
</Form.Item>
<Form.Item name="ai_model" label="AI Model">
<Input />
<Form.Item name="ai_api_key" label="APIKEY">
<Input.Password autoComplete="off" />
</Form.Item>
<Form.Item name="ai_configured" valuePropName="checked">
<Checkbox>已完成 AI 配置审批</Checkbox>
<Form.Item
name="ai_model"
label={<Flex align="center" gap="small">MODEL<Button type="link" size="small" icon={<ReloadOutlined />} disabled={!baseURL || !apiKey || loadingModels || busy} onClick={() => setRefreshToken((value) => value + 1)}>刷新模型</Button></Flex>}
>
<Select
showSearch={{ optionFilterProp: 'label' }}
options={modelOptions.map((model) => ({ label: model, value: model }))}
loading={loadingModels}
disabled={busy || loadingModels || modelOptions.length === 0}
placeholder={loadingModels ? '正在获取模型…' : '填写 BASEURL 和 APIKEY 后选择模型'}
/>
</Form.Item>
{modelError && <Alert type="error" showIcon title={modelError} />}
</div>
) },
),
},
]}
/>
<div style={{ display: 'flex', justifyContent: 'flex-end' }}>
<Button type="primary" htmlType="submit" loading={busy}>
{busy ? '保存中…' : '保存设置'}
</Button>
</div>
<Flex justify="flex-end">
<Button type="primary" htmlType="submit" loading={busy}>{busy ? '保存中…' : '保存设置'}</Button>
</Flex>
</Form>
</Card>
</div>
+121
View File
@@ -0,0 +1,121 @@
const test = require('node:test');
const assert = require('node:assert/strict');
const {readFileSync} = require('node:fs');
const {resolve} = require('node:path');
const Module = require('node:module');
const ts = require('typescript');
const React = require('react');
const defaults = {lookback_days:30, new_work_interval_seconds:1800, metric_initial_interval_seconds:3600, metric_multiplier:2, metric_max_interval_seconds:86400, metric_age_seconds:2592000, ai_base_url:'https://example.com/v1', ai_api_key:'saved-key', ai_model:'model-a', updated_at:'server-only'};
function deferred() { let resolve, reject; const promise = new Promise((yes, no) => {resolve=yes; reject=no;}); return {promise, resolve, reject}; }
function harness(options={}) {
const states=[], effects=[], queued=[], timers=new Map(), calls=[], saves=[], notices=[], logs=[];
let cursor=0, timerID=0;
const values={};
const form={setFieldsValue(value){Object.assign(values,value);}, getFieldValue(name){return values[name];}, setFieldValue(name,value){values[name]=value;}};
const antd={};
for(const name of ['Alert','Button','Card','Flex','Input','InputNumber','Select','Tabs']) antd[name]=()=>null;
antd.Input.Password=()=>null;
antd.Typography={Paragraph:()=>null,Text:()=>null};
antd.Form=()=>null; antd.Form.Item=()=>null;
antd.Form.useForm=()=>[form];
antd.Form.useWatch=(name)=>values[name];
antd.App={useApp:()=>({message:{success:v=>notices.push(v),error:v=>notices.push(v)}})};
const file=resolve(__dirname,'../src/pages/creator/settings/index.tsx');
const loaded=new Module(file,module); loaded.filename=file;
loaded.require=id=>{
if(id==='@test/timers') return {setTimeout(fn){const id=++timerID; timers.set(id,fn); return id;},clearTimeout(id){timers.delete(id);}};
if(id==='react') return {
useState(initial){const i=cursor++; if(!(i in states)) states[i]=typeof initial==='function'?initial():initial; return [states[i],value=>{states[i]=typeof value==='function'?value(states[i]):value;}];},
useRef(initial){const i=cursor++; if(!(i in states)) states[i]={current:initial}; return states[i];},
useCallback(callback,deps){const i=cursor++; const prior=states[i]; if(!prior||deps.some((v,j)=>!Object.is(v,prior.deps[j]))) states[i]={callback,deps}; return states[i].callback;},
useEffect(callback,deps){const i=cursor++; const prior=effects[i]; if(!prior||deps.some((v,j)=>!Object.is(v,prior.deps[j]))) {prior?.cleanup?.(); effects[i]={deps}; queued.push(()=>{effects[i].cleanup=callback();});}},
};
if(id==='antd') return antd;
if(id==='@ant-design/icons') return {ReloadOutlined:()=>null};
if(id==='@/utils/helpers') return {conflictMessage:(e,f)=>e?.message||f};
if(id==='@/services/api') return {
creatorGet:async()=>{if(options.loadError) throw new Error(options.loadError); return options.settings||defaults;},
creatorAction:async(path,payload)=>{calls.push({path,payload}); return options.models ? options.models(payload,calls.length) : {models:['model-a','model-b']};},
creatorUpdate:async(path,payload)=>{saves.push({path,payload}); if(options.saveError) throw new Error(options.saveError); return payload;},
};
return require(id);
};
const code=ts.transpileModule(readFileSync(file,'utf8'),{compilerOptions:{module:ts.ModuleKind.CommonJS,jsx:ts.JsxEmit.ReactJSX,target:ts.ScriptTarget.ES2020}}).outputText;
loaded._compile(`const {setTimeout,clearTimeout}=require('@test/timers');\n${code}`,file);
const h={antd,form,values,calls,saves,notices,logs,
render(){cursor=0; return loaded.exports.default();},
async flush(){for(let i=0;i<5;i++){h.render(); while(queued.length) queued.shift()(); for(let j=0;j<12;j++) await Promise.resolve();} return h.render();},
async tick(){const pending=[...timers.values()]; timers.clear(); pending.forEach(fn=>fn()); return h.flush();},
async change(changed){Object.assign(values,changed); const tree=h.render(); find(tree,antd.Form)?.props.onValuesChange?.(changed,values); return h.flush();},
dispose(){for(const effect of effects) effect?.cleanup?.();},
};
return h;
}
function all(node,predicate){const result=[]; function visit(value){if(Array.isArray(value)) return value.forEach(visit); if(!React.isValidElement(value)) return; if(predicate(value)) result.push(value); visit(value.props.children); visit(value.props.label); if(Array.isArray(value.props.items)) value.props.items.forEach(item=>visit(item.children));} visit(node); return result;}
function text(node){if(node==null) return ''; if(typeof node==='string'||typeof node==='number') return String(node); if(Array.isArray(node)) return node.map(text).join(''); return React.isValidElement(node)?text(node.props.children):'';}
const find=(tree,type)=>all(tree,n=>n.type===type)[0];
const alerts=(h,tree)=>all(tree,n=>n.type===h.antd.Alert).map(n=>`${text(n.props.title)} ${text(n.props.description)}`).join('\n');
const refresh=(h,tree)=>all(tree,n=>n.type===h.antd.Button).find(n=>text(n.props.children)==='刷新模型');
test('AI service uses password input, searchable discovered models and persists all settings',async t=>{
const h=harness(); t.after(h.dispose); await h.flush(); let tree=await h.tick();
assert.deepEqual(h.calls,[{path:'/creator/settings/models',payload:{base_url:defaults.ai_base_url,api_key:defaults.ai_api_key}}]);
assert.ok(find(tree,h.antd.Input.Password));
const select=find(tree,h.antd.Select); assert.ok(select.props.showSearch); assert.equal(select.props.mode,undefined);
assert.deepEqual(select.props.options.map(o=>o.value),['model-a','model-b']);
const fields=all(tree,n=>n.type===h.antd.Form.Item).map(n=>n.props.name);
assert.ok(fields.includes('ai_base_url')); assert.ok(fields.includes('ai_api_key')); assert.ok(fields.includes('ai_model'));
assert.ok(!fields.includes('ai_provider')); assert.ok(!fields.includes('ai_configured'));
assert.ok(find(tree,h.antd.Tabs).props.items.every(item=>item.forceRender));
await find(tree,h.antd.Form).props.onFinish({...h.values}); await h.flush();
assert.equal(h.saves[0].payload.ai_api_key,'saved-key'); assert.equal(h.saves[0].payload.ai_model,'model-a');
assert.equal(h.saves[0].payload.lookback_days,30); assert.equal(h.saves[0].payload.updated_at,undefined);
assert.equal(h.notices[0],'采集与 AI 配置已保存。');
});
test('connection edits clear selection and stale replies cannot override new models',async t=>{
const old=deferred(); const h=harness({models:(_,n)=>n===1?old.promise:{models:['new-model']}}); t.after(h.dispose);
await h.flush(); await h.tick();
await h.change({ai_base_url:'https://other.example/v1',ai_api_key:'new-key'});
assert.equal(h.values.ai_model,undefined);
let tree=await h.tick(); assert.deepEqual(find(tree,h.antd.Select).props.options.map(o=>o.value),['new-model']);
old.resolve({models:['stale-model']}); tree=await h.flush();
assert.deepEqual(find(tree,h.antd.Select).props.options.map(o=>o.value),['new-model']);
assert.equal(h.calls[1].payload.api_key,'new-key');
});
test('failed discovery is explicit and refresh recovers without a manual model fallback',async t=>{
const h=harness({models:(_,n)=>n===1?Promise.reject(new Error('HTTP 401: invalid key')):{models:['model-a']}}); t.after(h.dispose);
await h.flush(); let tree=await h.tick(); assert.match(alerts(h,tree),/401/);
assert.equal(find(tree,h.antd.Select).props.disabled,true); assert.deepEqual(find(tree,h.antd.Select).props.options,[]);
refresh(h,tree).props.onClick(); await h.flush(); tree=await h.tick();
assert.equal(alerts(h,tree),''); assert.equal(find(tree,h.antd.Select).props.disabled,false);
});
test('empty discovery and unavailable saved model are explicit',async t=>{
for(const models of [[],['different-model']]){
const h=harness({models:()=>({models})}); t.after(h.dispose); await h.flush(); const tree=await h.tick();
assert.ok(alerts(h,tree)); assert.equal(h.values.ai_model,undefined);
}
});
test('unconfigured AI does not fetch and collection settings remain saveable',async t=>{
const h=harness({settings:{...defaults,ai_base_url:'',ai_api_key:'',ai_model:''}}); t.after(h.dispose);
await h.flush(); const tree=await h.tick(); assert.equal(h.calls.length,0);
assert.equal(find(tree,h.antd.Select).props.disabled,true);
await find(tree,h.antd.Form).props.onFinish({...h.values}); assert.equal(h.saves.length,1);
});
test('partial connection or a model not from discovery cannot be saved',async t=>{
const h=harness(); t.after(h.dispose); await h.flush(); await h.tick();
await h.change({ai_api_key:'changed'}); let tree=h.render();
await find(tree,h.antd.Form).props.onFinish({...h.values,ai_model:'model-a'}); assert.equal(h.saves.length,0);
await h.tick(); tree=h.render(); await find(tree,h.antd.Form).props.onFinish({...h.values,ai_model:'invented-model'}); assert.equal(h.saves.length,0);
});
test('load and save failures remain visible',async t=>{
const broken=harness({loadError:'settings unavailable'}); t.after(broken.dispose);
assert.match(alerts(broken,await broken.flush()),/settings unavailable/);
const h=harness({saveError:'database unavailable'}); t.after(h.dispose); await h.flush(); const tree=await h.tick();
await find(tree,h.antd.Form).props.onFinish({...h.values}); assert.match(h.notices.join(' '),/database unavailable/);
});