Files
goauto/server/app/goauto/aimatching/suggest_test.go
T

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) }