178 lines
7.8 KiB
Go
178 lines
7.8 KiB
Go
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)
|
|
}
|
|
}
|
|
}
|