LLM Gateway & Routing: Маршрутизация запросов к LLM

llm-gatewayroutingload-balancingfallbackmulti-model
← Back to Blog

Зачем нужен LLM Gateway?

LLM Gateway — это прокси-слой между вашими приложениями и LLM-серверами. Он решает ключевые проблемы:

  • Multi-model routing: Маршрутизация запросов к разным моделям в зависимости от задачи
  • Load balancing: Распределение нагрузки между несколькими инстансами
  • Fallback: Автоматическое переключение на резервную модель при сбое
  • Rate limiting: Ограничение запросов для защиты от перегрузки
  • Cost optimization: Выбор оптимальной модели по стоимости/качеству
  • Observability: Единая точка для логирования и мониторинга

Архитектура Gateway

                    ┌─────────────┐
  App Request ────► │             │
                    │   Gateway   │
                    │             │
                    └──┬──┬──┬───┬┘
                       │  │  │   │
              ┌────────┘  │  │   └────────┐
              ▼           ▼  ▼            ▼
        ┌─────────┐ ┌────────┐ ┌────────┐
        │ Model A │ │Model B │ │Model C │
        │ (Fast)  │ │ (Smart)│ │(Cheap) │
        └─────────┘ └────────┘ └────────┘
        
        Routes:
        - Simple QA → Model A
        - Complex reasoning → Model B
        - Batch processing → Model C

Базовый Gateway на FastAPI

from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import StreamingResponse, JSONResponse
from typing import Optional, List, Dict, Any
import asyncio
import time
import json
import logging
from dataclasses import dataclass, field, asdict
from enum import Enum
from datetime import datetime, timedelta
import aiohttp
import uuid
from uuid import uuid4
import random

app = FastAPI(title="LLM Gateway", version="1.0.0")

class ModelProvider(str, Enum):
    OLLAMA = "ollama"
    VLLM = "vllm"
    TGI = "tgi"
    LM_STUDIO = "lm-studio"
    OPENAI_COMPATIBLE = "openai-compatible"

class ModelRoute(str, Enum):
    FAST = "fast"           # Быстрые, простые задачи
    SMART = "smart"         # Сложное рассуждение
    CODE = "code"           # Генерация кода
    CHAT = "chat"           # Обычный чат
    DEFAULT = "default"     # По умолчанию

@dataclass
class ModelEndpoint:
    """Эндпоинт модели"""
    name: str
    provider: ModelProvider
    url: str
    model_id: str
    route: ModelRoute
    max_tokens: int = 4096
    temperature: float = 0.7
    timeout: float = 60.0
    health: str = "healthy"  # healthy, degraded, unhealthy
    weight: int = 1  # Для load balancing
    rate_limit: int = 100  # Запросов в минуту
    cost_per_token: float = 0.0  # Стоимость за токен
    
    # Stats
    requests_count: int = 0
    error_count: int = 0
    avg_latency: float = 0.0
    last_request: Optional[datetime] = None

