summaryrefslogtreecommitdiff
path: root/internal/provider/forwarder.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/provider/forwarder.go')
-rw-r--r--internal/provider/forwarder.go56
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) {