diff --git a/README.md b/README.md index 5336ce3..502124f 100644 --- a/README.md +++ b/README.md @@ -30,6 +30,11 @@ This repository is in early foundation work. The current service supports: Routing policy, auth, telemetry, LM Studio lifecycle integration, and Omarchy integration are planned next. +Model aliases can also set basic concurrency guardrails with +`max_concurrent_requests`, `max_queue_size`, and `queue_timeout`. This lets heavy +local models wait or reject predictably instead of allowing multiple agents to +dogpile the same backend. + ## Quick Start Build and test locally: @@ -98,6 +103,9 @@ models: context_window: 65536 max_output_tokens: 4096 tool_calls: true + max_concurrent_requests: 2 + max_queue_size: 4 + queue_timeout: 30s backends: - id: lmstudio diff --git a/configs/router.docker.yaml b/configs/router.docker.yaml index 5702751..dc88093 100644 --- a/configs/router.docker.yaml +++ b/configs/router.docker.yaml @@ -9,6 +9,9 @@ models: context_window: 65536 max_output_tokens: 4096 tool_calls: true + max_concurrent_requests: 2 + max_queue_size: 4 + queue_timeout: 30s - id: local-coder-large name: Local Coder Large backend: mock-openai @@ -16,6 +19,9 @@ models: context_window: 131072 max_output_tokens: 4096 tool_calls: true + max_concurrent_requests: 1 + max_queue_size: 2 + queue_timeout: 2m backends: - id: mock-openai diff --git a/configs/router.example.yaml b/configs/router.example.yaml index 6f4e6a1..65fef21 100644 --- a/configs/router.example.yaml +++ b/configs/router.example.yaml @@ -9,6 +9,9 @@ models: context_window: 65536 max_output_tokens: 4096 tool_calls: true + max_concurrent_requests: 2 + max_queue_size: 4 + queue_timeout: 30s - id: local-coder-large name: Local Coder Large backend: lmstudio @@ -16,6 +19,9 @@ models: context_window: 131072 max_output_tokens: 4096 tool_calls: true + max_concurrent_requests: 1 + max_queue_size: 2 + queue_timeout: 2m backends: - id: lmstudio diff --git a/docs/architecture.md b/docs/architecture.md index 55abc45..2c12a0c 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -37,3 +37,28 @@ The first router is deliberately simple: Future routing can add deterministic policy, queueing, health-aware selection, RouteLLM-style strong/weak model routing, and second-pass review workflows. + +## Request Limits + +Model aliases can define optional concurrency and queue limits: + +```yaml +models: + - id: local-coder-large + backend: lmstudio + target_model: qwen/qwen3.6-35b-a3b + max_concurrent_requests: 1 + max_queue_size: 2 + queue_timeout: 2m +``` + +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: +the slot is not released while tokens are still flowing. + +If all slots are busy, requests can wait in the bounded queue. If the queue is +full, DevRail Router returns an OpenAI-shaped `429` error. If the request waits +longer than `queue_timeout`, it returns `503`. Successful queued requests get an +`X-Devrail-Queue-Wait-Ms` response header, and queue decisions are logged with +active count, queued count, and wait time. diff --git a/internal/config/config.go b/internal/config/config.go index d75fd5e..89271ff 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -5,6 +5,7 @@ import ( "fmt" "net/url" "os" + "time" "gopkg.in/yaml.v3" ) @@ -22,13 +23,16 @@ type ServerConfig struct { } type ModelConfig struct { - ID string `yaml:"id"` - Name string `yaml:"name"` - Backend string `yaml:"backend"` - TargetModel string `yaml:"target_model"` - ContextWindow int `yaml:"context_window"` - MaxOutputTokens int `yaml:"max_output_tokens"` - ToolCalls bool `yaml:"tool_calls"` + ID string `yaml:"id"` + Name string `yaml:"name"` + Backend string `yaml:"backend"` + TargetModel string `yaml:"target_model"` + ContextWindow int `yaml:"context_window"` + MaxOutputTokens int `yaml:"max_output_tokens"` + ToolCalls bool `yaml:"tool_calls"` + MaxConcurrentRequests int `yaml:"max_concurrent_requests"` + MaxQueueSize int `yaml:"max_queue_size"` + QueueTimeout string `yaml:"queue_timeout"` } type BackendConfig struct { @@ -102,6 +106,15 @@ func (cfg Config) Validate() error { if model.TargetModel == "" { return fmt.Errorf("model %q target_model is required", model.ID) } + if model.MaxConcurrentRequests < 0 { + return fmt.Errorf("model %q max_concurrent_requests must be non-negative", model.ID) + } + if model.MaxQueueSize < 0 { + return fmt.Errorf("model %q max_queue_size must be non-negative", model.ID) + } + if _, err := model.QueueTimeoutDuration(); err != nil { + return fmt.Errorf("model %q queue_timeout is invalid: %w", model.ID, err) + } if _, ok := models[model.ID]; ok { return fmt.Errorf("model %q is duplicated", model.ID) } @@ -130,3 +143,19 @@ func (cfg Config) Backend(id string) (BackendConfig, bool) { return BackendConfig{}, false } + +func (model ModelConfig) QueueTimeoutDuration() (time.Duration, error) { + if model.QueueTimeout == "" { + return 0, nil + } + + duration, err := time.ParseDuration(model.QueueTimeout) + if err != nil { + return 0, err + } + if duration < 0 { + return 0, errors.New("duration must be non-negative") + } + + return duration, nil +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go index ee49c40..8fae1f8 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -3,7 +3,9 @@ package config import ( "os" "path/filepath" + "strings" "testing" + "time" ) func TestLoadValidConfig(t *testing.T) { @@ -37,6 +39,62 @@ backends: } } +func TestValidateQueueSettings(t *testing.T) { + t.Parallel() + + cfg := Config{ + Models: []ModelConfig{{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "qwen/qwen3.6-35b-a3b", + MaxConcurrentRequests: 1, + MaxQueueSize: 2, + QueueTimeout: "250ms", + }}, + Backends: []BackendConfig{{ + ID: "lmstudio", + BaseURL: "http://127.0.0.1:1234/v1", + }}, + } + + if err := cfg.Validate(); err != nil { + t.Fatalf("validate config: %v", err) + } + + duration, err := cfg.Models[0].QueueTimeoutDuration() + if err != nil { + t.Fatalf("parse queue timeout: %v", err) + } + if duration != 250*time.Millisecond { + t.Fatalf("unexpected queue timeout: %s", duration) + } +} + +func TestValidateInvalidQueueTimeout(t *testing.T) { + t.Parallel() + + cfg := Config{ + Models: []ModelConfig{{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "qwen/qwen3.6-35b-a3b", + QueueTimeout: "eventually", + }}, + 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(), "queue_timeout") { + t.Fatalf("expected queue_timeout error, got: %v", err) + } +} + func TestValidateUnknownBackend(t *testing.T) { t.Parallel() diff --git a/internal/server/limiter.go b/internal/server/limiter.go new file mode 100644 index 0000000..e8da5e0 --- /dev/null +++ b/internal/server/limiter.go @@ -0,0 +1,125 @@ +package server + +import ( + "context" + "errors" + "sync" + "time" + + "github.com/devrail-dev/devrail-router/internal/config" +) + +var ( + errQueueFull = errors.New("model queue is full") + errQueueTimeout = errors.New("timed out waiting for model queue") +) + +type modelLimiter struct { + modelID string + slots chan struct{} + maxQueueSize int + queueTimeout time.Duration + + mu sync.Mutex + active int + queued int +} + +type limiterSnapshot struct { + active int + queued int +} + +func newModelLimiter(model config.ModelConfig) (*modelLimiter, error) { + queueTimeout, err := model.QueueTimeoutDuration() + if err != nil { + return nil, err + } + if model.MaxConcurrentRequests <= 0 { + return nil, nil + } + + return &modelLimiter{ + modelID: model.ID, + slots: make(chan struct{}, model.MaxConcurrentRequests), + maxQueueSize: model.MaxQueueSize, + queueTimeout: queueTimeout, + }, nil +} + +func (limiter *modelLimiter) acquire(ctx context.Context) (time.Duration, limiterSnapshot, func(), error) { + started := time.Now() + + select { + case limiter.slots <- struct{}{}: + snapshot := limiter.incrementActive() + return 0, snapshot, limiter.release, nil + default: + } + + if !limiter.joinQueue() { + return 0, limiter.snapshot(), nil, errQueueFull + } + defer limiter.leaveQueue() + + waitCtx := ctx + cancel := func() {} + if limiter.queueTimeout > 0 { + waitCtx, cancel = context.WithTimeout(ctx, limiter.queueTimeout) + } + defer cancel() + + select { + case limiter.slots <- struct{}{}: + waited := time.Since(started) + snapshot := limiter.incrementActive() + return waited, snapshot, limiter.release, nil + case <-waitCtx.Done(): + if errors.Is(waitCtx.Err(), context.DeadlineExceeded) { + return time.Since(started), limiter.snapshot(), nil, errQueueTimeout + } + return time.Since(started), limiter.snapshot(), nil, waitCtx.Err() + } +} + +func (limiter *modelLimiter) joinQueue() bool { + limiter.mu.Lock() + defer limiter.mu.Unlock() + + if limiter.queued >= limiter.maxQueueSize { + return false + } + limiter.queued++ + return true +} + +func (limiter *modelLimiter) leaveQueue() { + limiter.mu.Lock() + defer limiter.mu.Unlock() + + limiter.queued-- +} + +func (limiter *modelLimiter) incrementActive() limiterSnapshot { + limiter.mu.Lock() + defer limiter.mu.Unlock() + + limiter.active++ + return limiterSnapshot{active: limiter.active, queued: limiter.queued} +} + +func (limiter *modelLimiter) release() { + <-limiter.slots + + limiter.mu.Lock() + defer limiter.mu.Unlock() + + limiter.active-- +} + +func (limiter *modelLimiter) snapshot() limiterSnapshot { + limiter.mu.Lock() + defer limiter.mu.Unlock() + + return limiterSnapshot{active: limiter.active, queued: limiter.queued} +} diff --git a/internal/server/server.go b/internal/server/server.go index 67b2653..4375f16 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -3,6 +3,7 @@ package server import ( "bytes" "encoding/json" + "errors" "fmt" "io" "log/slog" @@ -16,7 +17,8 @@ import ( ) type Server struct { - cfg config.Config + cfg config.Config + limiters map[string]*modelLimiter } func New(cfg config.Config) (*Server, error) { @@ -24,7 +26,18 @@ func New(cfg config.Config) (*Server, error) { return nil, err } - return &Server{cfg: cfg}, nil + limiters := make(map[string]*modelLimiter, len(cfg.Models)) + for _, model := range cfg.Models { + limiter, err := newModelLimiter(model) + if err != nil { + return nil, err + } + if limiter != nil { + limiters[model.ID] = limiter + } + } + + return &Server{cfg: cfg, limiters: limiters}, nil } func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { @@ -91,6 +104,14 @@ func (s *Server) proxyOpenAI(w http.ResponseWriter, r *http.Request) { return } + release, ok := s.acquireModelSlot(w, r, model) + if !ok { + return + } + if release != nil { + defer release() + } + body, err = rewriteModel(body, model.TargetModel) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) @@ -124,6 +145,45 @@ func (s *Server) proxyOpenAI(w http.ResponseWriter, r *http.Request) { proxy.ServeHTTP(w, r) } +func (s *Server) acquireModelSlot(w http.ResponseWriter, r *http.Request, model config.ModelConfig) (func(), bool) { + limiter, ok := s.limiters[model.ID] + if !ok { + return nil, true + } + + waited, snapshot, release, err := limiter.acquire(r.Context()) + if err == nil { + w.Header().Set("X-Devrail-Queue-Wait-Ms", fmt.Sprintf("%d", waited.Milliseconds())) + slog.Info( + "acquired model slot", + "alias", model.ID, + "active", snapshot.active, + "queued", snapshot.queued, + "wait_ms", waited.Milliseconds(), + ) + return release, true + } + + switch { + case errors.Is(err, errQueueFull): + writeOpenAIError(w, http.StatusTooManyRequests, "model queue is full", "devrail_queue_full", "queue_full") + case errors.Is(err, errQueueTimeout): + writeOpenAIError(w, http.StatusServiceUnavailable, "timed out waiting for model queue", "devrail_queue_timeout", "queue_timeout") + default: + writeOpenAIError(w, http.StatusRequestTimeout, "request canceled while waiting for model queue", "devrail_queue_canceled", "queue_canceled") + } + + slog.Warn( + "rejected queued request", + "alias", model.ID, + "active", snapshot.active, + "queued", snapshot.queued, + "wait_ms", waited.Milliseconds(), + "error", err, + ) + return nil, false +} + func requestModel(r *http.Request) (string, []byte, error) { if r.Body == nil { return "", nil, fmt.Errorf("request body is required") @@ -198,3 +258,13 @@ func writeJSON(w http.ResponseWriter, status int, value any) { slog.Error("write response", "error", err) } } + +func writeOpenAIError(w http.ResponseWriter, status int, message, errorType, code string) { + writeJSON(w, status, map[string]any{ + "error": map[string]string{ + "message": message, + "type": errorType, + "code": code, + }, + }) +} diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 4b2380f..99cbc8c 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -2,10 +2,13 @@ package server import ( "encoding/json" + "io" "net/http" "net/http/httptest" "strings" + "sync/atomic" "testing" + "time" "github.com/devrail-dev/devrail-router/internal/config" ) @@ -97,19 +100,189 @@ func TestJoinOpenAIPathAvoidsDuplicateVersionPrefix(t *testing.T) { } } +func TestModelLimiterQueuesRequests(t *testing.T) { + t.Parallel() + + backendStarted := make(chan struct{}) + releaseBackend := make(chan struct{}) + var backendStartedCount atomic.Int32 + var active atomic.Int32 + var maxActive atomic.Int32 + + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + current := active.Add(1) + for { + previous := maxActive.Load() + if current <= previous || maxActive.CompareAndSwap(previous, current) { + break + } + } + defer active.Add(-1) + + if backendStartedCount.Add(1) == 1 { + close(backendStarted) + <-releaseBackend + } + + var payload struct { + Model string `json:"model"` + } + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + t.Errorf("decode backend request: %v", err) + http.Error(w, "bad request", http.StatusBadRequest) + return + } + writeJSON(w, http.StatusOK, map[string]string{"model": payload.Model}) + })) + t.Cleanup(backend.Close) + + srv := testServerWithBackend(t, backend.URL+"/v1", config.ModelConfig{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "target-model", + MaxConcurrentRequests: 1, + MaxQueueSize: 1, + QueueTimeout: "1s", + }) + + firstDone := make(chan int, 1) + go func() { + firstDone <- serveChat(t, srv, "local-coder") + }() + + select { + case <-backendStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first backend request") + } + + secondDone := make(chan int, 1) + go func() { + secondDone <- serveChat(t, srv, "local-coder") + }() + + select { + case status := <-secondDone: + t.Fatalf("second request finished before slot was released with status %d", status) + case <-time.After(50 * time.Millisecond): + } + + close(releaseBackend) + + if status := <-firstDone; status != http.StatusOK { + t.Fatalf("unexpected first status: %d", status) + } + if status := <-secondDone; status != http.StatusOK { + t.Fatalf("unexpected second status: %d", status) + } + if maxActive.Load() != 1 { + t.Fatalf("backend saw %d concurrent requests, want 1", maxActive.Load()) + } +} + +func TestModelLimiterRejectsFullQueue(t *testing.T) { + t.Parallel() + + backendStarted := make(chan struct{}) + releaseBackend := make(chan struct{}) + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + close(backendStarted) + <-releaseBackend + writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) + })) + t.Cleanup(backend.Close) + + srv := testServerWithBackend(t, backend.URL, config.ModelConfig{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "target-model", + MaxConcurrentRequests: 1, + MaxQueueSize: 0, + QueueTimeout: "1s", + }) + + firstDone := make(chan int, 1) + go func() { + firstDone <- serveChat(t, srv, "local-coder") + }() + + select { + case <-backendStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first backend request") + } + + if status := serveChat(t, srv, "local-coder"); status != http.StatusTooManyRequests { + t.Fatalf("unexpected status: %d", status) + } + + close(releaseBackend) + if status := <-firstDone; status != http.StatusOK { + t.Fatalf("unexpected first status: %d", status) + } +} + +func TestModelLimiterTimesOutQueuedRequest(t *testing.T) { + t.Parallel() + + backendStarted := make(chan struct{}) + releaseBackend := make(chan struct{}) + backend := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + close(backendStarted) + <-releaseBackend + writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) + })) + t.Cleanup(backend.Close) + + srv := testServerWithBackend(t, backend.URL, config.ModelConfig{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "target-model", + MaxConcurrentRequests: 1, + MaxQueueSize: 1, + QueueTimeout: "25ms", + }) + + firstDone := make(chan int, 1) + go func() { + firstDone <- serveChat(t, srv, "local-coder") + }() + + select { + case <-backendStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first backend request") + } + + if status := serveChat(t, srv, "local-coder"); status != http.StatusServiceUnavailable { + t.Fatalf("unexpected status: %d", status) + } + + close(releaseBackend) + if status := <-firstDone; status != http.StatusOK { + t.Fatalf("unexpected first status: %d", status) + } +} + func testServer(t *testing.T) *Server { t.Helper() + return testServerWithBackend(t, "http://127.0.0.1:1234/v1", config.ModelConfig{ + ID: "local-coder", + Backend: "lmstudio", + TargetModel: "qwen/qwen3.6-35b-a3b", + }) +} + +func testServerWithBackend(t *testing.T, backendURL string, model config.ModelConfig) *Server { + t.Helper() + srv, err := New(config.Config{ Server: config.ServerConfig{Address: "127.0.0.1:0"}, - Models: []config.ModelConfig{{ - ID: "local-coder", - Backend: "lmstudio", - TargetModel: "qwen/qwen3.6-35b-a3b", - }}, + Models: []config.ModelConfig{model}, Backends: []config.BackendConfig{{ ID: "lmstudio", - BaseURL: "http://127.0.0.1:1234/v1", + BaseURL: backendURL, }}, }) if err != nil { @@ -118,3 +291,18 @@ func testServer(t *testing.T) *Server { return srv } + +func serveChat(t *testing.T, srv *Server, model string) int { + t.Helper() + + req := httptest.NewRequest( + http.MethodPost, + "/v1/chat/completions", + strings.NewReader(`{"model":"`+model+`","messages":[]}`), + ) + rec := httptest.NewRecorder() + + srv.ServeHTTP(rec, req) + _, _ = io.Copy(io.Discard, rec.Result().Body) + return rec.Code +} diff --git a/test/mock-openai-backend/main.go b/test/mock-openai-backend/main.go index 0eeecd2..89777a1 100644 --- a/test/mock-openai-backend/main.go +++ b/test/mock-openai-backend/main.go @@ -37,12 +37,16 @@ func main() { return } var payload struct { - Model string `json:"model"` + Model string `json:"model"` + DevRailMockDelayMS int `json:"devrail_mock_delay_ms"` } if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { http.Error(w, fmt.Sprintf("decode request: %v", err), http.StatusBadRequest) return } + if payload.DevRailMockDelayMS > 0 { + time.Sleep(time.Duration(payload.DevRailMockDelayMS) * time.Millisecond) + } writeJSON(w, http.StatusOK, map[string]any{ "id": "chatcmpl-mock", "object": "chat.completion", diff --git a/test/smoke/docker-compose.sh b/test/smoke/docker-compose.sh index 19f06f0..e6c6610 100755 --- a/test/smoke/docker-compose.sh +++ b/test/smoke/docker-compose.sh @@ -10,12 +10,16 @@ compose() { } docker_helper_dir="" +tmp_dir="" cleanup() { compose down --remove-orphans --volumes >/dev/null 2>&1 || true if [ -n "$docker_helper_dir" ]; then rm -rf "$docker_helper_dir" fi + if [ -n "$tmp_dir" ]; then + rm -rf "$tmp_dir" + fi } trap cleanup EXIT @@ -78,4 +82,37 @@ large_response=$( echo "$large_response" | grep -q '"model":"qwen/qwen3.6-35b-a3b"' echo "$large_response" | grep -q 'ok from qwen/qwen3.6-35b-a3b' +tmp_dir=$(mktemp -d) +delayed_body="$tmp_dir/delayed.json" +queued_headers="$tmp_dir/queued.headers" +queued_body="$tmp_dir/queued.json" + +curl -fsS "$router_url/v1/chat/completions" \ + -H 'Content-Type: application/json' \ + -d '{ + "model": "local-coder-large", + "messages": [{"role": "user", "content": "Hold the slot."}], + "max_tokens": 16, + "devrail_mock_delay_ms": 400 + }' >"$delayed_body" & +delayed_pid=$! + +sleep 0.05 + +curl -fsS "$router_url/v1/chat/completions" \ + -D "$queued_headers" \ + -H 'Content-Type: application/json' \ + -d '{ + "model": "local-coder-large", + "messages": [{"role": "user", "content": "Queue behind the slot."}], + "max_tokens": 16 + }' >"$queued_body" + +wait "$delayed_pid" + +queue_wait_ms=$(awk -F': ' 'tolower($1)=="x-devrail-queue-wait-ms" {gsub("\r", "", $2); print $2}' "$queued_headers") +test "${queue_wait_ms:-0}" -gt 0 +grep -q 'ok from qwen/qwen3.6-35b-a3b' "$queued_body" +grep -q 'ok from qwen/qwen3.6-35b-a3b' "$delayed_body" + echo "docker compose smoke passed"