人人都会AI编程

10.2 多模型切换配置

更新时间:2026-06-28

生产环境中通常需要同时接入多个模型提供商(如 GPT-4、Claude、本地 LLM),根据成本、延迟或任务类型动态切换。以下是基于配置文件的标准实现方案。

10.2.1 配置文件设计

config/models.yaml 中定义模型池:

models:
  gpt4:
    provider: openai
    model_name: gpt-4o
    api_key: ${OPENAI_API_KEY}
    base_url: https://api.openai.com/v1
    priority: 1
    timeout: 30
    
  claude:
    provider: anthropic
    model_name: claude-3-sonnet-20240229
    api_key: ${ANTHROPIC_API_KEY}
    max_tokens: 4096
    priority: 2
    
  local:
    provider: ollama
    model_name: llama3.1:8b
    base_url: http://localhost:11434
    priority: 3
    fallback: true   # 作为降级备选

default_model: gpt4

10.2.2 模型工厂实现

from typing import Dict, Optional
import yaml
from enum import Enum

class ProviderType(Enum):
    OPENAI = "openai"
    ANTHROPIC = "anthropic"
    OLLAMA = "ollama"

class ModelFactory:
    def __init__(self, config_path: str = "config/models.yaml"):
        with open(config_path) as f:
            self.config = yaml.safe_load(f)
        self._clients: Dict[str, object] = {}
    
    def get_client(self, model_id: Optional[str] = None):
        """获取指定模型客户端,未指定则使用默认模型"""
        model_id = model_id or self.config["default_model"]
        
        if model_id not in self._clients:
            cfg = self.config["models"][model_id]
            self._clients[model_id] = self._create_client(cfg)
            
        return self._clients[model_id], self.config["models"][model_id]
    
    def _create_client(self, cfg: dict):
        provider = ProviderType(cfg["provider"])
        
        if provider == ProviderType.OPENAI:
            from openai import OpenAI
            return OpenAI(
                api_key=cfg["api_key"],
                base_url=cfg.get("base_url"),
                timeout=cfg.get("timeout", 30)
            )
        elif provider == ProviderType.ANTHROPIC:
            from anthropic import Anthropic
            return Anthropic(api_key=cfg["api_key"])
        elif provider == ProviderType.OLLAMA:
            import ollama
            return ollama.Client(host=cfg["base_url"])
        
        raise ValueError(f"Unknown provider: {provider}")

10.2.3 业务层调用示例

class LLMService:
    def __init__(self):
        self.factory = ModelFactory()
    
    async def chat(self, messages: list, model_id: str = None, **kwargs):
        client, cfg = self.factory.get_client(model_id)
        
        try:
            if cfg["provider"] == "openai":
                resp = client.chat.completions.create(
                    model=cfg["model_name"],
                    messages=messages,
                    **kwargs
                )
                return resp.choices[0].message.content
                
            elif cfg["provider"] == "anthropic":
                resp = client.messages.create(
                    model=cfg["model_name"],
                    messages=messages,
                    max_tokens=cfg.get("max_tokens", 1024)
                )
                return resp.content[0].text
                
        except Exception as e:
            # 自动降级到 fallback 模型
            if not model_id and cfg.get("fallback") is not True:
                fallback_id = self._get_fallback_model()
                print(f"Model {cfg['model_name']} failed, fallback to {fallback_id}")
                return await self.chat(messages, model_id=fallback_id)
            raise
    
    def _get_fallback_model(self):
        """获取优先级最低且标记为 fallback 的模型"""
        models = self.factory.config["models"]
        return min(
            (k for k, v in models.items() if v.get("fallback")),
            key=lambda x: models[x]["priority"],
            default=None
        )

10.2.4 关键注意事项

  1. 密钥管理:生产环境务必使用 ${ENV_VAR} 语法注入,避免硬编码
  2. 连接复用:工厂模式缓存客户端实例,避免每次请求重复建立 HTTP 连接
  3. 超时控制:不同模型响应速度差异大,建议在配置中单独设置 timeout
  4. 上下文隔离:切换模型时需注意 token 计算方式不同,长对话可能需要重新计算上下文长度
  5. 熔断机制:建议配合 tenacity 库实现重试和熔断,防止某一家服务故障拖垮整体
# 使用示例
service = LLMService()

# 显式指定模型
result = await service.chat(
    messages=[{"role": "user", "content": "分析这份代码"}], 
    model_id="claude"
)

# 使用默认模型,失败自动降级到本地模型
result = await service.chat(messages)