package provider import ( "bytes" "context" "encoding/json" "errors" "fmt" "io" "net" "net/http" "strings" "time" "aigw/internal/config" "aigw/internal/domain" "aigw/internal/providerhealth" "aigw/internal/telemetry" ) type Result struct { Response *http.Response Route domain.Route Attempts int AttemptStartedAt time.Time } type Forwarder struct { client *http.Client metrics *telemetry.Metrics health *providerhealth.Tracker } func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics, trackers ...*providerhealth.Tracker) *Forwarder { transport := &http.Transport{ Proxy: http.ProxyFromEnvironment, DialContext: (&net.Dialer{Timeout: 5 * time.Second, KeepAlive: 30 * time.Second}).DialContext, ForceAttemptHTTP2: true, MaxIdleConns: cfg.MaxIdleConnections, MaxIdleConnsPerHost: cfg.MaxIdleConnectionsPerHost, IdleConnTimeout: time.Duration(cfg.IdleConnectionTimeoutSecs) * time.Second, TLSHandshakeTimeout: 10 * time.Second, ResponseHeaderTimeout: time.Duration(cfg.ResponseHeaderTimeoutSecs) * time.Second, ExpectContinueTimeout: time.Second, } forwarder := &Forwarder{client: &http.Client{Transport: transport}, metrics: metrics} if len(trackers) > 0 { forwarder.health = trackers[0] } return forwarder } func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, requestID, modelID string, originalBody []byte, sourceHeaders http.Header, routes []domain.Route) (Result, error) { var lastErr error for i, route := range routes { if err := ctx.Err(); err != nil { return Result{Attempts: i}, err } body, err := rewriteRequestWithWireAPI(originalBody, route.UpstreamModel, protocol, route.Provider.EffectiveWireAPI()) if err != nil { return Result{Attempts: i}, err } request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpointURL(route.Provider, protocol), bytes.NewReader(body)) if err != nil { return Result{Attempts: i}, fmt.Errorf("build upstream request: %w", err) } setHeaders(request.Header, sourceHeaders, route.Provider, protocol, requestID) f.metrics.UpstreamAttempt() attemptStarted := time.Now() response, err := f.client.Do(request) if err != nil { if ctx.Err() != nil { return Result{Attempts: i + 1}, ctx.Err() } f.observe(modelID, route, 0, time.Since(attemptStarted), true) lastErr = err continue } attempts := i + 1 f.observe(modelID, route, response.StatusCode, time.Since(attemptStarted), retryableStatus(response.StatusCode)) if retryableStatus(response.StatusCode) && attempts < len(routes) { _, _ = io.CopyN(io.Discard, response.Body, 8<<10) _ = response.Body.Close() lastErr = fmt.Errorf("upstream %s returned %d", route.Provider.ID, response.StatusCode) continue } return Result{Response: response, Route: route, Attempts: attempts, AttemptStartedAt: attemptStarted}, nil } if lastErr == nil { lastErr = errors.New("all upstream routes failed") } return Result{Attempts: len(routes)}, lastErr } func (f *Forwarder) observe(modelID string, route domain.Route, statusCode int, latency time.Duration, failed bool) { if f.health == nil { return } f.health.Observe(providerhealth.RouteKey{ModelID: modelID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}, providerhealth.Observation{StatusCode: statusCode, Latency: latency, Failed: failed}) } // ObserveTTFT feeds the first user-visible output latency into adaptive route // selection. It is deliberately separate from the header/circuit observation // because streaming TTFT is only known after response forwarding begins. func (f *Forwarder) ObserveTTFT(modelID string, route domain.Route, latency time.Duration) { if f.health == nil || latency <= 0 { return } f.health.ObserveTTFT(providerhealth.RouteKey{ModelID: modelID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}, latency) if f.metrics != nil { f.metrics.UpstreamTTFT(latency) } } func rewriteModel(body []byte, upstreamModel string) ([]byte, error) { return rewriteRequest(body, upstreamModel, domain.ProtocolOpenAI) } func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) ([]byte, error) { wireAPI := "chat_completions" if protocol == domain.ProtocolAnthropic { wireAPI = "messages" } return rewriteRequestWithWireAPI(body, upstreamModel, protocol, wireAPI) } func rewriteRequestWithWireAPI(body []byte, upstreamModel string, protocol domain.Protocol, wireAPI string) ([]byte, error) { var object map[string]json.RawMessage if err := json.Unmarshal(body, &object); err != nil { return nil, fmt.Errorf("decode request body: %w", err) } encoded, _ := json.Marshal(upstreamModel) object["model"] = encoded if protocol == domain.ProtocolOpenAI && wireAPI == "chat_completions" { var stream bool _ = json.Unmarshal(object["stream"], &stream) if stream { var options map[string]json.RawMessage _ = json.Unmarshal(object["stream_options"], &options) if options == nil { options = map[string]json.RawMessage{} } options["include_usage"] = json.RawMessage("true") object["stream_options"], _ = json.Marshal(options) } } result, err := json.Marshal(object) if err != nil { return nil, fmt.Errorf("encode upstream request: %w", err) } return result, nil } func endpointURL(provider domain.Provider, _ domain.Protocol) string { baseURL := strings.TrimRight(provider.BaseURL, "/") switch provider.EffectiveWireAPI() { case "embeddings": return baseURL + "/embeddings" case "responses": return baseURL + "/responses" case "messages": return baseURL + "/messages" case "chat_completions": fallthrough default: return baseURL + "/chat/completions" } } func setHeaders(target, source http.Header, provider domain.Provider, protocol domain.Protocol, requestID string) { target.Set("Content-Type", "application/json") target.Set("Accept", source.Get("Accept")) if target.Get("Accept") == "" { target.Set("Accept", "application/json") } target.Set("User-Agent", "aigw/0.1") target.Set("X-Request-ID", requestID) if protocol == domain.ProtocolAnthropic { target.Set("x-api-key", provider.APIKey) version := source.Get("anthropic-version") if version == "" { version = "2023-06-01" } target.Set("anthropic-version", version) if beta := source.Get("anthropic-beta"); beta != "" { target.Set("anthropic-beta", beta) } return } target.Set("Authorization", "Bearer "+provider.APIKey) } func retryableStatus(status int) bool { switch status { case http.StatusTooManyRequests, http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout: return true default: return false } }