224 lines
6.3 KiB
Python
224 lines
6.3 KiB
Python
"""
|
|
Ollama client for model inference.
|
|
Handles both streaming and non-streaming requests.
|
|
"""
|
|
|
|
import httpx
|
|
import json
|
|
import logging
|
|
from typing import AsyncIterator, Dict, Any, Optional
|
|
from src.config import get_settings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
settings = get_settings()
|
|
|
|
|
|
class OllamaClient:
|
|
"""Client for interacting with Ollama API."""
|
|
|
|
def __init__(self):
|
|
self.base_url = settings.ollama_base_url
|
|
self.timeout = settings.ollama_timeout
|
|
self.client = httpx.AsyncClient(timeout=self.timeout)
|
|
logger.info(f"Initialized Ollama client: {self.base_url}")
|
|
|
|
async def close(self):
|
|
"""Close the HTTP client."""
|
|
await self.client.aclose()
|
|
|
|
def resolve_model(self, model_name: str) -> str:
|
|
"""
|
|
Resolve model alias to actual Ollama model.
|
|
|
|
Args:
|
|
model_name: Requested model name (e.g., "gpt-3.5-turbo")
|
|
|
|
Returns:
|
|
Actual Ollama model name (e.g., "gemma:7b")
|
|
"""
|
|
resolved = settings.model_aliases.get(model_name, model_name)
|
|
if resolved != model_name:
|
|
logger.info(f"Model resolution: {model_name} → {resolved}")
|
|
return resolved
|
|
|
|
async def generate_non_streaming(
|
|
self,
|
|
model: str,
|
|
prompt: str,
|
|
temperature: float = 0.7,
|
|
max_tokens: Optional[int] = None
|
|
) -> Dict[str, Any]:
|
|
"""
|
|
Generate non-streaming response from Ollama using chat endpoint.
|
|
|
|
Args:
|
|
model: Model name
|
|
prompt: User prompt
|
|
temperature: Sampling temperature
|
|
max_tokens: Maximum tokens to generate
|
|
|
|
Returns:
|
|
Dict with 'response' and 'tokens' keys
|
|
"""
|
|
actual_model = self.resolve_model(model)
|
|
|
|
payload = {
|
|
"model": actual_model,
|
|
"messages": [
|
|
{"role": "user", "content": prompt}
|
|
],
|
|
"stream": False,
|
|
"options": {
|
|
"temperature": temperature,
|
|
}
|
|
}
|
|
|
|
if max_tokens:
|
|
payload["options"]["num_predict"] = max_tokens
|
|
|
|
logger.debug(f"Ollama request to {actual_model}")
|
|
|
|
try:
|
|
response = await self.client.post(
|
|
f"{self.base_url}/api/chat",
|
|
json=payload
|
|
)
|
|
response.raise_for_status()
|
|
result = response.json()
|
|
|
|
return {
|
|
"response": result.get("message", {}).get("content", ""),
|
|
"tokens": {
|
|
"prompt": result.get("prompt_eval_count", 0),
|
|
"completion": result.get("eval_count", 0),
|
|
"total": result.get("prompt_eval_count", 0) + result.get("eval_count", 0)
|
|
}
|
|
}
|
|
|
|
except httpx.HTTPError as e:
|
|
logger.error(f"Ollama request failed: {e}")
|
|
raise
|
|
|
|
async def generate_streaming(
|
|
self,
|
|
model: str,
|
|
prompt: str,
|
|
temperature: float = 0.7,
|
|
max_tokens: Optional[int] = None
|
|
) -> AsyncIterator[str]:
|
|
"""
|
|
Generate streaming response from Ollama using chat endpoint.
|
|
|
|
Args:
|
|
model: Model name
|
|
prompt: User prompt
|
|
temperature: Sampling temperature
|
|
max_tokens: Maximum tokens to generate
|
|
|
|
Yields:
|
|
Token strings
|
|
"""
|
|
actual_model = self.resolve_model(model)
|
|
|
|
payload = {
|
|
"model": actual_model,
|
|
"messages": [
|
|
{"role": "user", "content": prompt}
|
|
],
|
|
"stream": True,
|
|
"options": {
|
|
"temperature": temperature,
|
|
}
|
|
}
|
|
|
|
if max_tokens:
|
|
payload["options"]["num_predict"] = max_tokens
|
|
|
|
logger.debug(f"Ollama streaming request to {actual_model}")
|
|
|
|
try:
|
|
async with self.client.stream(
|
|
"POST",
|
|
f"{self.base_url}/api/chat",
|
|
json=payload
|
|
) as response:
|
|
response.raise_for_status()
|
|
|
|
async for line in response.aiter_lines():
|
|
if not line:
|
|
continue
|
|
|
|
try:
|
|
chunk = json.loads(line)
|
|
if "message" in chunk:
|
|
content = chunk["message"].get("content", "")
|
|
if content:
|
|
yield content
|
|
|
|
# Check if done
|
|
if chunk.get("done", False):
|
|
break
|
|
|
|
except json.JSONDecodeError:
|
|
logger.warning(f"Failed to parse JSON: {line}")
|
|
continue
|
|
|
|
except httpx.HTTPError as e:
|
|
logger.error(f"Ollama streaming request failed: {e}")
|
|
raise
|
|
|
|
async def health_check(self) -> bool:
|
|
"""
|
|
Check if Ollama is healthy.
|
|
|
|
Returns:
|
|
True if healthy, False otherwise
|
|
"""
|
|
try:
|
|
response = await self.client.get(
|
|
f"{self.base_url}/api/tags",
|
|
timeout=5.0
|
|
)
|
|
return response.status_code == 200
|
|
except Exception as e:
|
|
logger.error(f"Ollama health check failed: {e}")
|
|
return False
|
|
|
|
async def list_models(self) -> Dict[str, Any]:
|
|
"""
|
|
List all available models in Ollama.
|
|
|
|
Returns:
|
|
Dict with 'models' key containing list of model info
|
|
"""
|
|
try:
|
|
response = await self.client.get(
|
|
f"{self.base_url}/api/tags",
|
|
timeout=5.0
|
|
)
|
|
response.raise_for_status()
|
|
return response.json()
|
|
except Exception as e:
|
|
logger.error(f"Failed to list Ollama models: {e}")
|
|
raise
|
|
|
|
|
|
# Global client instance
|
|
_ollama_client: Optional[OllamaClient] = None
|
|
|
|
|
|
def get_ollama_client() -> OllamaClient:
|
|
"""Get or create the global Ollama client instance."""
|
|
global _ollama_client
|
|
if _ollama_client is None:
|
|
_ollama_client = OllamaClient()
|
|
return _ollama_client
|
|
|
|
|
|
async def close_ollama_client():
|
|
"""Close the global Ollama client."""
|
|
global _ollama_client
|
|
if _ollama_client is not None:
|
|
await _ollama_client.close()
|
|
_ollama_client = None
|