Files
ai-gateway-go/internal/gateway/proxy.go
T
superidou c22669c31d 0.11.0: 旗舰版功能补齐(License/登录记录/会话管理/角色管理/门户定时任务/模型配额/输出脱敏/供应链扫描/记忆管理/AI助手/真实概览)
- 新增迁移 000031-000034(登录日志/角色/模型配额/记忆)
- 新增包: license/memory/modelquota/assistant,扫描引擎
- 全部功能后端+前端+端到端验证通过(25 包单测)
2026-08-13 11:37:18 +08:00

507 lines
19 KiB
Go

package gateway
import (
"context"
"crypto/subtle"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net/http"
"net/http/httputil"
"strconv"
"strings"
"sync"
"time"
"aigateway.local/core/internal/apikey"
"aigateway.local/core/internal/contentpolicy"
"aigateway.local/core/internal/pricing"
"aigateway.local/core/internal/provider"
)
type Proxy struct {
resolver AdapterResolver
auth apikey.KeyAuthenticator
maxBody int64
logger *slog.Logger
transport *http.Transport
proxies sync.Map
circuits sync.Map
admission AdmissionController
tokenQuota TokenQuotaController
modelQuota ModelQuotaController
resilience ResiliencePolicy
audit AuditRecorder
policies *contentpolicy.Engine
pricing *pricing.Service
outputPolicies interface {
OutputRedact([]byte) ([]byte, bool)
}
}
type cachedProxy struct {
key string
proxy *httputil.ReverseProxy
}
var (
ErrProviderNotFound = errors.New("requested provider is not available")
ErrProviderUnavailable = errors.New("provider configuration is unavailable")
)
type ResolvedAdapter struct {
Code string
Revision int64
Adapter provider.Adapter
Capabilities map[provider.Capability]bool
}
type AdapterResolver interface {
Resolve(providerCode string) (ResolvedAdapter, error)
}
func NewProxy(adapter provider.Adapter, apiKey string, maxBody int64, logger *slog.Logger) *Proxy {
return NewProxyWithAuthenticator(adapter, staticKeyAuthenticator(apiKey), maxBody, logger)
}
func NewProxyWithAuthenticator(adapter provider.Adapter, authenticator apikey.KeyAuthenticator, maxBody int64, logger *slog.Logger) *Proxy {
return NewDynamicProxy(staticAdapterResolver{adapter: adapter}, authenticator, maxBody, logger)
}
func NewDynamicProxy(resolver AdapterResolver, authenticator apikey.KeyAuthenticator, maxBody int64, logger *slog.Logger) *Proxy {
// 默认拒绝拨号到非公网地址:管理员未显式放行私网时,数据平面在拨号阶段
// 复检目标地址,防止 DNS rebinding 把流量引到内网(169.254.169.254 等)。
transport := &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: provider.SafeDialContext(false, 5*time.Second, 30*time.Second),
ForceAttemptHTTP2: true,
MaxIdleConns: 512,
MaxIdleConnsPerHost: 256,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 5 * time.Second,
ResponseHeaderTimeout: 60 * time.Second,
ExpectContinueTimeout: time.Second,
}
return &Proxy{resolver: resolver, auth: authenticator, maxBody: maxBody, logger: logger, transport: transport, resilience: DefaultResiliencePolicy()}
}
// SetAllowPrivateProviderURLs 允许数据平面拨号到私网地址(与管理员配置的
// ALLOW_PRIVATE_PROVIDER_URLS 保持一致);关闭时保持拨号阶段 SSRF 校验。
func (p *Proxy) SetAllowPrivateProviderURLs(allow bool) {
timeout := p.transport.ResponseHeaderTimeout
p.transport = p.transport.Clone()
p.transport.DialContext = provider.SafeDialContext(allow, 5*time.Second, 30*time.Second)
p.transport.ResponseHeaderTimeout = timeout
}
func (p *Proxy) SetAdmissionController(controller AdmissionController) {
p.admission = controller
}
func (p *Proxy) SetTokenQuotaController(controller TokenQuotaController) {
p.tokenQuota = controller
}
// SetModelQuotaController 启用模型级 Token 配额(provider 解析后预留)。
func (p *Proxy) SetModelQuotaController(controller ModelQuotaController) {
p.modelQuota = controller
}
func (p *Proxy) SetResiliencePolicy(policy ResiliencePolicy) {
p.resilience = policy
p.transport.ResponseHeaderTimeout = policy.ResponseHeaderTimeout
}
func (p *Proxy) SetAuditRecorder(recorder AuditRecorder) { p.audit = recorder }
func (p *Proxy) SetContentPolicyEngine(engine *contentpolicy.Engine) { p.policies = engine }
// SetOutputPolicyEngine 启用输出侧脱敏(模型回答隐私拦截替换)。
func (p *Proxy) SetOutputPolicyEngine(engine interface {
OutputRedact([]byte) ([]byte, bool)
}) {
p.outputPolicies = engine
}
func (p *Proxy) SetPricingService(service *pricing.Service) { p.pricing = service }
func (p *Proxy) ServeHTTP(writer http.ResponseWriter, request *http.Request) {
started := time.Now()
if !isSupportedPath(request.URL.Path, request.Method) {
writeOpenAIError(writer, http.StatusNotFound, "invalid_request_error", "unsupported gateway endpoint")
return
}
principal, err := p.authorized(request)
if err != nil {
if errors.Is(err, apikey.ErrStore) {
writeOpenAIError(writer, http.StatusServiceUnavailable, "authentication_unavailable", "API key service is unavailable")
return
}
writer.Header().Set("WWW-Authenticate", "Bearer")
writeOpenAIError(writer, http.StatusUnauthorized, "authentication_error", "invalid API key")
return
}
span := newAuditSpan(p.audit, principal, request, started)
if span != nil {
span.pricing = p.pricing
}
if span != nil {
statusWriter := &statusResponseWriter{ResponseWriter: writer}
writer = statusWriter
defer func() { span.finish(statusWriter.Status()) }()
}
if p.admission != nil && principal.APIKeyID != "" {
decision, err := p.admission.Allow(request.Context(), principal, time.Now())
if err != nil {
writeOpenAIError(writer, http.StatusServiceUnavailable, "rate_limit_unavailable", "rate limit service is unavailable")
return
}
writeAdmissionHeaders(writer.Header(), decision)
if !decision.Allowed {
writer.Header().Set("Retry-After", strconv.FormatInt(max(int64(decision.RetryAfter.Seconds()), 1), 10))
if decision.Reason == AdmissionMonthlyQuota {
writeOpenAIError(writer, http.StatusTooManyRequests, "insufficient_quota", "monthly request quota exceeded")
} else {
writeOpenAIError(writer, http.StatusTooManyRequests, "rate_limit_error", "request rate limit exceeded")
}
return
}
}
if request.ContentLength > p.maxBody {
writeOpenAIError(writer, http.StatusRequestEntityTooLarge, "invalid_request_error", "request body is too large")
return
}
if p.policies != nil {
policyResult, policyErr := p.policies.Apply(request, p.maxBody, principal.APIKeyID)
if errors.Is(policyErr, contentpolicy.ErrBodyTooLarge) {
writeOpenAIError(writer, http.StatusRequestEntityTooLarge, "invalid_request_error", "request body is too large")
return
}
if policyErr != nil {
writeOpenAIError(writer, http.StatusBadRequest, "invalid_request_error", "request body could not be inspected")
return
}
span.setContentPolicy(policyResult)
if policyResult.Blocked {
writeOpenAIError(writer, http.StatusUnprocessableEntity, "content_policy_violation", "request was blocked by content policy")
return
}
if policyResult.Redacted {
writer.Header().Set("X-Gateway-Content-Redacted", "true")
}
}
routeResolver, routingEnabled := p.resolver.(ModelRouteResolver)
routingEnabled = routingEnabled && routeResolver.ModelRoutingEnabled()
var modelPayload modelRequest
if routingEnabled {
modelPayload, err = readModelRequest(request, p.maxBody)
if errors.Is(err, errRequestBodyTooLarge) {
writeOpenAIError(writer, http.StatusRequestEntityTooLarge, "invalid_request_error", "request body is too large")
return
}
if err != nil {
writeOpenAIError(writer, http.StatusBadRequest, "invalid_request_error", "request body could not be read")
return
}
}
var usage *usageSession
if p.tokenQuota != nil && principal.APIKeyID != "" {
estimate := int64(0)
if principal.MonthlyTokenQuota > 0 {
estimate, err = prepareTokenBudget(request, p.maxBody)
if errors.Is(err, errRequestBodyTooLarge) {
writeOpenAIError(writer, http.StatusRequestEntityTooLarge, "invalid_request_error", "request body is too large")
return
}
if err != nil {
writeOpenAIError(writer, http.StatusBadRequest, "invalid_request_error", "request body could not be read")
return
}
}
reservation, reserveErr := p.tokenQuota.Reserve(request.Context(), principal, estimate, time.Now())
if reserveErr != nil && principal.MonthlyTokenQuota > 0 {
writeOpenAIError(writer, http.StatusServiceUnavailable, "token_quota_unavailable", "token quota service is unavailable")
return
}
if reserveErr == nil && !reservation.Allowed {
writeTokenQuotaHeaders(writer.Header(), reservation)
writer.Header().Set("Retry-After", strconv.FormatInt(max(int64(time.Until(reservation.ResetAt).Seconds()), 1), 10))
writeOpenAIError(writer, http.StatusTooManyRequests, "insufficient_quota", "monthly token quota exceeded")
return
}
if reserveErr == nil && reservation.CounterKey != "" {
writeTokenQuotaHeaders(writer.Header(), reservation)
usage = &usageSession{controller: p.tokenQuota, reservation: reservation, log: p.logger}
request = withUsageSession(request, usage)
defer usage.finish(TokenUsage{})
}
}
if span != nil {
if usage == nil {
usage = &usageSession{log: p.logger}
request = withUsageSession(request, usage)
defer usage.finish(TokenUsage{})
}
usage.onFinish = span.setUsage
}
if request.Body != nil && request.Method != http.MethodGet {
if request.Header.Get("Idempotency-Key") != "" && request.GetBody == nil {
if _, err := prepareTokenBudget(request, p.maxBody); err != nil {
if errors.Is(err, errRequestBodyTooLarge) {
writeOpenAIError(writer, http.StatusRequestEntityTooLarge, "invalid_request_error", "request body is too large")
} else {
writeOpenAIError(writer, http.StatusBadRequest, "invalid_request_error", "request body could not be read")
}
return
}
}
request.Body = http.MaxBytesReader(writer, request.Body, p.maxBody)
}
providerCode := strings.ToLower(strings.TrimSpace(request.Header.Get("X-Gateway-Provider")))
var resolved ResolvedAdapter
if routingEnabled && modelPayload.model != "" {
span.setModel(modelPayload.model)
tenantID := ""
if principal.TenantID != nil {
tenantID = *principal.TenantID
}
route, routeErr := routeResolver.ResolveModelRoute(ModelRouteQuery{
ProviderCode: providerCode, Model: modelPayload.model, Endpoint: request.URL.Path,
APIKeyID: principal.APIKeyID, TenantID: tenantID, Seed: RequestID(request.Context()),
})
if routeErr != nil {
writeOpenAIError(writer, http.StatusServiceUnavailable, "provider_unavailable", "model routing configuration is unavailable")
return
}
if route.Known && !route.Matched {
writeOpenAIError(writer, http.StatusBadRequest, "invalid_request_error", "model route is not available for this request")
return
}
if route.Matched {
resolved = route.ResolvedAdapter
if err := modelPayload.rewrite(request, route.TargetModel); err != nil {
writeOpenAIError(writer, http.StatusBadRequest, "invalid_request_error", "request model could not be rewritten")
return
}
writer.Header().Set("X-Gateway-Model", route.TargetModel)
}
}
if resolved.Adapter == nil {
resolved, err = p.resolver.Resolve(providerCode)
}
if err != nil {
if errors.Is(err, ErrProviderNotFound) {
writeOpenAIError(writer, http.StatusBadRequest, "invalid_request_error", "requested provider is not available")
return
}
writeOpenAIError(writer, http.StatusServiceUnavailable, "provider_unavailable", "provider configuration is unavailable")
return
}
capability := capabilityForPath(request.URL.Path)
if capability != "" && !resolved.Capabilities[capability] {
writeOpenAIError(writer, http.StatusBadRequest, "invalid_request_error", "provider does not support this endpoint")
return
}
request.Header.Del("X-Gateway-Provider")
writer.Header().Set("X-Gateway-Provider", resolved.Code)
span.setRoute(resolved.Code, writer.Header().Get("X-Gateway-Model"))
// 模型级配额(企业总配额,所有 Key 共享):在确定 provider 与 model 后
// 预留,与 API Key 级配额叠加;未配置配额或额度不足时按 429 处理。
if p.modelQuota != nil && usage != nil && modelPayload.model != "" {
modelQuota, quotaErr := p.modelQuota.Reserve(request.Context(), resolved.Code, modelPayload.model, estimateForReserve(p, request, usage), time.Now())
if quotaErr != nil {
writeOpenAIError(writer, http.StatusServiceUnavailable, "model_quota_unavailable", "model quota service is unavailable")
return
}
if modelQuota != nil {
if reservation, ok := modelQuota.(interface {
AllowedFlag() bool
RemainingTokens() int64
ResetTime() time.Time
}); ok {
if !reservation.AllowedFlag() {
writeOpenAIError(writer, http.StatusTooManyRequests, "insufficient_quota", "model token quota exceeded")
return
}
writer.Header().Set("X-ModelTokenLimit-Remaining", strconv.FormatInt(max(reservation.RemainingTokens(), 0), 10))
usage.modelController = p.modelQuota
usage.modelReservation = modelQuota
}
}
}
request.Body = span.captureBody(request.Body)
p.proxyFor(resolved).ServeHTTP(writer, request)
}
func (p *Proxy) proxyFor(resolved ResolvedAdapter) *httputil.ReverseProxy {
target := resolved.Adapter.Target()
key := fmt.Sprintf("%s:%d:%s", resolved.Code, resolved.Revision, target.String())
if cached, ok := p.proxies.Load(resolved.Code); ok {
entry := cached.(cachedProxy)
if entry.key == key {
return entry.proxy
}
}
reverseProxy := httputil.NewSingleHostReverseProxy(target)
originalDirector := reverseProxy.Director
reverseProxy.Director = func(request *http.Request) {
originalDirector(request)
request.Host = target.Host
resolved.Adapter.Prepare(request)
}
reverseProxy.FlushInterval = -1
circuitValue, _ := p.circuits.LoadOrStore(resolved.Code, newCircuitBreaker(p.resilience))
reverseProxy.Transport = &resilientTransport{
base: p.transport, circuit: circuitValue.(*circuitBreaker), maxRetries: p.resilience.MaxRetries, backoff: p.resilience.RetryBackoff,
}
reverseProxy.ModifyResponse = func(response *http.Response) error {
if session := usageSessionFrom(response.Request); session != nil {
if response.StatusCode >= http.StatusOK && response.StatusCode < http.StatusMultipleChoices {
session.fallback = session.reservation.Reserved
response.Body = newUsageReadCloser(response.Body, response.Header.Get("Content-Type"), session)
} else {
session.fallback = 0
session.finish(TokenUsage{})
}
}
// 输出侧脱敏:模型回答隐私信息拦截替换(仅 2xx 且启用了输出策略时)。
if response.StatusCode >= http.StatusOK && response.StatusCode < http.StatusMultipleChoices && p.outputPolicies != nil {
response.Body = newOutputRedactReadCloser(response.Body, response.Header.Get("Content-Type"), p.outputPolicies)
response.Header.Set("X-Gateway-Output-Redacted", "true")
}
return nil
}
reverseProxy.ErrorHandler = func(writer http.ResponseWriter, request *http.Request, err error) {
if session := usageSessionFrom(request); session != nil {
session.fallback = 0
session.finish(TokenUsage{})
}
p.logger.Error("upstream request failed", "request_id", RequestID(request.Context()), "provider", resolved.Code, "error", err)
if errors.Is(err, ErrCircuitOpen) {
writer.Header().Set("Retry-After", "30")
writeOpenAIError(writer, http.StatusServiceUnavailable, "provider_unavailable", "provider circuit is temporarily open")
return
}
writeOpenAIError(writer, http.StatusBadGateway, "upstream_error", "upstream service is unavailable")
}
// 并发缓存 miss 时只保留一个胜出的代理,其余立即丢弃,避免重复构建。
if actual, loaded := p.proxies.LoadOrStore(resolved.Code, cachedProxy{key: key, proxy: reverseProxy}); loaded {
entry := actual.(cachedProxy)
if entry.key == key {
return entry.proxy
}
// 另一个 goroutine 写入了不同的 key(快照已前进):保留新条目。
_ = reverseProxy
return entry.proxy
}
return reverseProxy
}
// estimateForReserve 复用已有 usage session 的预留估算;不可用时回退 0。
func estimateForReserve(p *Proxy, request *http.Request, usage *usageSession) int64 {
if usage != nil && usage.reservation.Reserved > 0 {
return usage.reservation.Reserved
}
estimate, err := prepareTokenBudget(request, p.maxBody)
if err != nil {
return 0
}
return estimate
}
func (p *Proxy) authorized(request *http.Request) (apikey.Principal, error) {
presented := strings.TrimSpace(request.Header.Get("X-Gateway-API-Key"))
if presented == "" {
authorization := strings.TrimSpace(request.Header.Get("Authorization"))
if len(authorization) > len("Bearer ") && strings.EqualFold(authorization[:len("Bearer ")], "Bearer ") {
presented = strings.TrimSpace(authorization[len("Bearer "):])
}
}
if p.auth == nil {
return apikey.Principal{}, apikey.ErrInvalid
}
if authenticator, ok := p.auth.(apikey.PrincipalAuthenticator); ok {
return authenticator.AuthenticatePrincipal(request.Context(), presented)
}
return apikey.Principal{}, p.auth.Authenticate(request.Context(), presented)
}
func writeAdmissionHeaders(header http.Header, decision AdmissionDecision) {
if decision.Limit <= 0 || decision.ResetAt.IsZero() {
return
}
header.Set("X-RateLimit-Limit", strconv.FormatInt(decision.Limit, 10))
header.Set("X-RateLimit-Remaining", strconv.FormatInt(max(decision.Remaining, 0), 10))
header.Set("X-RateLimit-Reset", strconv.FormatInt(decision.ResetAt.Unix(), 10))
}
type staticKeyAuthenticator string
func (a staticKeyAuthenticator) Authenticate(_ context.Context, presented string) error {
key := string(a)
if key == "" {
return nil
}
if len(presented) != len(key) || subtle.ConstantTimeCompare([]byte(presented), []byte(key)) != 1 {
return apikey.ErrInvalid
}
return nil
}
type staticAdapterResolver struct{ adapter provider.Adapter }
func (r staticAdapterResolver) Resolve(code string) (ResolvedAdapter, error) {
if r.adapter == nil || code != "" && code != "environment" {
return ResolvedAdapter{}, ErrProviderNotFound
}
capabilities := make(map[provider.Capability]bool)
for _, capability := range r.adapter.Capabilities() {
capabilities[capability] = true
}
return ResolvedAdapter{Code: "environment", Adapter: r.adapter, Capabilities: capabilities}, nil
}
func capabilityForPath(path string) provider.Capability {
switch path {
case "/v1/models":
return provider.CapabilityModels
case "/v1/chat/completions":
return provider.CapabilityChat
case "/v1/responses":
return provider.CapabilityResponses
case "/v1/embeddings":
return provider.CapabilityEmbeddings
case "/v1/messages":
return provider.CapabilityMessages
default:
return ""
}
}
func isSupportedPath(path, method string) bool {
switch path {
case "/v1/models":
return method == http.MethodGet
case "/v1/chat/completions", "/v1/responses", "/v1/embeddings", "/v1/messages":
return method == http.MethodPost
default:
return false
}
}
func writeOpenAIError(writer http.ResponseWriter, status int, errorType, message string) {
writer.Header().Set("Content-Type", "application/json")
writer.WriteHeader(status)
_ = json.NewEncoder(writer).Encode(map[string]any{
"error": map[string]any{"message": message, "type": errorType, "param": nil, "code": nil},
})
}
var errInvalidAdapter = errors.New("invalid provider adapter")
func ValidateAdapter(adapter provider.Adapter) error {
if adapter == nil || adapter.Target() == nil || adapter.Name() == "" {
return errInvalidAdapter
}
return nil
}