class ModelPool:
    """Управление пулом моделей"""
    
    def __init__(self):
        self.endpoints: Dict[str, ModelEndpoint] = {}
        self._lock = asyncio.Lock()
    
    def add_endpoint(self, endpoint: ModelEndpoint):
        """Добавление эндпоинта"""
        self.endpoints[endpoint.name] = endpoint
    
    def remove_endpoint(self, name: str):
        """Удаление эндпоинта"""
        self.endpoints.pop(name, None)
    
    def get_by_route(self, route: ModelRoute) -> List[ModelEndpoint]:
        """Получение эндпоинтов по маршруту"""
        return [
            ep for ep in self.endpoints.values()
            if ep.route == route and ep.health != "unhealthy"
        ]
    
    def get_healthy_endpoints(self) -> List[ModelEndpoint]:
        """Получение всех здоровых эндпоинтов"""
        return [
            ep for ep in self.endpoints.values()
            if ep.health != "unhealthy"
        ]
    
    def select_endpoint(
        self,
        route: ModelRoute = None,
        load_balance: bool = True
    ) -> Optional[ModelEndpoint]:
        """Выбор эндпоинта"""
        if route:
            candidates = self.get_by_route(route)
        else:
            candidates = self.get_healthy_endpoints()
        
        if not candidates:
            return None
        
        if load_balance:
            # Weighted random selection
            total_weight = sum(ep.weight for ep in candidates)
            if total_weight == 0:
                return candidates[0]
            
            r = random.uniform(0, total_weight)
            cumulative = 0
            for ep in candidates:
                cumulative += ep.weight
                if r <= cumulative:
                    return ep
            return candidates[-1]
        
        # Round-robin based on request count
        return min(candidates, key=lambda ep: ep.requests_count)
    
    def update_health(self, name: str, health: str):
        """Обновление статуса здоровья"""
        if name in self.endpoints:
            self.endpoints[name].health = health
    
    def update_stats(self, name: str, latency: float, success: bool):
        """Обновление статистики"""
        if name in self.endpoints:
            ep = self.endpoints[name]
            ep.requests_count += 1
            ep.last_request = datetime.now()
            if not success:
                ep.error_count += 1
            # Exponential moving average для latency
            alpha = 0.1
            ep.avg_latency = (1 - alpha) * ep.avg_latency + alpha * latency

# Глобальный пул моделей
model_pool = ModelPool()

# Инициализация моделей
model_pool.add_endpoint(ModelEndpoint(
    name="fast-local",
    provider=ModelProvider.OLLAMA,
    url="http://localhost:11434",
    model_id="llama3.1:8b",
    route=ModelRoute.FAST,
    max_tokens=2048,
    weight=3,
    rate_limit=200
))

model_pool.add_endpoint(ModelEndpoint(
    name="smart-local",
    provider=ModelProvider.VLLM,
    url="http://localhost:8000",
    model_id="Qwen2.5-72B-Instruct",
    route=ModelRoute.SMART,
    max_tokens=8192,
    weight=1,
    rate_limit=50
))

model_pool.add_endpoint(ModelEndpoint(
    name="code-local",
    provider=ModelProvider.OLLAMA,
    url="http://localhost:11434",
    model_id="codellama:13b",
    route=ModelRoute.CODE,
    max_tokens=4096,
    weight=2,
    rate_limit=100
))

model_pool.add_endpoint(ModelEndpoint(
    name="chat-local",
    provider=ModelProvider.OLLAMA,
    url="http://localhost:11434",
    model_id="mistral:7b",
    route=ModelRoute.CHAT,
    max_tokens=4096,
    weight=2,
    rate_limit=150
))

Routing Engine

import re
from collections import defaultdict

class RoutingEngine:
    """Engine для маршрутизации запросов к моделям"""
    
    def __init__(self, model_pool: ModelPool):
        self.model_pool = model_pool
        self.rules: List[Dict] = []
        self._default_route = ModelRoute.DEFAULT
    
    def add_rule(
        self,
        name: str,
        pattern: str,
        route: ModelRoute,
        priority: int = 0,
        regex: bool = True
    ):
        """Добавление правила маршрутизации"""
        compiled = re.compile(pattern) if regex else None
        self.rules.append({
            "name": name,
            "pattern": pattern,
            "regex": compiled,
            "route": route,
            "priority": priority
        })
        # Сортировка по приоритету
        self.rules.sort(key=lambda r: r["priority"], reverse=True)
    
    def route_request(
        self,
        messages: List[Dict],
        system_prompt: str = None,
        tags: List[str] = None
    ) -> ModelRoute:
        """Определение маршрута для запроса"""
        # Объединение всех сообщений для анализа
        full_text = " ".join(m.get("content", "") for m in messages)
        
        # Проверка правил
        for rule in self.rules:
            if rule["regex"]:
                if rule["regex"].search(full_text):
                    return rule["route"]
            else:
                if rule["pattern"].lower() in full_text.lower():
                    return rule["route"]
        
        # Анализ контента для автоматического определения
        return self._analyze_and_route(full_text, tags)
    
    def _analyze_and_route(self, text: str, tags: List[str] = None) -> ModelRoute:
        """Анализ текста для определения маршрута"""
        text_lower = text.lower()
        
        # Code detection
        code_indicators = [
            "def ", "function ", "class ", "import ", "const ",
            "let ", "var ", "return ", "if ", "for ", "while ",
            "```", "```python", "```javascript", "```typescript"
        ]
        if any(ind in text_lower for ind in code_indicators):
            return ModelRoute.CODE
        
        # Complex reasoning indicators
        reasoning_indicators = [
            "explain", "analyze", "compare", "evaluate",
            "why", "how does", "architecture", "design pattern"
        ]
        if any(ind in text_lower for ind in reasoning_indicators):
            return ModelRoute.SMART
        
        # Tags-based routing
        if tags:
            if "code" in tags:
                return ModelRoute.CODE
            if "reasoning" in tags:
                return ModelRoute.SMART
        
        return self._default_route

