229 lines
8.5 KiB
Go
229 lines
8.5 KiB
Go
package aimatching
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"go-admin/app/goauto/models"
|
|
|
|
"gorm.io/driver/sqlite"
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/logger"
|
|
)
|
|
|
|
func openSuggestTestDB(t *testing.T) *gorm.DB {
|
|
t.Helper()
|
|
db, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared&_foreign_keys=on", t.Name())), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
|
if err != nil {
|
|
t.Fatalf("open database: %v", err)
|
|
}
|
|
if err := db.AutoMigrate(&models.AIMatchingSetting{}); err != nil {
|
|
t.Fatalf("migrate: %v", err)
|
|
}
|
|
return db
|
|
}
|
|
|
|
func seedEnabledSetting(t *testing.T, db *gorm.DB, baseURL string) {
|
|
t.Helper()
|
|
setting := models.AIMatchingSetting{ID: 1, Enabled: true, Provider: ProviderOpenAICompatible, BaseURL: baseURL, Model: "test-model", APIKey: "test-key", TimeoutSeconds: 5, AutoConfirmMinConfidence: 0.9}
|
|
if err := db.Create(&setting).Error; err != nil {
|
|
t.Fatalf("seed setting: %v", err)
|
|
}
|
|
}
|
|
|
|
func chatCompletionResponder(content string) http.HandlerFunc {
|
|
return func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
|
"choices": []map[string]any{{"message": map[string]any{"content": content}}},
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestSuggestBatchOnlyAcceptsSuppliedCandidateIDs(t *testing.T) {
|
|
server := httptest.NewServer(chatCompletionResponder(`{"suggestions":[
|
|
{"sourceId":"s1","candidateId":"c2","confidence":0.95,"reason":"exact"},
|
|
{"sourceId":"s2","candidateId":"c9","confidence":0.99,"reason":"out of set"},
|
|
{"sourceId":"s3","candidateId":"","confidence":0.1,"reason":"no match"}
|
|
]}`))
|
|
defer server.Close()
|
|
|
|
db := openSuggestTestDB(t)
|
|
seedEnabledSetting(t, db, server.URL)
|
|
service := NewService(db)
|
|
|
|
result, err := service.SuggestBatch(context.Background(), SuggestRequest{
|
|
Dimension: "color",
|
|
Sources: []SuggestSource{{ID: "s1", Label: "黑色"}, {ID: "s2", Label: "白色"}, {ID: "s3", Label: "红色"}},
|
|
Candidates: []SuggestCandidate{
|
|
{ID: "c1", Label: "藏青色"}, {ID: "c2", Label: "黑色"},
|
|
},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if got := result.Decisions["s1"]; got.CandidateID != "c2" || got.Confidence != 0.95 {
|
|
t.Fatalf("s1 decision wrong: %+v", got)
|
|
}
|
|
if got := result.Decisions["s2"]; got.CandidateID != "" {
|
|
t.Fatalf("s2 must be treated as no-suggestion for out-of-set candidate id, got %+v", got)
|
|
}
|
|
if got := result.Decisions["s3"]; got.CandidateID != "" {
|
|
t.Fatalf("s3 decision wrong: %+v", got)
|
|
}
|
|
}
|
|
|
|
func TestSuggestBatchIgnoresUnknownAndDuplicateSourceIDs(t *testing.T) {
|
|
server := httptest.NewServer(chatCompletionResponder(`{"suggestions":[
|
|
{"sourceId":"s1","candidateId":"c1","confidence":0.5,"reason":"first"},
|
|
{"sourceId":"s1","candidateId":"c1","confidence":0.99,"reason":"second, must be ignored"},
|
|
{"sourceId":"unknown","candidateId":"c1","confidence":0.9,"reason":"must be dropped"}
|
|
]}`))
|
|
defer server.Close()
|
|
|
|
db := openSuggestTestDB(t)
|
|
seedEnabledSetting(t, db, server.URL)
|
|
service := NewService(db)
|
|
|
|
result, err := service.SuggestBatch(context.Background(), SuggestRequest{
|
|
Sources: []SuggestSource{{ID: "s1", Label: "黑色"}},
|
|
Candidates: []SuggestCandidate{{ID: "c1", Label: "黑色"}},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if len(result.Decisions) != 1 {
|
|
t.Fatalf("expected exactly one decision, got %+v", result.Decisions)
|
|
}
|
|
if got := result.Decisions["s1"]; got.Confidence != 0.5 {
|
|
t.Fatalf("expected first duplicate to win, got %+v", got)
|
|
}
|
|
}
|
|
|
|
func TestSuggestBatchRejectsOversizedRequest(t *testing.T) {
|
|
db := openSuggestTestDB(t)
|
|
seedEnabledSetting(t, db, "http://unused.invalid")
|
|
service := NewService(db)
|
|
|
|
sources := make([]SuggestSource, maxSuggestSources+1)
|
|
for i := range sources {
|
|
sources[i] = SuggestSource{ID: fmt.Sprintf("s%d", i), Label: "x"}
|
|
}
|
|
_, err := service.SuggestBatch(context.Background(), SuggestRequest{Sources: sources, Candidates: []SuggestCandidate{{ID: "c1", Label: "y"}}})
|
|
if err == nil {
|
|
t.Fatal("expected error for oversized request")
|
|
}
|
|
}
|
|
|
|
func TestSuggestBatchRequiresConfiguredProvider(t *testing.T) {
|
|
db := openSuggestTestDB(t)
|
|
service := NewService(db)
|
|
_, err := service.SuggestBatch(context.Background(), SuggestRequest{
|
|
Sources: []SuggestSource{{ID: "s1", Label: "黑色"}},
|
|
Candidates: []SuggestCandidate{{ID: "c1", Label: "黑色"}},
|
|
})
|
|
target, ok := err.(*Error)
|
|
if !ok {
|
|
t.Fatalf("expected *Error, got %v (%T)", err, err)
|
|
}
|
|
if target.Code != CodeNotConfigured {
|
|
t.Fatalf("expected CodeNotConfigured, got %v", target.Code)
|
|
}
|
|
}
|
|
|
|
func TestSuggestBatchLogsSafeDiagnosticForProvider502(t *testing.T) {
|
|
const sensitiveBody = "api-key-and-provider-body-must-not-be-logged"
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
http.Error(w, sensitiveBody, http.StatusBadGateway)
|
|
}))
|
|
defer server.Close()
|
|
|
|
db := openSuggestTestDB(t)
|
|
seedEnabledSetting(t, db, server.URL)
|
|
var diagnostic ProviderFailureDiagnostic
|
|
service := NewService(db)
|
|
service.ProviderFailureLogger = func(value ProviderFailureDiagnostic) { diagnostic = value }
|
|
|
|
_, err := service.SuggestBatch(context.Background(), SuggestRequest{
|
|
Sources: []SuggestSource{{ID: "s1", Label: "sensitive-source"}},
|
|
Candidates: []SuggestCandidate{{ID: "c1", Label: "sensitive-candidate"}},
|
|
})
|
|
if err == nil {
|
|
t.Fatal("expected provider failure")
|
|
}
|
|
if diagnostic.Operation != "suggest_batch" || diagnostic.Kind != "http_status" || diagnostic.StatusCode != http.StatusBadGateway || diagnostic.CallID == "" {
|
|
t.Fatalf("unexpected diagnostic: %+v", diagnostic)
|
|
}
|
|
printed := fmt.Sprintf("%+v", diagnostic)
|
|
for _, secret := range []string{sensitiveBody, "sensitive-source", "sensitive-candidate", "test-key", server.URL} {
|
|
if strings.Contains(printed, secret) {
|
|
t.Fatalf("diagnostic leaked %q: %s", secret, printed)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestSuggestBatchClassifiesProviderTimeout(t *testing.T) {
|
|
db := openSuggestTestDB(t)
|
|
seedEnabledSetting(t, db, "http://provider.invalid")
|
|
var diagnostic ProviderFailureDiagnostic
|
|
service := NewService(db)
|
|
service.HTTPClient = &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
|
return nil, context.DeadlineExceeded
|
|
})}
|
|
service.ProviderFailureLogger = func(value ProviderFailureDiagnostic) { diagnostic = value }
|
|
|
|
_, err := service.SuggestBatch(context.Background(), SuggestRequest{
|
|
Sources: []SuggestSource{{ID: "s1", Label: "黑色"}}, Candidates: []SuggestCandidate{{ID: "c1", Label: "黑色"}},
|
|
})
|
|
if err == nil || diagnostic.Kind != "timeout" || diagnostic.StatusCode != 0 {
|
|
t.Fatalf("timeout was not safely classified: diagnostic=%+v err=%v", diagnostic, err)
|
|
}
|
|
}
|
|
|
|
func TestSuggestBatchClassifiesInvalidResponseWithoutLoggingBody(t *testing.T) {
|
|
const sensitiveBody = "not-json-with-sensitive-provider-details"
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
_, _ = w.Write([]byte(sensitiveBody))
|
|
}))
|
|
defer server.Close()
|
|
db := openSuggestTestDB(t)
|
|
seedEnabledSetting(t, db, server.URL)
|
|
var diagnostic ProviderFailureDiagnostic
|
|
service := NewService(db)
|
|
service.ProviderFailureLogger = func(value ProviderFailureDiagnostic) { diagnostic = value }
|
|
|
|
_, err := service.SuggestBatch(context.Background(), SuggestRequest{
|
|
Sources: []SuggestSource{{ID: "s1", Label: "黑色"}}, Candidates: []SuggestCandidate{{ID: "c1", Label: "黑色"}},
|
|
})
|
|
if err == nil || diagnostic.Kind != "invalid_response" || strings.Contains(fmt.Sprintf("%+v", diagnostic), sensitiveBody) {
|
|
t.Fatalf("invalid response diagnostic is unsafe or missing: diagnostic=%+v err=%v", diagnostic, err)
|
|
}
|
|
}
|
|
|
|
func TestSuggestBatchAllowsProviderResponseAfterTwoSeconds(t *testing.T) {
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
time.Sleep(2100 * time.Millisecond)
|
|
chatCompletionResponder(`{"suggestions":[{"sourceId":"s1","candidateId":"c1","confidence":0.95,"reason":"match"}]}`)(w, r)
|
|
}))
|
|
defer server.Close()
|
|
|
|
db := openSuggestTestDB(t)
|
|
seedEnabledSetting(t, db, server.URL)
|
|
result, err := NewService(db).SuggestBatch(context.Background(), SuggestRequest{
|
|
Sources: []SuggestSource{{ID: "s1", Label: "黑色"}}, Candidates: []SuggestCandidate{{ID: "c1", Label: "黑色"}},
|
|
})
|
|
if err != nil || result.Decisions["s1"].CandidateID != "c1" {
|
|
t.Fatalf("delayed provider response failed: result=%+v err=%v", result, err)
|
|
}
|
|
}
|
|
|
|
type roundTripFunc func(*http.Request) (*http.Response, error)
|
|
|
|
func (fn roundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { return fn(request) }
|