summaryrefslogtreecommitdiff
path: root/internal/provider/forwarder.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-04 19:58:52 +1200
committerChia <Chia@93.nz>2026-08-04 20:43:23 +1200
commit5b651488b081b65fda8a323f228e139adb79a35d (patch)
tree08baf40efb8fe103b32721cd991ff712323e3173 /internal/provider/forwarder.go
Build AI gateway control plane and admin UI
Diffstat (limited to 'internal/provider/forwarder.go')
-rw-r--r--internal/provider/forwarder.go134
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
+ }
+}