# Инициализация routing engine
routing_engine = RoutingEngine(model_pool)

# Правила маршрутизации
routing_engine.add_rule(
    name="code_block",
    pattern=r"```",
    route=ModelRoute.CODE,
    priority=100
)
routing_engine.add_rule(
    name="explain",
    pattern="explain",
    route=ModelRoute.SMART,
    priority=50
)
routing_engine.add_rule(
    name="simple_qa",
    pattern=r"^(what|who|when|where|how)\b",
    route=ModelRoute.FAST,
    priority=30
)

Fallback Mechanism

class FallbackChain:
    """Цепочка fallback для отказоустойчивости"""
    
    def __init__(self):
        self.chains: Dict[str, List[ModelEndpoint]] = {}
    
    def add_chain(
        self,
        name: str,
        endpoints: List[ModelEndpoint],
        max_retries: int = 3,
        retry_delay: float = 1.0
    ):
        """Добавление цепочки fallback"""
        self.chains[name] = {
            "endpoints": endpoints,
            "max_retries": max_retries,
            "retry_delay": retry_delay
        }
    
    async def execute_with_fallback(
        self,
        chain_name: str,
        request_fn,
        *args,
        **kwargs
    ):
        """Выполнение с fallback"""
        chain = self.chains.get(chain_name)
        if not chain:
            raise ValueError(f"Chain {chain_name} not found")
        
        last_error = None
        endpoints = chain["endpoints"]
        
        for attempt, endpoint in enumerate(endpoints):
            try:
                # Health check
                if endpoint.health == "unhealthy":
                    continue
                
                # Execute request
                result = await request_fn(endpoint, *args, **kwargs)
                
                # Update stats
                model_pool.update_stats(endpoint.name, result["latency"], True)
                
                return result
                
            except Exception as e:
                last_error = e
                model_pool.update_stats(endpoint.name, 0, False)
                endpoint.health = "degraded"
                
                if attempt < chain["max_retries"]:
                    await asyncio.sleep(chain["retry_delay"])
        
        raise last_error

# Fallback chains
fallback_chain = FallbackChain()
fallback_chain.add_chain(
    name="default",
    endpoints=[
        model_pool.endpoints["fast-local"],
        model_pool.endpoints["chat-local"],
        model_pool.endpoints["smart-local"],
    ],
    max_retries=3,
    retry_delay=0.5
)

Rate Limiting & Quotas

import time
from collections import defaultdict

