155 lines
6.1 KiB
Go
155 lines
6.1 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
"io"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"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_兼容HTTP并阻止凭据和内网地址(t *testing.T) {
|
|
policy := NewAIEndpointPolicy(nil)
|
|
for _, raw := range []string{
|
|
"ftp://api.example.com/v1", "http://user:pass@api.example.com/v1", "https://user:pass@api.example.com/v1",
|
|
"http://api.example.com/v1?token=fake", "https://api.example.com/v1#fragment",
|
|
"http://127.0.0.1/v1", "https://127.0.0.1/v1", "http://169.254.169.254/latest",
|
|
"https://169.254.169.254/latest", "http://localhost/v1", "https://localhost/v1",
|
|
} {
|
|
if _, err := policy.ValidateSyntax(raw); err == nil {
|
|
t.Errorf("危险地址 %q 应被拒绝", raw)
|
|
}
|
|
}
|
|
for raw, want := range map[string]string{
|
|
"http://api.example.com/v1/": "http://api.example.com/v1",
|
|
"HTTPS://api.example.com/v1/": "https://api.example.com/v1",
|
|
} {
|
|
got, err := policy.ValidateSyntax(raw)
|
|
if err != nil || got != want {
|
|
t.Errorf("公网 HTTP/HTTPS 地址应通过并规范化: raw=%q got=%q err=%v", raw, got, err)
|
|
}
|
|
}
|
|
allowed := NewAIEndpointPolicy([]string{"10.0.0.8"})
|
|
for _, raw := range []string{"http://10.0.0.8/v1", "https://10.0.0.8/v1"} {
|
|
if _, err := allowed.ValidateSyntax(raw); err != nil {
|
|
t.Errorf("现有部署允许列表中的私有 HTTP/HTTPS 端点应通过: %q %v", raw, err)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestSafeAIHTTPClient_可调用允许列表中的HTTP端点(t *testing.T) {
|
|
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
|
|
if request.URL.Path != "/v1/chat/completions" {
|
|
t.Errorf("请求路径 = %q", request.URL.Path)
|
|
}
|
|
if request.Header.Get("Authorization") != "Bearer fake-http-key" {
|
|
t.Errorf("Authorization 未按既有方式发送")
|
|
}
|
|
response.WriteHeader(http.StatusOK)
|
|
}))
|
|
defer server.Close()
|
|
|
|
policy := NewAIEndpointPolicy([]string{"127.0.0.1"})
|
|
if err := policy.ValidateResolved(context.Background(), server.URL+"/v1"); err != nil {
|
|
t.Fatalf("测试级允许列表中的 HTTP 地址应通过完整校验: %v", err)
|
|
}
|
|
client := NewSafeAIHTTPClient(policy, time.Second)
|
|
if err := callAIHealthCheck(context.Background(), client, model.AIProviderConfig{
|
|
BaseURL: server.URL + "/v1", Model: "test-model",
|
|
}, "fake-http-key"); err != nil {
|
|
t.Fatalf("HTTP OpenAI 兼容端点调用失败: %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)
|
|
}
|
|
}
|
|
}
|