feat: implement AI model API proxy service
This commit is contained in:
151
backend/app/providers/litellm_wrapper.py
Normal file
151
backend/app/providers/litellm_wrapper.py
Normal 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"
|
||||
Reference in New Issue
Block a user