diff --git a/CHANGELOG.md b/CHANGELOG.md index ab96344..25c15e4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -19,6 +19,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 labels so prompt-size guardrail behavior can be tuned from real traffic. - `/metrics` now exposes `devrail_router_inflight_requests` so open streams can be distinguished from completed or stalled client-side requests. +- Model aliases can define `max_prompt_chars` to reject oversized prompts before + queueing, readiness hooks, or backend proxying. Rejections return an + OpenAI-shaped `context_length_exceeded` error with compact-context headers. - Opt-in command-backed model profile ensure hooks. - Roadmap for maturing DevRail Router from a single-backend gateway into an observable local inference control plane. diff --git a/README.md b/README.md index 2a0d406..832c839 100644 --- a/README.md +++ b/README.md @@ -128,6 +128,7 @@ models: target_model: qwen3-coder-30b-a3b-instruct context_window: 65536 max_output_tokens: 4096 + max_prompt_chars: 200000 tool_calls: true max_concurrent_requests: 2 max_queue_size: 4 diff --git a/configs/router.docker.yaml b/configs/router.docker.yaml index dc88093..fb201dc 100644 --- a/configs/router.docker.yaml +++ b/configs/router.docker.yaml @@ -8,6 +8,7 @@ models: target_model: qwen3-coder-30b-a3b-instruct context_window: 65536 max_output_tokens: 4096 + max_prompt_chars: 200000 tool_calls: true max_concurrent_requests: 2 max_queue_size: 4 diff --git a/configs/router.example.yaml b/configs/router.example.yaml index cf58f90..5bf5785 100644 --- a/configs/router.example.yaml +++ b/configs/router.example.yaml @@ -8,6 +8,7 @@ models: target_model: qwen3-coder-30b-a3b-instruct context_window: 65536 max_output_tokens: 4096 + max_prompt_chars: 200000 tool_calls: true max_concurrent_requests: 2 max_queue_size: 4 diff --git a/docs/architecture.md b/docs/architecture.md index dbf6b91..460575a 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -79,6 +79,24 @@ models: queue_timeout: 2m ``` +Aliases can also define `max_prompt_chars` to reject oversized requests before +queue acquisition, readiness hooks, or backend proxying: + +```yaml +models: + - id: local-coder + backend: lmstudio + target_model: qwen/qwen3.6-35b-a3b + max_prompt_chars: 200000 +``` + +Oversized prompt rejections return an OpenAI-shaped `400` response with +`error.code` set to `context_length_exceeded`. The response includes +`X-Devrail-Action: compact_context`, `X-Devrail-Prompt-Chars`, and +`X-Devrail-Max-Prompt-Chars` headers so coding-agent clients can compact or +trim context before retrying. Request metrics record these rejections with +`route_rule="prompt-limit"` and `status="400"`. + When `max_concurrent_requests` is unset or `0`, the alias is unlimited. When it is set, DevRail Router holds one slot for each proxied request until the upstream response is fully complete. That matters for streaming chat responses: diff --git a/internal/config/config.go b/internal/config/config.go index 248b003..10f58a2 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -30,6 +30,7 @@ type ModelConfig struct { TargetModel string `yaml:"target_model"` ContextWindow int `yaml:"context_window"` MaxOutputTokens int `yaml:"max_output_tokens"` + MaxPromptChars int `yaml:"max_prompt_chars"` ToolCalls bool `yaml:"tool_calls"` MaxConcurrentRequests int `yaml:"max_concurrent_requests"` MaxQueueSize int `yaml:"max_queue_size"` @@ -182,6 +183,9 @@ func (cfg Config) Validate() error { if model.MaxConcurrentRequests < 0 { return fmt.Errorf("model %q max_concurrent_requests must be non-negative", model.ID) } + if model.MaxPromptChars < 0 { + return fmt.Errorf("model %q max_prompt_chars must be non-negative", model.ID) + } if model.MaxQueueSize < 0 { return fmt.Errorf("model %q max_queue_size must be non-negative", model.ID) } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index c59a783..944b9a0 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -70,6 +70,52 @@ func TestValidateQueueSettings(t *testing.T) { } } +func TestValidateMaxPromptChars(t *testing.T) { + t.Parallel() + + cfg := Config{ + Models: []ModelConfig{{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "qwen/qwen3.6-35b-a3b", + MaxPromptChars: 120000, + }}, + Backends: []BackendConfig{{ + ID: "lmstudio", + BaseURL: "http://127.0.0.1:1234/v1", + }}, + } + + if err := cfg.Validate(); err != nil { + t.Fatalf("validate config: %v", err) + } +} + +func TestValidateRejectsNegativeMaxPromptChars(t *testing.T) { + t.Parallel() + + cfg := Config{ + Models: []ModelConfig{{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "qwen/qwen3.6-35b-a3b", + MaxPromptChars: -1, + }}, + Backends: []BackendConfig{{ + ID: "lmstudio", + BaseURL: "http://127.0.0.1:1234/v1", + }}, + } + + err := cfg.Validate() + if err == nil { + t.Fatal("expected validation error") + } + if !strings.Contains(err.Error(), "max_prompt_chars") { + t.Fatalf("expected max_prompt_chars error, got: %v", err) + } +} + func TestValidateRoutingRules(t *testing.T) { t.Parallel() diff --git a/internal/server/server.go b/internal/server/server.go index 31dd543..a212698 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -89,6 +89,7 @@ func (s *Server) handleModels(w http.ResponseWriter, _ *http.Request) { Name string `json:"name,omitempty"` ContextWindow int `json:"context_window,omitempty"` MaxOutput int `json:"max_output_tokens,omitempty"` + MaxPrompt int `json:"max_prompt_chars,omitempty"` ToolCall bool `json:"tool_call,omitempty"` TargetModel string `json:"target_model,omitempty"` } @@ -102,6 +103,7 @@ func (s *Server) handleModels(w http.ResponseWriter, _ *http.Request) { Name: model.Name, ContextWindow: model.ContextWindow, MaxOutput: model.MaxOutputTokens, + MaxPrompt: model.MaxPromptChars, ToolCall: model.ToolCalls, TargetModel: model.TargetModel, }) @@ -134,6 +136,21 @@ func (s *Server) proxyOpenAI(w http.ResponseWriter, r *http.Request) { } features := extractRoutingFeatures(body) + if rejectOversizedPrompt(w, model, features) { + metrics := requestMetricsFromModel(model, config.BackendConfig{}, http.StatusBadRequest, time.Now()) + metrics.RouteRule = "prompt-limit" + metrics.applyRoutingFeatures(features) + s.metrics.record(metrics) + slog.Warn( + "rejected oversized prompt", + "request_id", requestID, + "alias", model.ID, + "prompt_chars", features.PromptChars, + "max_prompt_chars", model.MaxPromptChars, + ) + return + } + backend, ok := s.cfg.Backend(model.Backend) if !ok { writeOpenAIError(w, http.StatusInternalServerError, fmt.Sprintf("unknown backend %q", model.Backend), "devrail_config_error", "unknown_backend") @@ -224,6 +241,24 @@ func (s *Server) proxyOpenAI(w http.ResponseWriter, r *http.Request) { proxy.ServeHTTP(w, r) } +func rejectOversizedPrompt(w http.ResponseWriter, model config.ModelConfig, features routingFeatures) bool { + if model.MaxPromptChars <= 0 || !features.Valid || features.PromptChars <= model.MaxPromptChars { + return false + } + + w.Header().Set("X-Devrail-Prompt-Chars", fmt.Sprintf("%d", features.PromptChars)) + w.Header().Set("X-Devrail-Max-Prompt-Chars", fmt.Sprintf("%d", model.MaxPromptChars)) + w.Header().Set("X-Devrail-Action", "compact_context") + writeOpenAIError( + w, + http.StatusBadRequest, + fmt.Sprintf("context length exceeded: prompt has %d characters, limit is %d for model alias %q; compact or trim context and retry", features.PromptChars, model.MaxPromptChars, model.ID), + "invalid_request_error", + "context_length_exceeded", + ) + return true +} + func (s *Server) acquireModelSlot(w http.ResponseWriter, r *http.Request, model config.ModelConfig, backend config.BackendConfig, requestID string, features routingFeatures) (time.Duration, func(), bool) { limiter, ok := s.limiters[model.ID] if !ok { diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 913bdc3..14df74d 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -180,6 +180,55 @@ func TestRoutingRuleSelectsTargetByPromptSize(t *testing.T) { } } +func TestMaxPromptCharsRejectsBeforeBackend(t *testing.T) { + t.Parallel() + + var backendHit atomic.Bool + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + backendHit.Store(true) + writeJSON(w, http.StatusOK, map[string]string{"model": "target-model"}) + })) + t.Cleanup(backend.Close) + + srv := testServerWithBackend(t, backend.URL, config.ModelConfig{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "target-model", + MaxPromptChars: 16, + }) + req := httptest.NewRequest( + http.MethodPost, + "/v1/chat/completions", + strings.NewReader(`{"model":"local-coder","messages":[{"role":"user","content":"please analyze this oversized prompt"}]}`), + ) + rec := httptest.NewRecorder() + + srv.ServeHTTP(rec, req) + + if rec.Code != http.StatusBadRequest { + t.Fatalf("unexpected status: %d", rec.Code) + } + assertOpenAIErrorCode(t, rec.Body.Bytes(), "context_length_exceeded") + if backendHit.Load() { + t.Fatal("backend should not receive oversized prompt") + } + if got := rec.Header().Get("X-Devrail-Action"); got != "compact_context" { + t.Fatalf("unexpected action header: %q", got) + } + if got := rec.Header().Get("X-Devrail-Max-Prompt-Chars"); got != "16" { + t.Fatalf("unexpected prompt limit header: %q", got) + } + + 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",target_model="target-model",route_rule="prompt-limit",status="400",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) { t.Parallel()