diff --git a/configs/config.yaml b/configs/config.yaml index 879211c..36a5a05 100644 --- a/configs/config.yaml +++ b/configs/config.yaml @@ -36,7 +36,7 @@ context: models: general-chat: provider: ollama - actual_model: qwen2.5:0.5b + actual_model: deepseek-r1:1.5b endpoint: http://127.0.0.1:11434 context_window: 32768 max_output_tokens: 4096 @@ -46,7 +46,7 @@ models: fast-chat: provider: ollama - actual_model: qwen2.5:0.5b + actual_model: deepseek-r1:1.5b endpoint: http://127.0.0.1:11434 context_window: 16384 max_output_tokens: 2048 @@ -54,6 +54,16 @@ models: residency: on_demand idle_unload_seconds: 600 + vllm-chat: + provider: vllm + actual_model: deepseek-r1:1.5b + endpoint: http://127.0.0.1:8000 + context_window: 32768 + max_output_tokens: 4096 + max_concurrency: 4 + residency: always + cancel_supported: true + routing: sensitive_data_local_only: true allow_cloud_fallback_by_default: false diff --git a/internal/adapter/vllm.go b/internal/adapter/vllm.go new file mode 100644 index 0000000..cddceac --- /dev/null +++ b/internal/adapter/vllm.go @@ -0,0 +1,309 @@ +package adapter + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/edgeai/gateway/pkg/api" +) + +// VLLMAdapter implements ModelAdapter for vLLM inference engine. +// vLLM exposes an OpenAI-compatible API at /v1/chat/completions and /v1/models. +type VLLMAdapter struct { + endpoint string + httpClient *http.Client +} + +// NewVLLMAdapter creates a new vLLM adapter. +func NewVLLMAdapter(endpoint string) *VLLMAdapter { + return &VLLMAdapter{ + endpoint: strings.TrimRight(endpoint, "/"), + httpClient: &http.Client{ + Timeout: 120 * time.Second, + }, + } +} + +func (a *VLLMAdapter) Name() string { + return "vllm" +} + +// vllmChatRequest is the OpenAI-compatible chat request for vLLM. +type vllmChatRequest struct { + Model string `json:"model"` + Messages []vllmMsg `json:"messages"` + Stream bool `json:"stream"` + MaxTokens int `json:"max_tokens,omitempty"` + Temperature *float64 `json:"temperature,omitempty"` + TopP *float64 `json:"top_p,omitempty"` +} + +type vllmMsg struct { + Role string `json:"role"` + Content string `json:"content"` +} + +// vllmChatResponse is the OpenAI-compatible non-streaming response. +type vllmChatResponse struct { + ID string `json:"id"` + Model string `json:"model"` + Choices []struct { + Index int `json:"index"` + Message vllmMsg `json:"message"` + FinishReason string `json:"finish_reason"` + } `json:"choices"` + Usage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + } `json:"usage"` +} + +// vllmStreamChunk is a single SSE chunk in vLLM streaming response. +type vllmStreamChunk struct { + ID string `json:"id"` + Model string `json:"model"` + Choices []struct { + Index int `json:"index"` + Delta vllmMsg `json:"delta"` + FinishReason *string `json:"finish_reason"` + } `json:"choices"` + Usage *struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + } `json:"usage,omitempty"` +} + +func (a *VLLMAdapter) ChatCompletion(ctx context.Context, req *ChatRequest) (*ChatResponse, error) { + vllmReq := a.buildRequest(req, false) + + body, err := json.Marshal(vllmReq) + if err != nil { + return nil, fmt.Errorf("marshal vllm request: %w", err) + } + + httpReq, err := http.NewRequestWithContext(ctx, "POST", a.endpoint+"/v1/chat/completions", bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("create vllm request: %w", err) + } + httpReq.Header.Set("Content-Type", "application/json") + + resp, err := a.httpClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("vllm request failed: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + bodyBytes, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("vllm returned status %d: %s", resp.StatusCode, string(bodyBytes)) + } + + var vllmResp vllmChatResponse + if err := json.NewDecoder(resp.Body).Decode(&vllmResp); err != nil { + return nil, fmt.Errorf("decode vllm response: %w", err) + } + + content := "" + finishReason := "stop" + if len(vllmResp.Choices) > 0 { + content = vllmResp.Choices[0].Message.Content + finishReason = vllmResp.Choices[0].FinishReason + if finishReason == "" { + finishReason = "stop" + } + } + + return &ChatResponse{ + Content: content, + FinishReason: finishReason, + InputTokens: vllmResp.Usage.PromptTokens, + OutputTokens: vllmResp.Usage.CompletionTokens, + ActualModel: vllmResp.Model, + }, nil +} + +func (a *VLLMAdapter) ChatCompletionStream(ctx context.Context, req *ChatRequest) (<-chan StreamChunk, error) { + vllmReq := a.buildRequest(req, true) + + body, err := json.Marshal(vllmReq) + if err != nil { + return nil, fmt.Errorf("marshal vllm stream request: %w", err) + } + + httpReq, err := http.NewRequestWithContext(ctx, "POST", a.endpoint+"/v1/chat/completions", bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("create vllm stream request: %w", err) + } + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Accept", "text/event-stream") + + resp, err := a.httpClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("vllm stream request failed: %w", err) + } + + if resp.StatusCode != http.StatusOK { + bodyBytes, _ := io.ReadAll(resp.Body) + resp.Body.Close() + return nil, fmt.Errorf("vllm stream returned status %d: %s", resp.StatusCode, string(bodyBytes)) + } + + ch := make(chan StreamChunk, 100) + go func() { + defer close(ch) + defer resp.Body.Close() + + scanner := bufio.NewScanner(resp.Body) + scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) + + for scanner.Scan() { + line := scanner.Text() + if !strings.HasPrefix(line, "data: ") { + continue + } + data := strings.TrimPrefix(line, "data: ") + if data == "[DONE]" { + ch <- StreamChunk{Done: true, FinishReason: "stop"} + return + } + + // Check for cancellation + select { + case <-req.CancelCh: + ch <- StreamChunk{Done: true, FinishReason: "cancelled"} + return + default: + } + + var chunk vllmStreamChunk + if err := json.Unmarshal([]byte(data), &chunk); err != nil { + ch <- StreamChunk{Error: fmt.Errorf("decode stream chunk: %w", err)} + return + } + + if len(chunk.Choices) > 0 { + choice := chunk.Choices[0] + if choice.Delta.Content != "" { + ch <- StreamChunk{Delta: choice.Delta.Content} + } + if choice.FinishReason != nil { + inputTokens, outputTokens := 0, 0 + if chunk.Usage != nil { + inputTokens = chunk.Usage.PromptTokens + outputTokens = chunk.Usage.CompletionTokens + } + ch <- StreamChunk{ + Done: true, + FinishReason: *choice.FinishReason, + InputTokens: inputTokens, + OutputTokens: outputTokens, + } + return + } + } + } + + if err := scanner.Err(); err != nil { + ch <- StreamChunk{Error: fmt.Errorf("stream read error: %w", err)} + return + } + + ch <- StreamChunk{Done: true, FinishReason: "stop"} + }() + + return ch, nil +} + +func (a *VLLMAdapter) ListModels(ctx context.Context) ([]ModelInfo, error) { + httpReq, err := http.NewRequestWithContext(ctx, "GET", a.endpoint+"/v1/models", nil) + if err != nil { + return nil, fmt.Errorf("create list models request: %w", err) + } + + resp, err := a.httpClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("list models failed: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("list models returned status %d", resp.StatusCode) + } + + var modelsResp struct { + Data []struct { + ID string `json:"id"` + Context int `json:"max_model_len,omitempty"` + } `json:"data"` + } + if err := json.NewDecoder(resp.Body).Decode(&modelsResp); err != nil { + return nil, fmt.Errorf("decode models response: %w", err) + } + + models := make([]ModelInfo, len(modelsResp.Data)) + for i, m := range modelsResp.Data { + models[i] = ModelInfo{Name: m.ID, ContextWindow: m.Context} + } + return models, nil +} + +func (a *VLLMAdapter) HealthCheck(ctx context.Context) error { + httpReq, err := http.NewRequestWithContext(ctx, "GET", a.endpoint+"/v1/models", nil) + if err != nil { + return fmt.Errorf("create health check request: %w", err) + } + + resp, err := a.httpClient.Do(httpReq) + if err != nil { + return fmt.Errorf("health check failed: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("health check returned status %d", resp.StatusCode) + } + return nil +} + +func (a *VLLMAdapter) Cancel(requestID string) error { + // vLLM supports cancellation via context cancellation (closing HTTP connection). + return nil +} + +func (a *VLLMAdapter) buildRequest(req *ChatRequest, stream bool) vllmChatRequest { + msgs := make([]vllmMsg, len(req.Messages)) + for i, m := range req.Messages { + content, _ := m.Content.(string) + msgs[i] = vllmMsg{Role: m.Role, Content: content} + } + + vllmReq := vllmChatRequest{ + Model: req.Model, + Messages: msgs, + Stream: stream, + } + + if req.MaxTokens > 0 { + vllmReq.MaxTokens = req.MaxTokens + } + if req.Temperature != nil { + vllmReq.Temperature = req.Temperature + } + if req.TopP != nil { + vllmReq.TopP = req.TopP + } + + return vllmReq +} + +// Ensure vllmMsg satisfies the api.Message content interface when needed. +var _ = api.Message{} diff --git a/internal/middleware/middleware.go b/internal/middleware/middleware.go index 871130a..5aa286d 100644 --- a/internal/middleware/middleware.go +++ b/internal/middleware/middleware.go @@ -95,6 +95,12 @@ func (rw *responseWriter) WriteHeader(code int) { rw.ResponseWriter.WriteHeader(code) } +func (rw *responseWriter) Flush() { + if f, ok := rw.ResponseWriter.(http.Flusher); ok { + f.Flush() + } +} + // GetRequestID extracts the request ID from context. func GetRequestID(ctx context.Context) string { if v, ok := ctx.Value(RequestIDKey).(string); ok { diff --git a/internal/server/handlers.go b/internal/server/handlers.go index 3b6d061..9f8a9de 100644 --- a/internal/server/handlers.go +++ b/internal/server/handlers.go @@ -122,7 +122,7 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) { } } -func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, target *router.ModelTarget, logicalModel string) { +func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, _ *router.ModelTarget, logicalModel string) { sse := handler.NewSSEWriter(w) if sse == nil { s.scheduler.Complete(tk.ID) @@ -158,7 +158,7 @@ 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, target *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, _ *router.ModelTarget, logicalModel string) { ctx, cancel := context.WithTimeout(r.Context(), time.Duration(s.cfg.Timeouts.DefaultInferenceMs)*time.Millisecond) defer cancel() diff --git a/internal/server/server.go b/internal/server/server.go index b9bff5a..77d6791 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -68,6 +68,8 @@ func New(cfg *config.Config, logger *observability.Logger) (*Server, error) { switch mc.Provider { case "ollama": registry.Register(mc.Provider, adapter.NewOllamaAdapter(mc.Endpoint)) + case "vllm": + registry.Register(mc.Provider, adapter.NewVLLMAdapter(mc.Endpoint)) } registered[key] = true }