diff options
| author | Chia <Chia@93.nz> | 2026-08-04 19:58:52 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-04 20:43:23 +1200 |
| commit | 5b651488b081b65fda8a323f228e139adb79a35d (patch) | |
| tree | 08baf40efb8fe103b32721cd991ff712323e3173 /internal/provider | |
Build AI gateway control plane and admin UI
Diffstat (limited to 'internal/provider')
| -rw-r--r-- | internal/provider/forwarder.go | 134 |
1 files changed, 134 insertions, 0 deletions
diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go new file mode 100644 index 0000000..a9d5734 --- /dev/null +++ b/internal/provider/forwarder.go @@ -0,0 +1,134 @@ +package provider + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "strings" + "time" + + "aigw/internal/config" + "aigw/internal/domain" + "aigw/internal/telemetry" +) + +type Result struct { + Response *http.Response + Route domain.Route + Attempts int +} + +type Forwarder struct { + client *http.Client + metrics *telemetry.Metrics +} + +func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics) *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, + } + return &Forwarder{client: &http.Client{Transport: transport}, metrics: metrics} +} + +func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, requestID 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 := rewriteModel(originalBody, route.UpstreamModel) + if err != nil { + return Result{Attempts: i}, err + } + request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpointURL(route.Provider.BaseURL, 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() + response, err := f.client.Do(request) + if err != nil { + lastErr = err + continue + } + attempts := i + 1 + 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}, nil + } + if lastErr == nil { + lastErr = errors.New("all upstream routes failed") + } + return Result{Attempts: len(routes)}, lastErr +} + +func rewriteModel(body []byte, upstreamModel 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 + result, err := json.Marshal(object) + if err != nil { + return nil, fmt.Errorf("encode upstream request: %w", err) + } + return result, nil +} + +func endpointURL(baseURL string, protocol domain.Protocol) string { + baseURL = strings.TrimRight(baseURL, "/") + if protocol == domain.ProtocolAnthropic { + return baseURL + "/messages" + } + 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 + } +} |
