feat: implement AI model API proxy service
This commit is contained in:
7
backend/app/services/__init__.py
Normal file
7
backend/app/services/__init__.py
Normal 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"]
|
||||
36
backend/app/services/auth.py
Normal file
36
backend/app/services/auth.py
Normal 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
|
||||
60
backend/app/services/key_selector.py
Normal file
60
backend/app/services/key_selector.py
Normal 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()
|
||||
99
backend/app/services/logger.py
Normal file
99
backend/app/services/logger.py
Normal 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()
|
||||
76
backend/app/services/router.py
Normal file
76
backend/app/services/router.py
Normal 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()
|
||||
Reference in New Issue
Block a user