From e7e98271d4141a22f3eca0b5f5b43dc41fb45354 Mon Sep 17 00:00:00 2001 From: freedakgmail Date: Mon, 3 Aug 2026 08:12:28 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=E8=BF=87=E8=BD=BD?= =?UTF-8?q?=E4=BF=9D=E6=8A=A4=E3=80=81=E8=83=8C=E5=8E=8B=E5=92=8C=E7=86=94?= =?UTF-8?q?=E6=96=AD=E6=9C=BA=E5=88=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/router/overload.go | 178 ++++++++++++++++++++++ internal/scheduler/backpressure.go | 92 ++++++++++++ internal/scheduler/circuit_breaker.go | 145 ++++++++++++++++++ internal/scheduler/scheduler.go | 34 +++-- internal/server/handlers.go | 206 +++++++++++++++++++++++++- internal/server/server.go | 65 +++++--- internal/task/store.go | 2 +- internal/task/task.go | 18 ++- 8 files changed, 704 insertions(+), 36 deletions(-) create mode 100644 internal/router/overload.go create mode 100644 internal/scheduler/backpressure.go create mode 100644 internal/scheduler/circuit_breaker.go diff --git a/internal/router/overload.go b/internal/router/overload.go new file mode 100644 index 0000000..4f3ee62 --- /dev/null +++ b/internal/router/overload.go @@ -0,0 +1,178 @@ +package router + +import ( + "fmt" + "strings" + + "github.com/edgeai/gateway/internal/config" +) + +// OverloadStrategy defines how to handle requests when the primary model is overloaded. +type OverloadStrategy int + +const ( + StrategySameModelOtherInstance OverloadStrategy = iota // try same model on another instance + StrategySmallerLocalModel // fall back to a smaller local model + StrategyBackupEdgeNode // route to a backup edge node + StrategyReject // reject the request +) + +// OverloadResolver resolves overload situations using configured strategies. +type OverloadResolver struct { + strategies []OverloadStrategy + modelMap *LogicalModelMapping +} + +// NewOverloadResolver creates an overload resolver from routing config. +func NewOverloadResolver(cfg *config.RoutingConfig, modelMap *LogicalModelMapping) *OverloadResolver { + r := &OverloadResolver{ + modelMap: modelMap, + } + for _, s := range cfg.OverloadStrategy { + switch strings.ToLower(s) { + case "same_model_other_instance": + r.strategies = append(r.strategies, StrategySameModelOtherInstance) + case "smaller_local_model": + r.strategies = append(r.strategies, StrategySmallerLocalModel) + case "backup_edge_node": + r.strategies = append(r.strategies, StrategyBackupEdgeNode) + case "reject": + r.strategies = append(r.strategies, StrategyReject) + } + } + // Default: reject if no strategies configured + if len(r.strategies) == 0 { + r.strategies = []OverloadStrategy{StrategyReject} + } + return r +} + +// OverloadResult contains the result of an overload resolution attempt. +type OverloadResult struct { + Strategy OverloadStrategy + Target *ModelTarget // resolved fallback target, nil if rejected + Rejected bool + Reason string +} + +// Resolve attempts to find a fallback target when the primary model is overloaded. +// currentModel is the logical model that is overloaded. +// excludeEndpoints is a set of endpoints already tried (to avoid loops). +func (r *OverloadResolver) Resolve(currentModel string, excludeEndpoints map[string]bool) OverloadResult { + for _, strategy := range r.strategies { + switch strategy { + case StrategySameModelOtherInstance: + // Find same actual model on a different endpoint + target := r.findSameModelDifferentEndpoint(currentModel, excludeEndpoints) + if target != nil { + return OverloadResult{ + Strategy: StrategySameModelOtherInstance, + Target: target, + } + } + case StrategySmallerLocalModel: + // Find a model with smaller context window (proxy for "smaller") + target := r.findSmallerModel(currentModel, excludeEndpoints) + if target != nil { + return OverloadResult{ + Strategy: StrategySmallerLocalModel, + Target: target, + } + } + case StrategyBackupEdgeNode: + // Find any available model not yet tried + target := r.findAnyAvailable(excludeEndpoints) + if target != nil { + return OverloadResult{ + Strategy: StrategyBackupEdgeNode, + Target: target, + } + } + case StrategyReject: + return OverloadResult{ + Strategy: StrategyReject, + Rejected: true, + Reason: "all overload strategies exhausted, request rejected", + } + } + } + return OverloadResult{ + Strategy: StrategyReject, + Rejected: true, + Reason: fmt.Sprintf("no fallback available for model %s", currentModel), + } +} + +func (r *OverloadResolver) findSameModelDifferentEndpoint(currentModel string, exclude map[string]bool) *ModelTarget { + r.modelMap.mu.RLock() + defer r.modelMap.mu.RUnlock() + + current, ok := r.modelMap.mapping[currentModel] + if !ok { + return nil + } + + for logical, target := range r.modelMap.mapping { + if logical == currentModel { + continue + } + if target.ActualModel == current.ActualModel && !exclude[target.Endpoint] { + return target + } + } + return nil +} + +func (r *OverloadResolver) findSmallerModel(currentModel string, exclude map[string]bool) *ModelTarget { + r.modelMap.mu.RLock() + defer r.modelMap.mu.RUnlock() + + current, ok := r.modelMap.mapping[currentModel] + if !ok { + return nil + } + + var best *ModelTarget + for logical, target := range r.modelMap.mapping { + if logical == currentModel { + continue + } + if exclude[target.Endpoint] { + continue + } + // Pick model with smaller context window as "smaller" proxy + if target.ContextWindow < current.ContextWindow { + if best == nil || target.ContextWindow > best.ContextWindow { + best = target + } + } + } + return best +} + +func (r *OverloadResolver) findAnyAvailable(exclude map[string]bool) *ModelTarget { + r.modelMap.mu.RLock() + defer r.modelMap.mu.RUnlock() + + for _, target := range r.modelMap.mapping { + if !exclude[target.Endpoint] { + return target + } + } + return nil +} + +// StrategyName returns a human-readable name for a strategy. +func StrategyName(s OverloadStrategy) string { + switch s { + case StrategySameModelOtherInstance: + return "same_model_other_instance" + case StrategySmallerLocalModel: + return "smaller_local_model" + case StrategyBackupEdgeNode: + return "backup_edge_node" + case StrategyReject: + return "reject" + } + return "unknown" +} diff --git a/internal/scheduler/backpressure.go b/internal/scheduler/backpressure.go new file mode 100644 index 0000000..8b59e00 --- /dev/null +++ b/internal/scheduler/backpressure.go @@ -0,0 +1,92 @@ +package scheduler + +import ( + "sync/atomic" +) + +// BackpressureManager tracks system load and provides admission control. +// Level 1 (70%): accept but log warning +// Level 2 (85%): accept only high priority (P0-P1) +// Level 3 (95%): reject all new requests +type BackpressureManager struct { + maxRunning int64 + maxQueued int64 + level1Threshold float64 + level2Threshold float64 + level3Threshold float64 + + runningCount int64 // atomic + queuedCount int64 // atomic +} + +// NewBackpressureManager creates a new backpressure manager. +func NewBackpressureManager(maxRunning, maxQueued int, l1, l2, l3 float64) *BackpressureManager { + return &BackpressureManager{ + maxRunning: int64(maxRunning), + maxQueued: int64(maxQueued), + level1Threshold: l1, + level2Threshold: l2, + level3Threshold: l3, + } +} + +// Update sets the current running and queued counts. +func (bp *BackpressureManager) Update(running, queued int) { + atomic.StoreInt64(&bp.runningCount, int64(running)) + atomic.StoreInt64(&bp.queuedCount, int64(queued)) +} + +// LoadRatio returns the current system load ratio (0.0 to 1.0). +func (bp *BackpressureManager) LoadRatio() float64 { + running := atomic.LoadInt64(&bp.runningCount) + queued := atomic.LoadInt64(&bp.queuedCount) + total := running + queued + capacity := bp.maxRunning + bp.maxQueued + if capacity == 0 { + return 0 + } + return float64(total) / float64(capacity) +} + +// Level returns the current backpressure level (0=normal, 1=warning, 2=restricted, 3=reject). +func (bp *BackpressureManager) Level() int { + load := bp.LoadRatio() + if load >= bp.level3Threshold { + return 3 + } + if load >= bp.level2Threshold { + return 2 + } + if load >= bp.level1Threshold { + return 1 + } + return 0 +} + +// ShouldAccept decides whether to accept a request based on priority and load. +// priority: 0=P0(highest) to 4=P4(lowest) +func (bp *BackpressureManager) ShouldAccept(priority int) bool { + level := bp.Level() + switch level { + case 0, 1: + return true + case 2: + // Only accept P0 and P1 + return priority <= 1 + case 3: + return false + } + return true +} + +// RejectReason returns a human-readable reason if the request should be rejected. +func (bp *BackpressureManager) RejectReason(priority int) string { + level := bp.Level() + if level == 3 { + return "system overloaded, please retry later" + } + if level == 2 && priority > 1 { + return "backpressure active, only high-priority requests accepted" + } + return "" +} diff --git a/internal/scheduler/circuit_breaker.go b/internal/scheduler/circuit_breaker.go new file mode 100644 index 0000000..efeaaa6 --- /dev/null +++ b/internal/scheduler/circuit_breaker.go @@ -0,0 +1,145 @@ +package scheduler + +import ( + "sync" + "time" +) + +// CircuitBreaker implements a sliding-window circuit breaker for adapter health. +type CircuitBreaker struct { + mu sync.Mutex + errorRateThreshold float64 + minRequests int + windowSeconds int + openDuration time.Duration + halfOpenMax int + + // sliding window state + requests []time.Time + errors []time.Time + + // breaker state + state breakerState + openedAt time.Time + halfOpenCount int +} + +type breakerState int + +const ( + breakerClosed breakerState = iota + breakerOpen + breakerHalfOpen +) + +// NewCircuitBreaker creates a new circuit breaker from config. +func NewCircuitBreaker(errorRate float64, minRequests, windowSec, openSec, halfOpenMax int) *CircuitBreaker { + return &CircuitBreaker{ + errorRateThreshold: errorRate, + minRequests: minRequests, + windowSeconds: windowSec, + openDuration: time.Duration(openSec) * time.Second, + halfOpenMax: halfOpenMax, + state: breakerClosed, + } +} + +// AllowRequest checks if a request should be allowed through. +func (cb *CircuitBreaker) AllowRequest() bool { + cb.mu.Lock() + defer cb.mu.Unlock() + + now := time.Now() + cb.prune(now) + + switch cb.state { + case breakerClosed: + return true + case breakerOpen: + if now.Sub(cb.openedAt) >= cb.openDuration { + cb.state = breakerHalfOpen + cb.halfOpenCount = 0 + return true + } + return false + case breakerHalfOpen: + if cb.halfOpenCount < cb.halfOpenMax { + cb.halfOpenCount++ + return true + } + return false + } + return true +} + +// RecordSuccess records a successful request. +func (cb *CircuitBreaker) RecordSuccess() { + cb.mu.Lock() + defer cb.mu.Unlock() + + now := time.Now() + cb.requests = append(cb.requests, now) + + if cb.state == breakerHalfOpen { + cb.state = breakerClosed + cb.requests = nil + cb.errors = nil + } +} + +// RecordError records a failed request and may trip the breaker. +func (cb *CircuitBreaker) RecordError() { + cb.mu.Lock() + defer cb.mu.Unlock() + + now := time.Now() + cb.requests = append(cb.requests, now) + cb.errors = append(cb.errors, now) + + if cb.state == breakerHalfOpen { + cb.state = breakerOpen + cb.openedAt = now + return + } + + if cb.state == breakerClosed && len(cb.requests) >= cb.minRequests { + errorRate := float64(len(cb.errors)) / float64(len(cb.requests)) + if errorRate >= cb.errorRateThreshold { + cb.state = breakerOpen + cb.openedAt = now + } + } +} + +// State returns the current breaker state name. +func (cb *CircuitBreaker) State() string { + cb.mu.Lock() + defer cb.mu.Unlock() + switch cb.state { + case breakerClosed: + return "closed" + case breakerOpen: + return "open" + case breakerHalfOpen: + return "half_open" + } + return "unknown" +} + +// prune removes entries outside the sliding window. +func (cb *CircuitBreaker) prune(now time.Time) { + cutoff := now.Add(-time.Duration(cb.windowSeconds) * time.Second) + cb.requests = pruneBefore(cb.requests, cutoff) + cb.errors = pruneBefore(cb.errors, cutoff) +} + +func pruneBefore(times []time.Time, cutoff time.Time) []time.Time { + idx := 0 + for idx < len(times) && times[idx].Before(cutoff) { + idx++ + } + if idx > 0 { + times = times[idx:] + } + return times +} diff --git a/internal/scheduler/scheduler.go b/internal/scheduler/scheduler.go index 24070ff..1207129 100644 --- a/internal/scheduler/scheduler.go +++ b/internal/scheduler/scheduler.go @@ -14,15 +14,15 @@ import ( // Scheduler manages task queuing and execution with priority-based scheduling. type Scheduler struct { - mu sync.Mutex - queue *priorityQueue - running map[string]*task.Task - maxRunning int - maxQueued int - notifyCh chan struct{} - logger *observability.Logger - ctx context.Context - cancel context.CancelFunc + mu sync.Mutex + queue *priorityQueue + running map[string]*task.Task + maxRunning int + maxQueued int + notifyCh chan struct{} + logger *observability.Logger + ctx context.Context + cancel context.CancelFunc } // NewScheduler creates a new scheduler. @@ -114,6 +114,22 @@ func (s *Scheduler) RunningCount() int { return len(s.running) } +// GetTask returns a running task by ID, or a queued task by ID. +func (s *Scheduler) GetTask(taskID string) (*task.Task, bool) { + s.mu.Lock() + defer s.mu.Unlock() + if t, ok := s.running[taskID]; ok { + return t, true + } + for i := 0; i < s.queue.Len(); i++ { + t := (*s.queue)[i] + if t.ID == taskID { + return t, true + } + } + return nil, false +} + // Stop shuts down the scheduler. func (s *Scheduler) Stop() { s.cancel() diff --git a/internal/server/handlers.go b/internal/server/handlers.go index 9f8a9de..3565d53 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -3,6 +3,7 @@ package server import ( "context" "encoding/json" + "fmt" "net/http" "strings" "time" @@ -10,10 +11,12 @@ import ( "github.com/edgeai/gateway/internal/adapter" "github.com/edgeai/gateway/internal/auth" "github.com/edgeai/gateway/internal/config" + ctxasm "github.com/edgeai/gateway/internal/context" "github.com/edgeai/gateway/internal/handler" "github.com/edgeai/gateway/internal/middleware" "github.com/edgeai/gateway/internal/observability" "github.com/edgeai/gateway/internal/router" + "github.com/edgeai/gateway/internal/session" "github.com/edgeai/gateway/internal/task" "github.com/edgeai/gateway/pkg/api" "github.com/google/uuid" @@ -70,9 +73,26 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) { priority = int(task.PriorityNormal) // downgrade to P2 if not allowed P0 } + // Backpressure check: reject low-priority requests under high load + s.backpressure.Update(s.scheduler.RunningCount(), s.scheduler.QueueLength()) + if !s.backpressure.ShouldAccept(priority) { + reason := s.backpressure.RejectReason(priority) + s.metrics.IncRequest("backpressure_rejected") + handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrQueueFull, reason, requestID)) + return + } + + // Circuit breaker check + if !s.breaker.AllowRequest() { + s.metrics.IncRequest("circuit_open") + handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, "circuit breaker open, please retry later", requestID)) + return + } + // Create task taskID := uuid.New().String() tk := task.NewTask(taskID, requestID, identity.AppID, identity.TenantID, req.Model, task.TaskPriority(priority), req.Stream) + tk.SessionID = req.SessionID // Submit to scheduler if err := s.scheduler.Submit(tk); err != nil { @@ -84,6 +104,7 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) { s.metrics.SetQueueLength(s.scheduler.QueueLength()) // Wait for task to be dequeued + queueStart := time.Now() ctx, cancel := context.WithTimeout(r.Context(), time.Duration(s.cfg.Timeouts.DefaultQueueMs)*time.Millisecond) defer cancel() @@ -91,19 +112,37 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) { if err != nil { s.scheduler.Complete(tk.ID) s.metrics.IncRequest("queue_timeout") + tk.Transition(task.StateTimedOut) handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrQueueTimeout, "queue timeout", requestID)) return } + // Record queue time + dequeued.QueueMs = int(time.Since(queueStart) / time.Millisecond) + // Transition to RUNNING dequeued.Transition(task.StateRunning) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) - // Build adapter request + // Build adapter request — assemble context if session_id provided + messages := req.Messages + if req.SessionID != "" { + sess, err := s.sessions.Get(req.SessionID) + if err == nil && sess != nil { + // Load session history and assemble context + history := s.loadSessionHistory(sess) + if len(history) > 0 { + assembler := ctxasm.NewAssembler(&s.cfg.Context) + result := assembler.Assemble(history, req.Messages, target.ContextWindow, target.MaxOutputTokens, req.ContextPolicy) + messages = result.Messages + } + } + } + adapterReq := &adapter.ChatRequest{ RequestID: requestID, Model: target.ActualModel, - Messages: req.Messages, + Messages: messages, MaxTokens: target.MaxOutputTokens, Temperature: req.Temperature, TopP: req.TopP, @@ -131,6 +170,7 @@ func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapter } tk.Transition(task.StateStreaming) + inferenceStart := time.Now() ch, err := adapterInst.ChatCompletionStream(r.Context(), req) if err != nil { @@ -146,11 +186,16 @@ func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapter s.logger.Error("streaming error", observability.F().Event("stream_error").TaskID(tk.ID).Reason(err.Error())) tk.Transition(task.StateFailed) s.metrics.IncTask("failed") + s.breaker.RecordError() } else { tk.Transition(task.StateCompleted) s.metrics.IncTask("completed") + s.breaker.RecordSuccess() } + tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond) + tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond) + s.metrics.AddTokens(inputTokens, outputTokens) s.scheduler.Complete(tk.ID) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) @@ -158,18 +203,46 @@ func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapter s.metrics.IncRequest("stream_ok") } -func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, _ *router.ModelTarget, logicalModel string) { +func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, target *router.ModelTarget, logicalModel string) { ctx, cancel := context.WithTimeout(r.Context(), time.Duration(s.cfg.Timeouts.DefaultInferenceMs)*time.Millisecond) defer cancel() + inferenceStart := time.Now() + resp, err := adapterInst.ChatCompletion(ctx, req) if err != nil { + // Try overload fallback strategies + fallbackResp, fbErr := s.tryOverloadFallback(r.Context(), req, target, logicalModel) + if fbErr == nil && fallbackResp != nil { + tk.Degraded = true + tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond) + tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond) + s.scheduler.Complete(tk.ID) + tk.Transition(task.StateCompleted) + s.metrics.IncTask("completed") + s.metrics.IncRequest("ok") + s.metrics.IncRequest("overload_fallback") + s.metrics.AddTokens(fallbackResp.InputTokens, fallbackResp.OutputTokens) + s.breaker.RecordSuccess() + s.scheduler.Complete(tk.ID) + s.metrics.SetRunningTasks(s.scheduler.RunningCount()) + s.metrics.SetQueueLength(s.scheduler.QueueLength()) + + chatResp := handler.BuildChatResponse(requestID, tk.ID, logicalModel, fallbackResp) + chatResp.Degraded = true + chatResp.Timing = s.buildTiming(tk) + handler.WriteJSON(w, http.StatusOK, chatResp) + return + } + s.scheduler.Complete(tk.ID) tk.Transition(task.StateFailed) s.metrics.IncTask("failed") s.metrics.IncRequest("error") + s.breaker.RecordError() if strings.Contains(err.Error(), "timeout") || ctx.Err() != nil { + tk.Transition(task.StateTimedOut) handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInferenceTimeout, "inference timeout", requestID)) } else { handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID)) @@ -177,18 +250,61 @@ func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adap return } + tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond) + tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond) + tk.Transition(task.StateCompleted) s.metrics.IncTask("completed") s.metrics.IncRequest("ok") s.metrics.AddTokens(resp.InputTokens, resp.OutputTokens) + s.breaker.RecordSuccess() s.scheduler.Complete(tk.ID) s.metrics.SetRunningTasks(s.scheduler.RunningCount()) s.metrics.SetQueueLength(s.scheduler.QueueLength()) chatResp := handler.BuildChatResponse(requestID, tk.ID, logicalModel, resp) + chatResp.Timing = s.buildTiming(tk) handler.WriteJSON(w, http.StatusOK, chatResp) } +// tryOverloadFallback attempts fallback strategies when the primary model fails. +func (s *Server) tryOverloadFallback(ctx context.Context, req *adapter.ChatRequest, target *router.ModelTarget, logicalModel string) (*adapter.ChatResponse, error) { + exclude := map[string]bool{target.Endpoint: true} + + result := s.overload.Resolve(logicalModel, exclude) + if result.Rejected { + return nil, fmt.Errorf("overload: %s", result.Reason) + } + + fallbackAdapter, err := s.registry.Get(result.Target.Provider) + if err != nil { + return nil, fmt.Errorf("overload: adapter not found: %w", err) + } + + fallbackReq := &adapter.ChatRequest{ + RequestID: req.RequestID, + Model: result.Target.ActualModel, + Messages: req.Messages, + MaxTokens: result.Target.MaxOutputTokens, + Temperature: req.Temperature, + TopP: req.TopP, + Stream: false, + CancelCh: req.CancelCh, + } + + s.logger.Info("overload fallback", + observability.F(). + Event("overload_fallback"). + Set("strategy", router.StrategyName(result.Strategy)). + Set("from", logicalModel). + Set("to", result.Target.LogicalModel)) + + fbCtx, fbCancel := context.WithTimeout(ctx, time.Duration(s.cfg.Timeouts.DefaultInferenceMs)*time.Millisecond) + defer fbCancel() + + return fallbackAdapter.ChatCompletion(fbCtx, fallbackReq) +} + func (s *Server) handleModels(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet { handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) @@ -282,3 +398,87 @@ func (s *Server) handleSessionByID(w http.ResponseWriter, r *http.Request) { handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) } } + +// loadSessionHistory returns the message history from a session. +func (s *Server) loadSessionHistory(sess *session.Session) []api.Message { + if sess == nil { + return nil + } + return sess.Messages +} + +// buildTiming constructs Timing metadata from a task. +func (s *Server) buildTiming(tk *task.Task) *api.Timing { + return &api.Timing{ + QueueMs: tk.QueueMs, + FirstTokenMs: tk.FirstTokenMs, + InferenceMs: tk.InferenceMs, + TotalMs: tk.TotalMs, + } +} + +func (s *Server) handleTasks(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) + return + } + handler.WriteJSON(w, http.StatusOK, map[string]any{"tasks": []any{}}) +} + +func (s *Server) handleTaskByID(w http.ResponseWriter, r *http.Request) { + requestID := middleware.GetRequestID(r.Context()) + taskID := strings.TrimPrefix(r.URL.Path, "/v1/tasks/") + + switch r.Method { + case http.MethodGet: + tk, ok := s.scheduler.GetTask(taskID) + if !ok { + handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "task not found", requestID)) + return + } + resp := map[string]any{ + "task_id": tk.ID, + "request_id": tk.RequestID, + "session_id": tk.SessionID, + "state": string(tk.GetState()), + "priority": int(tk.Priority), + "logical_model": tk.LogicalModel, + "stream": tk.Stream, + "created_at": tk.CreatedAt.Format(time.RFC3339), + "degraded": tk.Degraded, + } + if tk.StartedAt != nil { + resp["started_at"] = tk.StartedAt.Format(time.RFC3339) + } + if tk.CompletedAt != nil { + resp["completed_at"] = tk.CompletedAt.Format(time.RFC3339) + } + if tk.CancelReason != "" { + resp["cancel_reason"] = tk.CancelReason + } + if tk.ErrorMessage != "" { + resp["error_message"] = tk.ErrorMessage + } + resp["timing"] = s.buildTiming(tk) + handler.WriteJSON(w, http.StatusOK, resp) + + case http.MethodDelete: + tk, ok := s.scheduler.GetTask(taskID) + if !ok { + handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "task not found", requestID)) + return + } + if tk.IsTerminal() { + handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "task already in terminal state", requestID)) + return + } + tk.Cancel("client requested cancellation") + s.scheduler.Complete(tk.ID) + s.metrics.SetRunningTasks(s.scheduler.RunningCount()) + s.metrics.SetQueueLength(s.scheduler.QueueLength()) + handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "cancelled"}) + + default: + handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed")) + } +} diff --git a/internal/server/server.go b/internal/server/server.go index 77d6791..5c78fca 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -22,15 +22,18 @@ import ( // Server is the main HTTP server for the AI gateway. type Server struct { - cfg *config.Config - logger *observability.Logger - metrics *observability.Metrics - HTTPSrv *http.Server - auth *auth.Authenticator - registry *adapter.Registry - modelMap *router.LogicalModelMapping - scheduler *scheduler.Scheduler - sessions *session.Store + cfg *config.Config + logger *observability.Logger + metrics *observability.Metrics + HTTPSrv *http.Server + auth *auth.Authenticator + registry *adapter.Registry + modelMap *router.LogicalModelMapping + scheduler *scheduler.Scheduler + sessions *session.Store + breaker *scheduler.CircuitBreaker + backpressure *scheduler.BackpressureManager + overload *router.OverloadResolver } // New creates a new Server instance with all components wired. @@ -60,6 +63,9 @@ func New(cfg *config.Config, logger *observability.Logger) (*Server, error) { // Initialize logical model mapping modelMap := router.NewLogicalModelMapping(cfg) + // Initialize overload resolver + overloadResolver := router.NewOverloadResolver(&cfg.Routing, modelMap) + // Register adapters for each unique endpoint registered := make(map[string]bool) for _, mc := range cfg.Models { @@ -78,18 +84,39 @@ func New(cfg *config.Config, logger *observability.Logger) (*Server, error) { // Initialize scheduler sched := scheduler.NewScheduler(&cfg.Scheduler, logger) + // Initialize circuit breaker + breaker := scheduler.NewCircuitBreaker( + cfg.CircuitBreaker.ErrorRateThreshold, + cfg.CircuitBreaker.MinRequests, + cfg.CircuitBreaker.WindowSeconds, + cfg.CircuitBreaker.OpenDurationSeconds, + cfg.CircuitBreaker.HalfOpenMaxRequests, + ) + + // Initialize backpressure manager + bp := scheduler.NewBackpressureManager( + cfg.Scheduler.MaxRunningTasks, + cfg.Scheduler.MaxQueuedTasks, + cfg.Backpressure.Level1Threshold, + cfg.Backpressure.Level2Threshold, + cfg.Backpressure.Level3Threshold, + ) + // Initialize metrics metrics := observability.NewMetrics() s := &Server{ - cfg: cfg, - logger: logger, - metrics: metrics, - auth: authenticator, - registry: registry, - modelMap: modelMap, - scheduler: sched, - sessions: sessionStore, + cfg: cfg, + logger: logger, + metrics: metrics, + auth: authenticator, + registry: registry, + modelMap: modelMap, + scheduler: sched, + sessions: sessionStore, + breaker: breaker, + backpressure: bp, + overload: overloadResolver, } mux := http.NewServeMux() @@ -128,6 +155,10 @@ func (s *Server) registerRoutes(mux *http.ServeMux) { // Session management mux.HandleFunc("/v1/sessions", s.handleSessions) mux.HandleFunc("/v1/sessions/", s.handleSessionByID) + + // Task management + mux.HandleFunc("/v1/tasks", s.handleTasks) + mux.HandleFunc("/v1/tasks/", s.handleTaskByID) } // Authenticator returns the authenticator instance (for testing/management). diff --git a/internal/task/store.go b/internal/task/store.go index 47a241e..c1f78db 100644 --- a/internal/task/store.go +++ b/internal/task/store.go @@ -157,7 +157,7 @@ func (s *Store) RecoverPendingTasks() (int, error) { defer s.mu.Unlock() result, err := s.db.Exec( - `UPDATE tasks SET state = 'FAILED', error_message = 'gateway restart' WHERE state IN ('RUNNING', 'STREAMING')`) + `UPDATE tasks SET state = 'FAILED', error_message = 'gateway restart' WHERE state IN ('RUNNING', 'STREAMING', 'QUEUED')`) if err != nil { return 0, err } diff --git a/internal/task/task.go b/internal/task/task.go index e94a66c..3865cc2 100644 --- a/internal/task/task.go +++ b/internal/task/task.go @@ -19,6 +19,7 @@ const ( StateCompleted TaskState = "COMPLETED" StateFailed TaskState = "FAILED" StateCancelled TaskState = "CANCELLED" + StateTimedOut TaskState = "TIMED_OUT" ) // TaskPriority levels (P0 highest, P4 lowest). @@ -53,6 +54,10 @@ type Task struct { OutputTokens int NodeID string Degraded bool + QueueMs int + FirstTokenMs int + InferenceMs int + TotalMs int cancelCh chan struct{} cancelOnce sync.Once mu sync.RWMutex @@ -76,12 +81,13 @@ func NewTask(id, requestID, appID, tenantID, logicalModel string, priority TaskP // AllowedTransitions defines valid state transitions. var allowedTransitions = map[TaskState][]TaskState{ - StateQueued: {StateRunning, StateFailed, StateCancelled}, - StateRunning: {StateStreaming, StateCompleted, StateFailed, StateCancelled}, - StateStreaming: {StateCompleted, StateFailed, StateCancelled}, + StateQueued: {StateRunning, StateFailed, StateCancelled, StateTimedOut}, + StateRunning: {StateStreaming, StateCompleted, StateFailed, StateCancelled, StateTimedOut}, + StateStreaming: {StateCompleted, StateFailed, StateCancelled, StateTimedOut}, StateCompleted: {}, StateFailed: {}, StateCancelled: {}, + StateTimedOut: {}, } // Transition changes the task state if the transition is valid. @@ -112,7 +118,7 @@ func (t *Task) Transition(to TaskState) error { switch to { case StateRunning: t.StartedAt = &now - case StateCompleted, StateFailed, StateCancelled: + case StateCompleted, StateFailed, StateCancelled, StateTimedOut: t.CompletedAt = &now } @@ -129,7 +135,7 @@ func (t *Task) Cancel(reason string) error { t.mu.Lock() defer t.mu.Unlock() - if t.State == StateCompleted || t.State == StateFailed || t.State == StateCancelled { + if t.State == StateCompleted || t.State == StateFailed || t.State == StateCancelled || t.State == StateTimedOut { return errors.New("task already in terminal state") } @@ -165,7 +171,7 @@ func (t *Task) GetState() TaskState { // IsTerminal returns true if the task is in a terminal state. func (t *Task) IsTerminal() bool { s := t.GetState() - return s == StateCompleted || s == StateFailed || s == StateCancelled + return s == StateCompleted || s == StateFailed || s == StateCancelled || s == StateTimedOut } // StateMachineLogger logs state transitions.