""" 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