Files
selfrelease da9c8334d8
CI / lint (push) Has been cancelled
CI / test (push) Has been cancelled
CI / build (push) Has been cancelled
CI / security-scan (push) Has been cancelled
feat: 十轮网关优化 - 安全加固/可观测性/性能/可靠性
- SSE Keepalive Ping (15s心跳防止代理断连)
- Timing HTTP 头 (X-Timing-Queue/Inference/Total-Ms)
- Adapter Request-ID 传播到后端
- Session 清理日志回调
- Server 安全加固 (ReadHeaderTimeout/MaxHeaderBytes 防 slowloris)
- Usage Tracker 数据保留清理 (retentionDays + 定期清理)
- Config Reload 后 Adapter Registry 更新 (RegisterIfAbsent + RWMutex)
- Rate Limiter 空闲 Bucket 清理 (30分钟过期)
- Shutdown Drain 超时可配置 (ShutdownDrainSeconds)
- Config 模型字段校验增强 (provider/endpoint/actual_model)
- Auth 过期 Key 自动清理 (5分钟扫描)
- Admin API Rate Limiting
- Adapter Health Check 独立超时 (每个 adapter 3s)
- TCP 连接阶段超时 (DialContext 5s + KeepAlive 30s)
- 幂等键缓存、审计日志、Gzip 中间件、CORS Expose Headers
- Backpressure 响应头、熔断器 Prometheus 指标
- 连接池优化、Trace-ID 全链路传播
2026-08-03 15:43:11 +08:00

314 lines
8.2 KiB
Go

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: newHTTPClient(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")
if req.RequestID != "" {
httpReq.Header.Set("X-Request-ID", req.RequestID)
}
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")
if req.RequestID != "" {
httpReq.Header.Set("X-Request-ID", req.RequestID)
}
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{}