Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 18 additions & 4 deletions internal/server/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,7 @@ func (s *Server) proxyOpenAI(w http.ResponseWriter, r *http.Request) {
setBackendAuth(req, backend)
}
proxy.ModifyResponse = func(resp *http.Response) error {
instrumentBackendResponse(resp, started, waited, routedModel, backend, requestID, s.metrics)
instrumentBackendResponse(resp, started, waited, routedModel, backend, requestID, routeRule, s.metrics)
return nil
}
proxy.ErrorHandler = func(rw http.ResponseWriter, req *http.Request, proxyErr error) {
Expand All @@ -202,7 +202,9 @@ func (s *Server) proxyOpenAI(w http.ResponseWriter, r *http.Request) {
"duration_ms", time.Since(started).Milliseconds(),
"error", proxyErr,
)
s.metrics.record(requestMetricsFromModel(routedModel, backend, http.StatusBadGateway, started))
metrics := requestMetricsFromModel(routedModel, backend, http.StatusBadGateway, started)
metrics.RouteRule = routeRule
s.metrics.record(metrics)
writeOpenAIError(rw, http.StatusBadGateway, "backend request failed", "devrail_backend_error", "backend_request_failed")
}

Expand Down Expand Up @@ -341,6 +343,7 @@ type responseTelemetry struct {
TargetModel string
Backend string
UpstreamModel string
RouteRule string
Status int
Streaming bool
FirstEvent time.Time
Expand Down Expand Up @@ -460,6 +463,7 @@ func logTelemetry(telemetry *responseTelemetry) {
Alias: telemetry.Alias,
TargetModel: telemetry.TargetModel,
Backend: telemetry.Backend,
RouteRule: telemetry.RouteRule,
Status: telemetry.Status,
Streaming: telemetry.Streaming,
FirstEvent: telemetry.FirstEvent,
Expand Down Expand Up @@ -495,10 +499,13 @@ func instrumentBackendResponse(
model config.ModelConfig,
backend config.BackendConfig,
requestID string,
routeRule string,
metrics *metricsRegistry,
) {
if resp.Body == nil {
metrics.record(requestMetricsFromModel(model, backend, resp.StatusCode, started))
requestMetrics := requestMetricsFromModel(model, backend, resp.StatusCode, started)
requestMetrics.RouteRule = routeRule
metrics.record(requestMetrics)
return
}

Expand All @@ -507,6 +514,7 @@ func instrumentBackendResponse(
Alias: model.ID,
TargetModel: model.TargetModel,
Backend: backend.ID,
RouteRule: routeRule,
Status: resp.StatusCode,
Started: started,
QueueWait: queueWait,
Expand Down Expand Up @@ -574,6 +582,7 @@ type requestMetrics struct {
Alias string
TargetModel string
Backend string
RouteRule string
Status int
Streaming bool
FirstEvent time.Time
Expand Down Expand Up @@ -607,6 +616,7 @@ type metricLabels struct {
Alias string
Backend string
TargetModel string
RouteRule string
Status string
Streaming string
Le string
Expand Down Expand Up @@ -667,6 +677,7 @@ func (registry *metricsRegistry) record(metrics requestMetrics) {
Alias: metrics.Alias,
Backend: metrics.Backend,
TargetModel: metrics.TargetModel,
RouteRule: metrics.RouteRule,
Status: strconv.Itoa(metrics.Status),
Streaming: strconv.FormatBool(metrics.Streaming),
}
Expand Down Expand Up @@ -823,7 +834,7 @@ func sortedHistogramSeries(series map[string]*histogramSeries) []*histogramSerie
}

func (labels metricLabels) key() string {
return labels.Alias + "\xff" + labels.Backend + "\xff" + labels.TargetModel + "\xff" + labels.Status + "\xff" + labels.Streaming
return labels.Alias + "\xff" + labels.Backend + "\xff" + labels.TargetModel + "\xff" + labels.RouteRule + "\xff" + labels.Status + "\xff" + labels.Streaming
}

func (labels metricLabels) with(name, value string) metricLabels {
Expand All @@ -845,6 +856,9 @@ func (labels metricLabels) prometheus() string {
if labels.TargetModel != "" {
parts = append(parts, `target_model="`+escapePrometheusLabel(labels.TargetModel)+`"`)
}
if labels.RouteRule != "" {
parts = append(parts, `route_rule="`+escapePrometheusLabel(labels.RouteRule)+`"`)
}
if labels.Status != "" {
parts = append(parts, `status="`+escapePrometheusLabel(labels.Status)+`"`)
}
Expand Down
35 changes: 31 additions & 4 deletions internal/server/server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -169,6 +169,15 @@ func TestRoutingRuleSelectsTargetByPromptSize(t *testing.T) {
if backendModel != "deep-model" {
t.Fatalf("unexpected backend model: %q", backendModel)
}

metricsReq := httptest.NewRequest(http.MethodGet, "/metrics", nil)
metricsRec := httptest.NewRecorder()
srv.ServeHTTP(metricsRec, metricsReq)
body := metricsRec.Body.String()
want := `devrail_router_requests_total{alias="local-coder-auto",backend="lmstudio",target_model="deep-model",route_rule="large-prompt",status="200",streaming="false"} 1`
if !strings.Contains(body, want) {
t.Fatalf("expected metrics to contain %s, got:\n%s", want, body)
}
}

func TestRoutingRuleFallsBackToDefaultTarget(t *testing.T) {
Expand Down Expand Up @@ -298,6 +307,15 @@ func TestRoutingClassifierSelectsTarget(t *testing.T) {
if backendModel != "deep-model" {
t.Fatalf("unexpected backend model: %q", backendModel)
}

metricsReq := httptest.NewRequest(http.MethodGet, "/metrics", nil)
metricsRec := httptest.NewRecorder()
srv.ServeHTTP(metricsRec, metricsReq)
body := metricsRec.Body.String()
want := `devrail_router_requests_total{alias="local-coder-auto",backend="lmstudio",target_model="deep-model",route_rule="classifier",status="200",streaming="false"} 1`
if !strings.Contains(body, want) {
t.Fatalf("expected metrics to contain %s, got:\n%s", want, body)
}
}

func TestRoutingClassifierFallsBackToDefaultTarget(t *testing.T) {
Expand Down Expand Up @@ -449,6 +467,15 @@ func TestRoutingPreclassifierSelectsTargetBeforeClassifier(t *testing.T) {
if backendModel != "deep-model" {
t.Fatalf("unexpected backend model: %q", backendModel)
}

metricsReq := httptest.NewRequest(http.MethodGet, "/metrics", nil)
metricsRec := httptest.NewRecorder()
srv.ServeHTTP(metricsRec, metricsReq)
body := metricsRec.Body.String()
want := `devrail_router_requests_total{alias="local-coder-auto",backend="lmstudio",target_model="deep-model",route_rule="preclassifier",status="200",streaming="false"} 1`
if !strings.Contains(body, want) {
t.Fatalf("expected metrics to contain %s, got:\n%s", want, body)
}
}

func TestRoutingPreclassifierFallsThroughForNegatedKeyword(t *testing.T) {
Expand Down Expand Up @@ -981,8 +1008,8 @@ func TestMetricsEndpointExposesRequestTelemetry(t *testing.T) {
body := rec.Body.String()
for _, want := range []string{
"# TYPE devrail_router_requests_total counter",
`devrail_router_requests_total{alias="local-coder",backend="lmstudio",target_model="target-model",status="200",streaming="false"} 1`,
`devrail_router_request_duration_seconds_bucket{alias="local-coder",backend="lmstudio",target_model="target-model",status="200",streaming="false",le="+Inf"} 1`,
`devrail_router_requests_total{alias="local-coder",backend="lmstudio",target_model="target-model",route_rule="default",status="200",streaming="false"} 1`,
`devrail_router_request_duration_seconds_bucket{alias="local-coder",backend="lmstudio",target_model="target-model",route_rule="default",status="200",streaming="false",le="+Inf"} 1`,
"devrail_router_prompt_tokens_total 9",
"devrail_router_completion_tokens_total 3",
"devrail_router_total_tokens_total 12",
Expand Down Expand Up @@ -1019,8 +1046,8 @@ func TestMetricsEndpointExposesStreamingFirstEventLatency(t *testing.T) {
body := metricsRec.Body.String()

for _, want := range []string{
`devrail_router_requests_total{alias="local-coder",backend="lmstudio",target_model="target-model",status="200",streaming="true"} 1`,
`devrail_router_first_event_latency_seconds_bucket{alias="local-coder",backend="lmstudio",target_model="target-model",status="200",streaming="true",le="+Inf"} 1`,
`devrail_router_requests_total{alias="local-coder",backend="lmstudio",target_model="target-model",route_rule="default",status="200",streaming="true"} 1`,
`devrail_router_first_event_latency_seconds_bucket{alias="local-coder",backend="lmstudio",target_model="target-model",route_rule="default",status="200",streaming="true",le="+Inf"} 1`,
} {
if !strings.Contains(body, want) {
t.Fatalf("expected metrics to contain %s, got:\n%s", want, body)
Expand Down
Loading