2026-08-14 09:49:11 +08:00
package service
import (
"context"
"io"
"net/http"
2026-08-14 11:15:50 +08:00
"net/http/httptest"
2026-08-14 09:49:11 +08:00
"strings"
"testing"
"time"
"cmautobuy/admin/model"
2026-08-14 11:04:51 +08:00
"cmautobuy/admin/repository"
2026-08-14 09:49:11 +08:00
)
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 )
}
}
}
2026-08-14 11:04:51 +08:00
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 )
}
}
2026-08-14 11:15:50 +08:00
func TestAIEndpointPolicy_兼容HTTP并阻止凭据和内网地址 ( t * testing . T ) {
2026-08-14 09:49:11 +08:00
policy := NewAIEndpointPolicy ( nil )
for _ , raw := range [] string {
2026-08-14 11:15:50 +08:00
"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" ,
2026-08-14 09:49:11 +08:00
} {
if _ , err := policy . ValidateSyntax ( raw ); err == nil {
t . Errorf ( "危险地址 %q 应被拒绝" , raw )
}
}
2026-08-14 11:15:50 +08:00
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 )
}
2026-08-14 09:49:11 +08:00
}
allowed := NewAIEndpointPolicy ([] string { "10.0.0.8" })
2026-08-14 11:15:50 +08:00
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 )
2026-08-14 09:49:11 +08:00
}
}
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 )
}
}
}