class RateLimiter:
    """Rate limiter с поддержкой sliding window"""
    
    def __init__(self):
        self.requests: Dict[str, List[float]] = defaultdict(list)
        self.buckets: Dict[str, Dict] = {}
    
    def add_bucket(
        self,
        name: str,
        rate: int,  # Запросов за период
        window: int = 60  # Период в секундах
    ):
        """Добавление bucket'а"""
        self.buckets[name] = {
            "rate": rate,
            "window": window
        }
    
    def is_allowed(
        self,
        key: str,
        bucket: str = "default"
    ) -> bool:
        """Проверка что запрос разрешён"""
        if bucket not in self.buckets:
            return True
        
        config = self.buckets[bucket]
        now = time.time()
        window_start = now - config["window"]
        
        # Очистка старых запросов
        self.requests[key] = [
            t for t in self.requests[key] if t > window_start
        ]
        
        if len(self.requests[key]) >= config["rate"]:
            return False
        
        self.requests[key].append(now)
        return True
    
    def get_remaining(
        self,
        key: str,
        bucket: str = "default"
    ) -> int:
        """Получение оставшихся запросов"""
        if bucket not in self.buckets:
            return float("inf")
        
        config = self.buckets[bucket]
        now = time.time()
        window_start = now - config["window"]
        
        self.requests[key] = [
            t for t in self.requests[key] if t > window_start
        ]
        
        return max(0, config["rate"] - len(self.requests[key]))

# API Key management
class APIKeyManager:
    """Управление API ключами"""
    
    def __init__(self):
        self.keys: Dict[str, Dict] = {}
    
    def create_key(
        self,
        name: str,
        rate_limit: int = 60,
        max_tokens_per_day: int = 1_000_000,
        allowed_models: List[str] = None
    ) -> str:
        """Создание нового ключа"""
        import secrets
        key = f"llm_{secrets.token_hex(24)}"
        
        self.keys[key] = {
            "name": name,
            "rate_limit": rate_limit,
            "max_tokens_per_day": max_tokens_per_day,
            "allowed_models": allowed_models,
            "created_at": datetime.now(),
            "tokens_used_today": 0,
            "is_active": True
        }
        
        return key
    
    def validate_key(self, key: str) -> bool:
        """Валидация ключа"""
        if key not in self.keys:
            return False
        
        k = self.keys[key]
        if not k["is_active"]:
            return False
        
        # Проверка daily quota
        if k["tokens_used_today"] >= k["max_tokens_per_day"]:
            return False
        
        return True
    
    def record_usage(self, key: str, tokens: int):
        """Запись использования"""
        if key in self.keys:
            self.keys[key]["tokens_used_today"] += tokens

rate_limiter = RateLimiter()
api_key_manager = APIKeyManager()

# Default bucket
rate_limiter.add_bucket("default", rate=60, window=60)

Observability & Logging

import structlog
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import List, Dict
import json
import uuid

logger = structlog.get_logger()

class RequestLogger:
    """Логирование запросов и ответов"""
    
    def __init__(self, log_dir: str = "logs"):
        self.log_dir = Path(log_dir)
        self.log_dir.mkdir(parents=True, exist_ok=True)
    
    def log_request(
        self,
        request_id: str,
        model: str,
        messages: List[Dict],
        input_tokens: int,
        output_tokens: int,
        latency: float,
        status: str
    ):
        """Логирование запроса"""
        log_entry = {
            "request_id": request_id,
            "timestamp": datetime.now().isoformat(),
            "model": model,
            "input_tokens": input_tokens,
            "output_tokens": output_tokens,
            "total_tokens": input_tokens + output_tokens,
            "latency": latency,
            "ttft": latency / 2,  # Time to first token (в реальном коде отдельно)
            "tps": output_tokens / (latency / 2) if latency > 0 else 0,  # Tokens per second
            "status": status,
            "cost": self._estimate_cost(model, input_tokens, output_tokens)
        }
        
        # Write to file
        log_file = self.log_dir / f"{datetime.now().strftime('%Y-%m-%d')}.jsonl"
        with open(log_file, "a") as f:
            f.write(json.dumps(log_entry) + "\n")
        
        return log_entry
    
    def _estimate_cost(
        self,
        model: str,
        input_tokens: int,
        output_tokens: int
    ) -> float:
        """Оценка стоимости запроса"""
        costs = {
            "gpt-4": {"input": 0.03 / 1000, "output": 0.06 / 1000},
            "gpt-3.5": {"input": 0.0015 / 1000, "output": 0.002 / 1000},
            "local": {"input": 0.0, "output": 0.0},
        }
        
        rate = costs.get(model, {"input": 0.0, "output": 0.0})
        return (
            input_tokens * rate["input"] +
            output_tokens * rate["output"]
        )

