Files
superidou 9501751792 0.10.1: 安全与业务逻辑加固、新品牌与部署加固
三轮审查修复(60+ 项),相对远端 main(b536672)的关键变更:
- 安全: 数据面 SSRF 拨号防护(防 DNS rebinding)/上游凭据剥离/登录防枚举
  与锁定态统一/可信代理(X-Forwarded-For)限流加固/会话版本失效机制/
  撤销即时传播/弱密钥拒绝启动/脱敏字节级重写(保签名契约)
- 业务逻辑: 裸 body 上传 panic/bootstrap 审计管线卡死/定价通配符优先级/
  全局工具可见性/调度器停机补跑/TOTP 挑战令牌消费顺序/熔断探针语义/
  >4MB 响应 token 计量/管理员重置密码作废会话 等
- 前端: 新 logo(语枢 AI 网关主题)/Provider 凭据异常警示/删除入口/
  后端错误消息透传/localStorage 敏感数据收敛
- 部署: CREDENTIAL_MASTER_KEY 持久化与弱值拒绝/Provider DELETE 接口/
  nginx 安全头/worker 内存限制
- 新增迁移 000029(key_hash 索引)/000030(usage_daily 币种维度)
2026-08-13 10:50:51 +08:00

240 lines
8.1 KiB
Go

