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 }