package service import ( "context" "io" "net/http" "strings" "testing" "time" "cmautobuy/admin/model" "cmautobuy/admin/repository" ) func TestParseConfidenceThresholdBPS_精确转换(t *testing.T) { for raw, want := range map[string]int{"0": 0, "85": 8500, "85.5": 8550, "99.99": 9999, "100.00": 10000} { got, err := ParseConfidenceThresholdBPS(raw) if err != nil || got != want { t.Errorf("%q => %d,%v,期望 %d", raw, got, err, want) } } for _, raw := range []string{"", "-1", "100.01", "1.234", "abc"} { if _, err := ParseConfidenceThresholdBPS(raw); err == nil { t.Errorf("非法置信度 %q 应被拒绝", raw) } } } func TestSaveAIProviderConfig_新增后可查询并与审计同事务(t *testing.T) { db := newTestDB(t) now := time.Date(2026, 8, 14, 1, 2, 3, 0, time.UTC) actor := model.User{UserID: "USR-AI-CONFIG", Username: "ai-admin", PasswordHash: "test-hash", Role: model.RoleAdmin, Status: model.UserActive, PasswordChangedAt: model.NowISO(), CreatedAt: model.NowISO(), UpdatedAt: model.NowISO()} if err := repository.CreateUser(db, actor); err != nil { t.Fatal(err) } input := AIProviderInput{Name: "主线路", BaseURL: "https://api.example.com/v1", Model: "model-1", TimeoutSeconds: 30, MaxConcurrency: 2, ConfidenceThresholdBPS: 8500} providerID, err := SaveAIProviderConfig(db, &actor, input, NewAIEndpointPolicy(nil), now) if err != nil || providerID == "" { t.Fatalf("新增 AI 配置失败: id=%q err=%v", providerID, err) } items, err := repository.ListAIProviders(db) if err != nil || len(items) != 1 || items[0].ProviderID != providerID || items[0].Name != input.Name { t.Fatalf("新增后列表未回显配置: items=%+v err=%v", items, err) } var auditCount int if err := db.QueryRow(`SELECT COUNT(*) FROM ai_provider_audits WHERE provider_id=? AND action='create'`, providerID).Scan(&auditCount); err != nil || auditCount != 1 { t.Fatalf("创建审计数量=%d err=%v", auditCount, err) } if _, err := SaveAIProviderConfig(db, &actor, input, NewAIEndpointPolicy(nil), now.Add(time.Minute)); !IsValidationError(err) { t.Fatalf("重名配置应返回可展示的校验错误: %v", err) } if err := db.QueryRow(`SELECT COUNT(*) FROM ai_provider_configs`).Scan(&auditCount); err != nil || auditCount != 1 { t.Fatalf("失败请求不得留下半条配置: count=%d err=%v", auditCount, err) } } func TestAIEndpointPolicy_阻止凭据和内网地址(t *testing.T) { policy := NewAIEndpointPolicy(nil) for _, raw := range []string{ "http://api.example.com/v1", "https://user:pass@api.example.com/v1", "https://127.0.0.1/v1", "https://169.254.169.254/latest", "https://localhost/v1", } { if _, err := policy.ValidateSyntax(raw); err == nil { t.Errorf("危险地址 %q 应被拒绝", raw) } } got, err := policy.ValidateSyntax("https://api.example.com/v1/") if err != nil || got != "https://api.example.com/v1" { t.Fatalf("公网 HTTPS 地址应通过并去掉尾斜杠: %q %v", got, err) } allowed := NewAIEndpointPolicy([]string{"10.0.0.8"}) if _, err := allowed.ValidateSyntax("https://10.0.0.8/v1"); err != nil { t.Fatalf("部署允许的私有端点应通过: %v", err) } } func TestSafeAIHTTPClient_不使用环境代理(t *testing.T) { client := NewSafeAIHTTPClient(NewAIEndpointPolicy(nil), time.Second) transport, ok := client.Transport.(*http.Transport) if !ok { t.Fatalf("Transport 类型 = %T", client.Transport) } if transport.Proxy != nil { t.Fatal("AI HTTP 客户端不能使用环境代理,否则目标地址会绕过拨号阶段的 SSRF 校验") } } type recordingAIHTTPDoer struct { request *http.Request status int } func (d *recordingAIHTTPDoer) Do(request *http.Request) (*http.Response, error) { d.request = request return &http.Response{StatusCode: d.status, Body: io.NopCloser(strings.NewReader(`{"ok":true}`))}, nil } func TestCallAIHealthCheck_最小请求且错误不泄露密钥(t *testing.T) { doer := &recordingAIHTTPDoer{status: http.StatusUnauthorized} const secret = "test-secret-never-log" err := callAIHealthCheck(context.Background(), doer, model.AIProviderConfig{ BaseURL: "https://api.example.com/v1", Model: "test-model", }, secret) if err == nil || strings.Contains(err.Error(), secret) || !strings.Contains(err.Error(), "401") { t.Fatalf("错误必须脱敏且保留 HTTP 状态: %v", err) } if got := doer.request.Header.Get("Authorization"); got != "Bearer "+secret { t.Fatalf("测试请求未使用密钥: %q", got) } body, _ := io.ReadAll(doer.request.Body) text := string(body) for _, forbidden := range []string{"order_no", "syb_id", "address", secret} { if strings.Contains(text, forbidden) { t.Fatalf("最小测试请求包含业务数据或密钥 %q: %s", forbidden, text) } } }