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,151 @@
"""LiteLLM wrapper for unified model access."""
import time
import json
from typing import AsyncIterator, Optional, Dict, Any, List
from litellm import acompletion
from ..models import Provider
from ..services.logger import async_logger, LogEntry
class LiteLLMWrapper:
"""Wrapper around LiteLLM for proxy functionality."""
@staticmethod
def _build_model_string(provider: Provider, model: str) -> str:
"""Build LiteLLM model string.
Args:
provider: The provider configuration
model: The model name
Returns:
LiteLLM-compatible model string
"""
if provider.api_type == "anthropic":
return f"anthropic/{model}"
else:
return f"openai/{model}"
@staticmethod
async def chat_completion(
provider: Provider,
api_key: str,
model: str,
messages: List[Dict[str, Any]],
stream: bool = False,
**kwargs
) -> tuple[Any, int, int]:
"""Execute a chat completion request.
Args:
provider: The provider configuration
api_key: The API key to use
model: The model name
messages: The chat messages
stream: Whether to stream the response
**kwargs: Additional parameters
Returns:
Tuple of (response, prompt_tokens, completion_tokens)
"""
model_str = LiteLLMWrapper._build_model_string(provider, model)
# Set API key for this request
api_base = provider.base_url if provider.api_type == "openai" else None
start_time = time.time()
try:
response = await acompletion(
model=model_str,
messages=messages,
api_key=api_key,
api_base=api_base,
stream=stream,
**kwargs
)
if stream:
# For streaming, return the async iterator
# Token counts will be collected during iteration
return response, 0, 0
else:
# Extract token usage from response
usage = getattr(response, "usage", None) or {}
prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0
completion_tokens = getattr(usage, "completion_tokens", 0) or 0
return response, prompt_tokens, completion_tokens
except Exception as e:
# Re-raise with context
raise
@staticmethod
async def stream_response(
stream_iterator: AsyncIterator,
client_key_id: int,
provider_id: int,
model: str,
start_time: float
) -> AsyncIterator[str]:
"""Process streaming response and yield SSE-formatted data.
Args:
stream_iterator: The LiteLLM stream iterator
client_key_id: Client key ID for logging
provider_id: Provider ID for logging
model: The model name
start_time: Request start time
Yields:
SSE-formatted strings
"""
prompt_tokens = 0
completion_tokens = 0
try:
async for chunk in stream_iterator:
# Extract token counts if available
if hasattr(chunk, "usage") and chunk.usage:
prompt_tokens = getattr(chunk.usage, "prompt_tokens", 0) or 0
completion_tokens = getattr(chunk.usage, "completion_tokens", 0) or 0
# Convert to SSE format
if hasattr(chunk, "model_dump"):
chunk_dict = chunk.model_dump()
else:
chunk_dict = chunk
yield f"data: {json.dumps(chunk_dict)}\n\n"
yield "data: [DONE]\n\n"
# Log successful request
latency_ms = int((time.time() - start_time) * 1000)
await async_logger.log(LogEntry(
client_key_id=client_key_id,
provider_id=provider_id,
model=model,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
latency_ms=latency_ms,
success=True
))
except Exception as e:
# Log failed request
latency_ms = int((time.time() - start_time) * 1000)
await async_logger.log(LogEntry(
client_key_id=client_key_id,
provider_id=provider_id,
model=model,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
latency_ms=latency_ms,
success=False,
error_message=str(e)
))
# Send error as SSE
yield f"data: {json.dumps({'error': str(e)})}\n\n"