feat: configure OpenAI-compatible AI services and discover models
douyin-release-gate / verify (push) Failing after 19m29s
douyin-release-gate / verify (push) Failing after 19m29s
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user