feat: implement AI model API proxy service

This commit is contained in:
2026-03-11 18:22:16 +08:00
commit 26738973bd
33 changed files with 4607 additions and 0 deletions

View File

@@ -0,0 +1,7 @@
"""Business logic services."""
from .auth import AuthService
from .key_selector import KeySelector
from .router import ModelRouter
from .logger import AsyncLogger
__all__ = ["AuthService", "KeySelector", "ModelRouter", "AsyncLogger"]

View File

@@ -0,0 +1,36 @@
"""Authentication service."""
import aiosqlite
from typing import Optional
from ..database import db_connection
from ..models import ClientKey
class AuthService:
"""Service for client key authentication."""
@staticmethod
async def validate_key(api_key: str) -> Optional[ClientKey]:
"""Validate a client API key.
Args:
api_key: The API key to validate
Returns:
ClientKey if valid and active, None otherwise
"""
async with db_connection() as db:
cursor = await db.execute(
"""SELECT id, key, name, is_active, created_at
FROM client_keys WHERE key = ?""",
(api_key,)
)
row = await cursor.fetchone()
if row and row["is_active"]:
return ClientKey(
id=row["id"],
key=row["key"],
name=row["name"],
is_active=row["is_active"],
created_at=row["created_at"]
)
return None

View File

@@ -0,0 +1,60 @@
"""Provider key selector with round-robin strategy."""
import asyncio
from typing import Optional, Dict
from ..database import db_connection
from ..models import ProviderKey
class KeySelector:
"""Round-robin key selector for provider API keys."""
def __init__(self):
self._indices: Dict[int, int] = {} # provider_id -> current index
self._lock = asyncio.Lock()
async def get_next_key(self, provider_id: int) -> Optional[ProviderKey]:
"""Get the next available API key for a provider.
Args:
provider_id: The provider ID
Returns:
ProviderKey if available, None if no keys configured
"""
async with self._lock:
async with db_connection() as db:
cursor = await db.execute(
"""SELECT id, provider_id, key, is_active, created_at
FROM provider_keys
WHERE provider_id = ? AND is_active = TRUE""",
(provider_id,)
)
keys = await cursor.fetchall()
if not keys:
return None
# Get or initialize index
if provider_id not in self._indices:
self._indices[provider_id] = 0
# Round-robin selection
index = self._indices[provider_id] % len(keys)
self._indices[provider_id] = index + 1
row = keys[index]
return ProviderKey(
id=row["id"],
provider_id=row["provider_id"],
key=row["key"],
is_active=row["is_active"],
created_at=row["created_at"]
)
def reset_index(self, provider_id: int):
"""Reset the round-robin index for a provider."""
self._indices.pop(provider_id, None)
# Global instance
key_selector = KeySelector()

View File

@@ -0,0 +1,99 @@
"""Async batch logger for request logs."""
import asyncio
from typing import Optional, List
from dataclasses import dataclass
from ..database import db_connection
from ..config import LOG_BATCH_SIZE, LOG_FLUSH_INTERVAL
@dataclass
class LogEntry:
"""A single log entry."""
client_key_id: int
provider_id: int
model: str
prompt_tokens: Optional[int] = None
completion_tokens: Optional[int] = None
latency_ms: Optional[int] = None
success: bool = True
error_message: Optional[str] = None
class AsyncLogger:
"""Asynchronous batch logger for request logs."""
def __init__(self):
self._queue: asyncio.Queue[LogEntry] = asyncio.Queue()
self._task: Optional[asyncio.Task] = None
async def start(self):
"""Start the background flush task."""
if self._task is None:
self._task = asyncio.create_task(self._flush_loop())
async def stop(self):
"""Stop the background flush task."""
if self._task:
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
self._task = None
# Flush remaining logs
await self._flush_all()
async def log(self, entry: LogEntry):
"""Queue a log entry."""
await self._queue.put(entry)
async def _flush_loop(self):
"""Background loop to flush logs periodically."""
while True:
try:
await asyncio.sleep(LOG_FLUSH_INTERVAL)
await self._flush_batch()
except asyncio.CancelledError:
raise
async def _flush_batch(self):
"""Flush a batch of logs to database."""
entries: List[LogEntry] = []
# Collect up to batch_size entries
while len(entries) < LOG_BATCH_SIZE:
try:
entry = self._queue.get_nowait()
entries.append(entry)
except asyncio.QueueEmpty:
break
if not entries:
return
# Insert batch
async with db_connection() as db:
await db.executemany(
"""INSERT INTO request_logs
(client_key_id, provider_id, model, prompt_tokens,
completion_tokens, latency_ms, success, error_message)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
[
(
e.client_key_id, e.provider_id, e.model,
e.prompt_tokens, e.completion_tokens,
e.latency_ms, e.success, e.error_message
)
for e in entries
]
)
await db.commit()
async def _flush_all(self):
"""Flush all remaining logs."""
while not self._queue.empty():
await self._flush_batch()
# Global instance
async_logger = AsyncLogger()

View File

@@ -0,0 +1,76 @@
"""Model to provider router."""
from typing import Optional, Tuple
from ..database import db_connection
from ..models import Provider
from ..config import MODEL_PROVIDER_MAP
class ModelRouter:
"""Routes model names to providers."""
@staticmethod
def get_provider_name(model: str) -> Optional[str]:
"""Determine provider name from model name.
Args:
model: The model name (e.g., 'gpt-4o', 'claude-3-opus')
Returns:
Provider name or None if unknown
"""
model_lower = model.lower()
for prefix, provider in MODEL_PROVIDER_MAP.items():
if model_lower.startswith(prefix):
return provider
return None
@staticmethod
async def get_provider(provider_name: str) -> Optional[Provider]:
"""Get provider configuration by name.
Args:
provider_name: The provider name (e.g., 'openai', 'anthropic')
Returns:
Provider if found and active, None otherwise
"""
async with db_connection() as db:
cursor = await db.execute(
"""SELECT id, name, base_url, api_type, is_active
FROM providers WHERE name = ? AND is_active = TRUE""",
(provider_name,)
)
row = await cursor.fetchone()
if row:
return Provider(
id=row["id"],
name=row["name"],
base_url=row["base_url"],
api_type=row["api_type"],
is_active=row["is_active"]
)
return None
@staticmethod
async def route_model(model: str) -> Tuple[Optional[Provider], Optional[str]]:
"""Route a model to its provider.
Args:
model: The model name
Returns:
Tuple of (Provider, error_message)
"""
provider_name = ModelRouter.get_provider_name(model)
if not provider_name:
return None, f"Unsupported model: {model}"
provider = await ModelRouter.get_provider(provider_name)
if not provider:
return None, f"Provider not found: {provider_name}"
return provider, None
# Global instance
model_router = ModelRouter()