diff options
Diffstat (limited to 'internal/provider/forwarder.go')
| -rw-r--r-- | internal/provider/forwarder.go | 56 |
1 files changed, 45 insertions, 11 deletions
diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go index 0d40bb1..d7d850d 100644 --- a/internal/provider/forwarder.go +++ b/internal/provider/forwarder.go @@ -14,6 +14,7 @@ import ( "aigw/internal/config" "aigw/internal/domain" + "aigw/internal/providerhealth" "aigw/internal/telemetry" ) @@ -26,9 +27,10 @@ type Result struct { type Forwarder struct { client *http.Client metrics *telemetry.Metrics + health *providerhealth.Tracker } -func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics) *Forwarder { +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, @@ -40,31 +42,41 @@ func New(cfg config.UpstreamHTTPConfig, metrics *telemetry.Metrics) *Forwarder { ResponseHeaderTimeout: time.Duration(cfg.ResponseHeaderTimeoutSecs) * time.Second, ExpectContinueTimeout: time.Second, } - return &Forwarder{client: &http.Client{Transport: transport}, metrics: metrics} + 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 string, originalBody []byte, sourceHeaders http.Header, routes []domain.Route) (Result, error) { +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 := rewriteRequest(originalBody, route.UpstreamModel, protocol) + 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.BaseURL, protocol), bytes.NewReader(body)) + 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() @@ -79,18 +91,34 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque 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}) +} + func rewriteModel(body []byte, upstreamModel string) ([]byte, error) { - return rewriteRequest(body, upstreamModel, "") + 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 { + if protocol == domain.ProtocolOpenAI && wireAPI == "chat_completions" { var stream bool _ = json.Unmarshal(object["stream"], &stream) if stream { @@ -110,12 +138,18 @@ func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) return result, nil } -func endpointURL(baseURL string, protocol domain.Protocol) string { - baseURL = strings.TrimRight(baseURL, "/") - if protocol == domain.ProtocolAnthropic { +func endpointURL(provider domain.Provider, _ domain.Protocol) string { + baseURL := strings.TrimRight(provider.BaseURL, "/") + switch provider.EffectiveWireAPI() { + case "responses": + return baseURL + "/responses" + case "messages": return baseURL + "/messages" + case "chat_completions": + fallthrough + default: + return baseURL + "/chat/completions" } - return baseURL + "/chat/completions" } func setHeaders(target, source http.Header, provider domain.Provider, protocol domain.Protocol, requestID string) { |
