LLM Gateway & Routing: Маршрутизация запросов к LLM
llm-gatewayroutingload-balancingfallbackmulti-model
Зачем нужен 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,