package llm import ( "bufio" "context" "encoding/json" "fmt" "io" "strings" "sync" "time" "github.com/jackc/pgx/v5/pgxpool" ) // Manager manages multiple LLM providers and routes requests. type Manager struct { providers map[string]Provider fallback string pool *pgxpool.Pool cache map[string]providerCacheEntry cacheMu sync.RWMutex cacheTTL time.Duration } type providerCacheEntry struct { provider Provider updatedAt time.Time } func NewManager() *Manager { return &Manager{ providers: make(map[string]Provider), cache: make(map[string]providerCacheEntry), cacheTTL: 5 * time.Minute, // 缓存 5 分钟 } } // SetPool 设置数据库连接池,用于从数据库加载 providers func (m *Manager) SetPool(pool *pgxpool.Pool) { m.pool = pool } func (m *Manager) Register(name string, provider Provider) { m.providers[name] = provider if m.fallback == "" { m.fallback = name } } func (m *Manager) SetFallback(name string) { m.fallback = name } // LoadProvidersFromDB 从数据库加载所有激活的 providers func (m *Manager) LoadProvidersFromDB(ctx context.Context) error { if m.pool == nil { return fmt.Errorf("database pool not set") } rows, err := m.pool.Query(ctx, ` SELECT id, name, base_url, api_key_encrypted, models, config FROM model_providers WHERE is_active = true ORDER BY priority DESC, created_at `) if err != nil { return fmt.Errorf("query providers: %w", err) } defer rows.Close() for rows.Next() { var ( id, name, baseURL, apiKey string modelsJSON, configJSON []byte ) if err := rows.Scan(&id, &name, &baseURL, &apiKey, &modelsJSON, &configJSON); err != nil { return fmt.Errorf("scan provider: %w", err) } var config map[string]any if err := json.Unmarshal(configJSON, &config); err != nil { return fmt.Errorf("unmarshal config: %w", err) } // 根据 config.provider 类型创建对应的 provider providerType, _ := config["provider"].(string) if providerType == "" { providerType = "openai" // 默认 OpenAI 兼容 } var provider Provider switch providerType { case "openai": provider = NewOpenAIProvider(apiKey, baseURL, "") case "anthropic": provider = NewAnthropicProvider(apiKey, baseURL, "") default: return fmt.Errorf("unsupported provider type: %s", providerType) } // 使用 provider ID 作为 key m.Register(id, provider) // 缓存 m.cacheMu.Lock() m.cache[id] = providerCacheEntry{ provider: provider, updatedAt: time.Now(), } m.cacheMu.Unlock() } return rows.Err() } // GetActiveProvider 获取优先级最高的激活 provider func (m *Manager) GetActiveProvider(ctx context.Context) (Provider, error) { if m.pool == nil { // 如果没有数据库连接,使用注册的 fallback if p, ok := m.providers[m.fallback]; ok { return p, nil } return nil, fmt.Errorf("no active provider") } // 从数据库获取优先级最高的激活 provider var id, baseURL, apiKey string var configJSON []byte err := m.pool.QueryRow(ctx, ` SELECT id, base_url, api_key_encrypted, config FROM model_providers WHERE is_active = true ORDER BY priority DESC, created_at LIMIT 1 `).Scan(&id, &baseURL, &apiKey, &configJSON) if err != nil { return nil, fmt.Errorf("query active provider: %w", err) } // 检查缓存 m.cacheMu.RLock() if entry, ok := m.cache[id]; ok && time.Since(entry.updatedAt) < m.cacheTTL { m.cacheMu.RUnlock() return entry.provider, nil } m.cacheMu.RUnlock() // 创建新的 provider var config map[string]any if err := json.Unmarshal(configJSON, &config); err != nil { return nil, fmt.Errorf("unmarshal config: %w", err) } providerType, _ := config["provider"].(string) if providerType == "" { providerType = "openai" } var provider Provider switch providerType { case "openai": provider = NewOpenAIProvider(apiKey, baseURL, "") case "anthropic": provider = NewAnthropicProvider(apiKey, baseURL, "") default: return nil, fmt.Errorf("unsupported provider type: %s", providerType) } // 更新缓存 m.cacheMu.Lock() m.cache[id] = providerCacheEntry{ provider: provider, updatedAt: time.Now(), } m.cacheMu.Unlock() return provider, nil } func (m *Manager) GetProvider(name string) (Provider, error) { if p, ok := m.providers[name]; ok { return p, nil } if p, ok := m.providers[m.fallback]; ok { return p, nil } return nil, fmt.Errorf("no provider found: %s", name) } // Chat performs a blocking chat completion using the specified provider. func (m *Manager) Chat(ctx context.Context, provider Provider, req *ChatRequest) (*ChatResponse, error) { return provider.ChatCompletion(ctx, req) } // ChatStream performs a streaming chat and returns the raw SSE body. func (m *Manager) ChatStream(ctx context.Context, provider Provider, req *ChatRequest) (io.ReadCloser, error) { req.Stream = true return provider.ChatStream(ctx, req) } // StreamEvent represents a normalized SSE event for the frontend. type StreamEvent struct { Event string `json:"event"` Answer string `json:"answer,omitempty"` MessageID string `json:"message_id,omitempty"` Usage *Usage `json:"usage,omitempty"` } type Usage struct { PromptTokens int `json:"prompt_tokens"` CompletionTokens int `json:"completion_tokens"` TotalTokens int `json:"total_tokens"` Model string `json:"model"` } // TransformOpenAIStream reads an OpenAI SSE stream and writes normalized events to the writer. func TransformOpenAIStream(reader io.Reader, write func(event StreamEvent)) error { scanner := bufio.NewScanner(reader) scanner.Buffer(make([]byte, 64*1024), 256*1024) var totalContent string var model string for scanner.Scan() { line := scanner.Text() if !strings.HasPrefix(line, "data: ") { continue } data := strings.TrimPrefix(line, "data: ") if data == "[DONE]" { write(StreamEvent{ Event: "message_end", Usage: &Usage{ TotalTokens: estimateTokens(totalContent), Model: model, }, }) break } var chunk struct { ID string `json:"id"` Model string `json:"model"` Choices []struct { Delta struct { Content string `json:"content"` } `json:"delta"` FinishReason *string `json:"finish_reason"` } `json:"choices"` } if err := json.Unmarshal([]byte(data), &chunk); err != nil { continue } model = chunk.Model if len(chunk.Choices) > 0 && chunk.Choices[0].Delta.Content != "" { content := chunk.Choices[0].Delta.Content totalContent += content write(StreamEvent{ Event: "message", Answer: content, MessageID: chunk.ID, }) } } return scanner.Err() } // TransformAnthropicStream reads an Anthropic SSE stream and writes normalized events. func TransformAnthropicStream(reader io.Reader, write func(event StreamEvent)) error { scanner := bufio.NewScanner(reader) scanner.Buffer(make([]byte, 64*1024), 256*1024) var model string for scanner.Scan() { line := scanner.Text() if !strings.HasPrefix(line, "data: ") { if strings.HasPrefix(line, "event: ") { continue } continue } data := strings.TrimPrefix(line, "data: ") var event map[string]any if err := json.Unmarshal([]byte(data), &event); err != nil { continue } eventType, _ := event["type"].(string) switch eventType { case "message_start": if msg, ok := event["message"].(map[string]any); ok { if m, ok := msg["model"].(string); ok { model = m } } case "content_block_delta": if delta, ok := event["delta"].(map[string]any); ok { if text, ok := delta["text"].(string); ok { write(StreamEvent{ Event: "message", Answer: text, }) } } case "message_delta": if usage, ok := event["usage"].(map[string]any); ok { outputTokens := int(getFloat(usage, "output_tokens")) write(StreamEvent{ Event: "message_end", Usage: &Usage{ CompletionTokens: outputTokens, TotalTokens: outputTokens, Model: model, }, }) } } } return scanner.Err() } func getFloat(m map[string]any, key string) float64 { if v, ok := m[key].(float64); ok { return v } return 0 } func estimateTokens(text string) int { return len(text) / 4 }