package controlplane
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"strconv"
"strings"
"time"
"aigateway.local/core/internal/provider"
provideropenai "aigateway.local/core/internal/provider/openai"
)
const maxModelsResponseBytes int64 = 4 << 20
type Service struct {
repository *provider.Repository
cipher *provider.CredentialCipher
allowPrivate bool
client *http.Client
}
func NewService(repository *provider.Repository, cipher *provider.CredentialCipher, allowPrivate bool) *Service {
return &Service{
repository: repository,
cipher: cipher,
allowPrivate: allowPrivate,
client: &http.Client{
Timeout: 10 * time.Second,
Transport: &http.Transport{
DialContext: safeDialContext(allowPrivate),
ForceAttemptHTTP2: true,
MaxIdleConns: 32,
MaxIdleConnsPerHost: 4,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 5 * time.Second,
ResponseHeaderTimeout: 8 * time.Second,
},
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
return errors.New("upstream redirects are not allowed")
},
},
}
}
func (s *Service) TestConnection(ctx context.Context, providerID string) (provider.ConnectionTestResult, error) {
record, adapter, err := s.load(ctx, providerID)
if err != nil {
return provider.ConnectionTestResult{}, err
}
started := time.Now()
response, err := s.doModelsRequest(ctx, adapter)
latency := time.Since(started).Milliseconds()
if err != nil {
return provider.ConnectionTestResult{}, fmt.Errorf("%w: request provider %s: %v", provider.ErrProviderUpstream, record.Code, err)
}
defer response.Body.Close()
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 32<<10))
result := provider.ConnectionTestResult{
Connected: response.StatusCode >= 200 && response.StatusCode < 300,
StatusCode: response.StatusCode,
LatencyMS: latency,
}
if result.Connected {
result.Message = "连接成功"
} else {
result.Message = "上游返回 HTTP " + strconv.Itoa(response.StatusCode)
}
return result, nil
}
func (s *Service) SyncModels(ctx context.Context, providerID, actorID string) (provider.ModelSyncResult, error) {
record, adapter, err := s.load(ctx, providerID)
if err != nil {
return provider.ModelSyncResult{}, err
}
if !hasCapability(record.Capabilities, string(provider.CapabilityModels)) {
return provider.ModelSyncResult{}, errors.New("供应商未启用 models 能力")
}
response, err := s.doModelsRequest(ctx, adapter)
if err != nil {
return provider.ModelSyncResult{}, fmt.Errorf("%w: request provider %s: %v", provider.ErrProviderUpstream, record.Code, err)
}
defer response.Body.Close()
if response.StatusCode < 200 || response.StatusCode >= 300 {
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 32<<10))
return provider.ModelSyncResult{}, fmt.Errorf("%w: provider %s returned HTTP %d", provider.ErrProviderUpstream, record.Code, response.StatusCode)
}
payload, err := io.ReadAll(io.LimitReader(response.Body, maxModelsResponseBytes+1))
if err != nil {
return provider.ModelSyncResult{}, fmt.Errorf("%w: read provider %s response: %v", provider.ErrProviderUpstream, record.Code, err)
}
if int64(len(payload)) > maxModelsResponseBytes {
return provider.ModelSyncResult{}, fmt.Errorf("%w: provider %s model response exceeds 4 MiB", provider.ErrProviderUpstream, record.Code)
}
models, err := decodeModels(payload)
if err != nil {
return provider.ModelSyncResult{}, fmt.Errorf("%w: provider %s returned invalid model data: %v", provider.ErrProviderUpstream, record.Code, err)
}
return s.repository.SyncModels(ctx, providerID, actorID, models)
}
func (s *Service) ListModels(ctx context.Context, providerID string) ([]provider.Model, error) {
if _, err := s.repository.Get(ctx, providerID); err != nil {
return nil, err
}
return s.repository.ListModels(ctx, providerID)
}
func (s *Service) RotateCredentials(ctx context.Context, actorID string) (provider.CredentialRotationResult, error) {
result := provider.CredentialRotationResult{
ActiveVersion: s.cipher.ActiveVersion(), LoadedVersions: s.cipher.Versions(),
}
records, err := s.repository.List(ctx)
if err != nil {
return provider.CredentialRotationResult{}, err
}
rotations := make([]provider.CredentialRotation, 0, len(records))
for _, record := range records {
if record.CredentialKEKVersion == result.ActiveVersion {
result.Skipped++
continue
}
plaintext, err := s.cipher.Decrypt(record.EncryptedCredentials, record.CredentialKEKVersion)
if err != nil {
// 单条损坏(如 KEK 版本被删除)不阻塞其余 Provider 的轮换:
// 跳过并计数,管理端从 Skipped 明细中定位问题记录。
result.Skipped++
continue
}
ciphertext, version, err := s.cipher.Encrypt(plaintext)
if err != nil {
return provider.CredentialRotationResult{}, err
}
rotations = append(rotations, provider.CredentialRotation{
ProviderID: record.ID, FromVersion: record.CredentialKEKVersion,
ToVersion: version, Ciphertext: ciphertext,
})
}
if err := s.repository.RotateCredentials(ctx, actorID, rotations); err != nil {
return provider.CredentialRotationResult{}, err
}
result.Rotated = len(rotations)
return result, nil
}
func (s *Service) load(ctx context.Context, providerID string) (provider.Record, provider.Adapter, error) {
record, err := s.repository.Get(ctx, providerID)
if err != nil {
return provider.Record{}, nil, err
}
validatedURL, err := provider.ValidateBaseURL(ctx, record.BaseURL, s.allowPrivate)
if err != nil {
return provider.Record{}, nil, fmt.Errorf("供应商地址校验失败: %w", err)
}
plaintext, err := s.cipher.Decrypt(record.EncryptedCredentials, record.CredentialKEKVersion)
if err != nil {
return provider.Record{}, nil, err
}
var credentials provider.Credentials
if err := json.Unmarshal(plaintext, &credentials); err != nil {
return provider.Record{}, nil, errors.New("供应商凭据格式无效")
}
switch record.Adapter {
case "openai-compatible":
adapter, err := provideropenai.New(validatedURL, credentials.APIKey)
return record, adapter, err
default:
return provider.Record{}, nil, fmt.Errorf("不支持的供应商适配器 %q", record.Adapter)
}
}
func (s *Service) doModelsRequest(ctx context.Context, adapter provider.Adapter) (*http.Response, error) {
target := adapter.Target()
target.Path = strings.TrimRight(target.Path, "/") + "/v1/models"
target.RawPath = ""
request, err := http.NewRequestWithContext(ctx, http.MethodGet, target.String(), nil)
if err != nil {
return nil, err
}
request.Header.Set("Accept", "application/json")
adapter.Prepare(request)
return s.client.Do(request)
}
func decodeModels(payload []byte) ([]provider.DiscoveredModel, error) {
var response struct {
Data []json.RawMessage `json:"data"`
}
if err := json.Unmarshal(payload, &response); err != nil {
return nil, err
}
seen := make(map[string]struct{}, len(response.Data))
models := make([]provider.DiscoveredModel, 0, len(response.Data))
for _, raw := range response.Data {
var item struct {
ID string `json:"id"`
OwnedBy string `json:"owned_by"`
}
if err := json.Unmarshal(raw, &item); err != nil {
return nil, err
}
item.ID = strings.TrimSpace(item.ID)
if item.ID == "" || len(item.ID) > 512 {
return nil, errors.New("model id must contain 1 to 512 characters")
}
if _, exists := seen[item.ID]; exists {
continue
}
seen[item.ID] = struct{}{}
models = append(models, provider.DiscoveredModel{
ProviderModelID: item.ID,
OwnedBy: strings.TrimSpace(item.OwnedBy),
Metadata: append(json.RawMessage(nil), raw...),
})
}
return models, nil
}
func hasCapability(capabilities []string, expected string) bool {
for _, capability := range capabilities {
if capability == expected {
return true
}
}
return false
}
func safeDialContext(allowPrivate bool) func(context.Context, string, string) (net.Conn, error) {
// 统一复用 provider.IsPublicAddress 的完整网段判定(含 CGNAT/6to4/NAT64 等)。
return provider.SafeDialContext(allowPrivate, 5*time.Second, 30*time.Second)
}
var _ provider.AdminOperations = (*Service)(nil)