request_logger = RequestLogger()

Health Check & Self-Healing

class HealthChecker:
    """Проверка здоровья моделей"""
    
    def __init__(self, model_pool: ModelPool, interval: int = 30):
        self.model_pool = model_pool
        self.interval = interval
        self._running = False
        self._task = None
    
    async def start(self):
        """Запуск health checker"""
        self._running = True
        self._task = asyncio.create_task(self._check_loop())
    
    async def stop(self):
        """Остановка"""
        self._running = False
        if self._task:
            self._task.cancel()
    
    async def _check_loop(self):
        """Цикл проверок"""
        while self._running:
            for name, endpoint in self.model_pool.endpoints.items():
                await self._check_endpoint(endpoint)
                await asyncio.sleep(1)  # Задержка между проверками
            
            await asyncio.sleep(self.interval - 1)
    
    async def _check_endpoint(self, endpoint: ModelEndpoint):
        """Проверка одного эндпоинта"""
        try:
            async with aiohttp.ClientSession() as session:
                if endpoint.provider == ModelProvider.OLLAMA:
                    async with session.get(
                        f"{endpoint.url}/api/tags",
                        timeout=aiohttp.ClientTimeout(total=5)
                    ) as resp:
                        if resp.status != 200:
                            self.model_pool.update_health(name, "unhealthy")
                            return
                
                elif endpoint.provider == ModelProvider.VLLM:
                    async with session.get(
                        f"{endpoint.url}/health",
                        timeout=aiohttp.ClientTimeout(total=5)
                    ) as resp:
                        if resp.status != 200:
                            self.model_pool.update_health(name, "unhealthy")
                            return
                
                # Check error rate
                ep = self.model_pool.endpoints[endpoint.name]
                if ep.requests_count > 10:
                    error_rate = ep.error_count / ep.requests_count
                    if error_rate > 0.5:
                        self.model_pool.update_health(name, "degraded")
                    elif error_rate > 0.1:
                        self.model_pool.update_health(name, "degraded")
                    else:
                        self.model_pool.update_health(name, "healthy")
                else:
                    self.model_pool.update_health(name, "healthy")
                    
        except Exception:
            self.model_pool.update_health(endpoint.name, "unhealthy")

# Self-healing: автоматический перезапуск упавших моделей
class SelfHealingManager:
    """Менеджер самовосстановления"""
    
    def __init__(self, health_checker: HealthChecker):
        self.health_checker = health_checker
        self.recovery_actions: Dict[str, callable] = {}
    
    def register_recovery(self, model_name: str, action: callable):
        """Регистрация действия для восстановления"""
        self.recovery_actions[model_name] = action
    
    async def recover_model(self, model_name: str):
        """Восстановление модели"""
        action = self.recovery_actions.get(model_name)
        if action:
            logger.info("recovering", model=model_name)
            try:
                await action()
                model_pool.update_health(model_name, "healthy")
                logger.info("recovered", model=model_name)
            except Exception as e:
                logger.error("recovery_failed", model=model_name, error=str(e))

API Endpoints

@app.get("/health")
async def health_check():
    """Общий статус gateway"""
    healthy = sum(
        1 for ep in model_pool.get_healthy_endpoints()
    )
    total = len(model_pool.endpoints)
    
    return {
        "status": "healthy" if healthy == total else "degraded",
        "healthy": healthy,
        "total": total,