"""LLM handler for Ollama, Groq, and HuggingFace Transformers providers.""" import logging import os from collections.abc import Iterator from typing import Any from langchain_core.callbacks import CallbackManagerForLLMRun from langchain_core.language_models.base import BaseLanguageModel from langchain_core.language_models.llms import LLM from pydantic import PrivateAttr from .config_loader import get_config logger = logging.getLogger(__name__) # Available Ollama models (common options) OLLAMA_MODELS = [ "llama3.2:3b", "llama3.2:1b", "llama3.1:8b", "phi3:mini", "gemma2:2b", "mistral:7b", "qwen2.5:3b", ] DEFAULT_OLLAMA_MODEL = "llama3.2:3b" # Available Groq models (free tier) GROQ_MODELS = [ "openai/gpt-oss-120b", "llama-3.3-70b-versatile", "llama-3.1-8b-instant", "mixtral-8x7b-32768", "gemma2-9b-it", ] DEFAULT_GROQ_MODEL = "openai/gpt-oss-120b" DEFAULT_HF_MODEL = "Qwen/Qwen3.5-4B" HF_MODELS = [ "Qwen/Qwen3.5-4B", "Qwen/Qwen2.5-3B-Instruct", "Qwen/Qwen2.5-1.5B-Instruct", "microsoft/Phi-3.5-mini-instruct", ] class HuggingFaceTransformersLLM(LLM): """LangChain-compatible wrapper around a local HuggingFace Causal LM.""" model_name: str = DEFAULT_HF_MODEL temperature: float = 0.7 max_new_tokens: int = 800 top_p: float = 0.9 _tokenizer: Any = PrivateAttr(default=None) _model: Any = PrivateAttr(default=None) @property def _llm_type(self) -> str: return "huggingface_transformers" def _ensure_loaded(self) -> None: """Lazily load tokenizer and model onto the best available device.""" if self._model is not None and self._tokenizer is not None: return import torch from transformers import AutoModelForCausalLM, AutoTokenizer logger.info(f"Loading HuggingFace model: {self.model_name}") self._tokenizer = AutoTokenizer.from_pretrained(self.model_name, trust_remote_code=True) if self._tokenizer.pad_token is None: self._tokenizer.pad_token = self._tokenizer.eos_token dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32 device_map = "auto" if torch.cuda.is_available() else None load_kwargs = { "trust_remote_code": True, "device_map": device_map, } # Prefer `dtype=` (torch_dtype is deprecated in recent transformers) model = None for dtype_key in ("dtype", "torch_dtype"): try: model = AutoModelForCausalLM.from_pretrained( self.model_name, **{dtype_key: dtype}, **load_kwargs, ) break except TypeError: continue except Exception as causal_err: logger.warning( "AutoModelForCausalLM failed for %s (%s); trying image-text class", self.model_name, causal_err, ) break if model is None: try: from transformers import AutoModelForImageTextToText model = AutoModelForImageTextToText.from_pretrained( self.model_name, dtype=dtype, **load_kwargs, ) except Exception: model = AutoModelForCausalLM.from_pretrained( self.model_name, torch_dtype=dtype, **load_kwargs, ) self._model = model if device_map is None: self._model = self._model.to("cpu") self._model.eval() logger.info(f"HuggingFace model loaded: {self.model_name}") def _build_messages(self, prompt: str) -> list[dict[str, str]]: """Use chat template when available; fall back to raw prompt.""" return [{"role": "user", "content": prompt}] def _tokenize_prompt(self, prompt: str) -> tuple[Any, Any]: """Tokenize prompt into (input_ids, attention_mask) tensors.""" import torch messages = self._build_messages(prompt) attention_mask = None enable_thinking = bool(get_config().get("llm.enable_thinking", False)) if hasattr(self._tokenizer, "apply_chat_template"): template_kwargs: dict[str, Any] = { "add_generation_prompt": True, "return_tensors": "pt", "return_dict": True, } # Qwen3 / Qwen3.5 support enable_thinking in chat template try: encoded = self._tokenizer.apply_chat_template( messages, enable_thinking=enable_thinking, **template_kwargs, ) except TypeError: try: encoded = self._tokenizer.apply_chat_template( messages, **template_kwargs, ) except TypeError: encoded = self._tokenizer.apply_chat_template( messages, add_generation_prompt=True, return_tensors="pt", ) if hasattr(encoded, "input_ids"): input_ids = encoded["input_ids"] attention_mask = encoded.get("attention_mask") elif isinstance(encoded, dict): input_ids = encoded["input_ids"] attention_mask = encoded.get("attention_mask") else: input_ids = encoded else: encoded = self._tokenizer(prompt, return_tensors="pt") input_ids = encoded["input_ids"] attention_mask = encoded.get("attention_mask") if not torch.is_tensor(input_ids): input_ids = torch.tensor(input_ids) if attention_mask is None: attention_mask = torch.ones_like(input_ids) elif not torch.is_tensor(attention_mask): attention_mask = torch.tensor(attention_mask) return input_ids, attention_mask def _generate_text(self, prompt: str) -> str: import torch self._ensure_loaded() assert self._tokenizer is not None and self._model is not None input_ids, attention_mask = self._tokenize_prompt(prompt) model_device = next(self._model.parameters()).device input_ids = input_ids.to(model_device) attention_mask = attention_mask.to(model_device) with torch.inference_mode(): output_ids = self._model.generate( input_ids=input_ids, attention_mask=attention_mask, max_new_tokens=self.max_new_tokens, temperature=max(self.temperature, 1e-5), top_p=self.top_p, do_sample=self.temperature > 0, pad_token_id=self._tokenizer.pad_token_id, eos_token_id=self._tokenizer.eos_token_id, ) generated = output_ids[0][input_ids.shape[-1] :] return self._tokenizer.decode(generated, skip_special_tokens=True).strip() def _call( self, prompt: str, stop: list[str] | None = None, run_manager: CallbackManagerForLLMRun | None = None, # noqa: ARG002 **kwargs: Any, # noqa: ARG002 ) -> str: text = self._generate_text(prompt) if stop: for token in stop: if token in text: text = text.split(token)[0] return text def _stream( self, prompt: str, stop: list[str] | None = None, run_manager: CallbackManagerForLLMRun | None = None, **kwargs: Any, ) -> Iterator[Any]: """Non-token streaming fallback: yield the full generation once.""" from langchain_core.outputs import GenerationChunk text = self._call(prompt, stop=stop, run_manager=run_manager, **kwargs) chunk = GenerationChunk(text=text) if run_manager: run_manager.on_llm_new_token(text, chunk=chunk) yield chunk class LLMHandler: """Handle LLM initialization and configuration.""" def __init__( self, provider_override: str | None = None, api_key_override: str | None = None, model_override: str | None = None, ): """Initialize LLM handler with configuration. Args: provider_override: Override provider from config (ollama/groq/transformers) api_key_override: Override API key from environment model_override: Override model from config """ self.config = get_config() self._api_key_override = api_key_override self._model_override = model_override self.llm: BaseLanguageModel | None = None # Auto-detect provider based on overrides / env / config if provider_override: self.provider = provider_override else: groq_api_key = self._get_groq_api_key() config_provider = self.config.get("llm.provider", "ollama") # Prefer explicit transformers (HF Spaces) over Groq auto-detect if config_provider == "transformers": self.provider = "transformers" elif groq_api_key: self.provider = "groq" logger.info("GROQ_API_KEY detected, using Groq as default provider") else: self.provider = config_provider def _get_groq_api_key(self) -> str | None: """Get Groq API key from override or environment.""" if self._api_key_override: return self._api_key_override return self.config.get_env("GROQ_API_KEY") def get_ollama_llm(self) -> BaseLanguageModel: """Initialize Ollama LLM. Returns: OllamaLLM instance """ from langchain_ollama import OllamaLLM model = self._model_override or self.config.get("llm.model", "llama3.2:3b") temperature = self.config.get("llm.temperature", 0.7) base_url = self.config.get_env("OLLAMA_BASE_URL", "http://localhost:11434") logger.info(f"Initializing Ollama with model: {model}") try: llm = OllamaLLM( model=model, temperature=temperature, base_url=base_url, num_predict=self.config.get("llm.max_tokens", 512), top_p=self.config.get("llm.top_p", 0.9), ) # Test the connection logger.debug("Testing Ollama connection...") test_response = llm.invoke("Hi") logger.debug(f"Ollama test response: {test_response[:50]}...") return llm except Exception as e: # Check for "model not found" or 404 error error_str = str(e).lower() if "not found" in error_str or "404" in error_str: logger.warning(f"Ollama model '{model}' not found.") # We raise a ValueError with a helpful message that app.py can display nicely msg = ( f"Model '{model}' not found in Ollama.\n" f"Please run this command in your terminal:\n" f"ollama pull {model}" ) raise ValueError(msg) from e logger.error(f"Error initializing Ollama: {e}") logger.error( f"Make sure Ollama is running and the model is pulled. Try: ollama pull {model}" ) # Re-raise the original exception if it's not a missing model issue raise def get_groq_llm(self) -> BaseLanguageModel: """Initialize Groq LLM. Returns: ChatGroq instance """ from langchain_groq import ChatGroq api_key = self._get_groq_api_key() if not api_key: msg = "GROQ_API_KEY not found. Set it in .env or provide via UI." raise ValueError(msg) model = self._model_override or self.config.get("llm.groq_model", DEFAULT_GROQ_MODEL) temperature = self.config.get("llm.temperature", 0.7) max_tokens = self.config.get("llm.max_tokens", 800) logger.info(f"Initializing Groq with model: {model}") try: llm = ChatGroq( api_key=api_key, model=model, temperature=temperature, max_tokens=max_tokens, ) # Test the connection logger.debug("Testing Groq connection...") test_response = llm.invoke("Hi") logger.debug(f"Groq test response: {str(test_response.content)[:50]}...") return llm except Exception as e: logger.error(f"Error initializing Groq: {e}") raise def get_transformers_llm(self) -> HuggingFaceTransformersLLM: """Initialize HuggingFace Transformers LLM (ZeroGPU / local). Returns: HuggingFaceTransformersLLM instance """ model = self._model_override or self.config.get("llm.hf_model", DEFAULT_HF_MODEL) temperature = float(self.config.get("llm.temperature", 0.7)) max_tokens = int(self.config.get("llm.max_tokens", 800)) top_p = float(self.config.get("llm.top_p", 0.9)) logger.info(f"Initializing Transformers LLM with model: {model}") return HuggingFaceTransformersLLM( model_name=model, temperature=temperature, max_new_tokens=max_tokens, top_p=top_p, ) def get_llm(self) -> BaseLanguageModel: """Get LLM instance based on configured provider. Returns: LLM instance """ if self.llm is not None: return self.llm if self.provider == "ollama": self.llm = self.get_ollama_llm() elif self.provider == "groq": self.llm = self.get_groq_llm() elif self.provider == "transformers": self.llm = self.get_transformers_llm() else: msg = f"Unsupported LLM provider: {self.provider}" raise ValueError(msg) return self.llm def get_system_prompt(self) -> str: """Get formatted system prompt from configuration. Returns: Formatted system prompt """ template = self.config.get( "llm.system_prompt", "You are a helpful AI assistant. Answer questions based on the provided context.", ) return self.config.format_template(template) def get_provider(self) -> str: """Get the current provider name.""" return self.provider def get_model(self) -> str: """Get the current model name.""" if self._model_override: return self._model_override if self.provider == "groq": return self.config.get("llm.groq_model", DEFAULT_GROQ_MODEL) if self.provider == "transformers": return self.config.get("llm.hf_model", DEFAULT_HF_MODEL) return self.config.get("llm.model", "llama3.2:3b") def get_llm_handler( provider_override: str | None = None, api_key_override: str | None = None, model_override: str | None = None, ) -> LLMHandler: """Get LLM handler instance with optional overrides.""" return LLMHandler( provider_override=provider_override, api_key_override=api_key_override, model_override=model_override, ) def get_available_groq_models() -> list[str]: """Get list of available Groq models.""" return GROQ_MODELS.copy() def get_default_groq_model() -> str: """Get default Groq model.""" return DEFAULT_GROQ_MODEL def get_available_hf_models() -> list[str]: """Get list of supported HuggingFace models.""" return HF_MODELS.copy() def get_default_hf_model() -> str: """Get default HuggingFace model for ZeroGPU.""" return DEFAULT_HF_MODEL def get_available_ollama_models() -> list[str]: """Get list of currently pulled Ollama models. Fetches the list from Ollama API. Falls back to static list if unavailable. """ try: import ollama config = get_config() base_url = config.get_env("OLLAMA_BASE_URL", "http://localhost:11434") # Create client with configured host client = ollama.Client(host=base_url) response = client.list() # Extract model names - response.models is a list of Model objects models = [] for model in response.models: # Each model has a .model attribute with the name name = getattr(model, "model", None) or getattr(model, "name", None) if name: models.append(name) if models: logger.debug(f"Found {len(models)} pulled Ollama models: {models}") return models # No models pulled - return empty list logger.warning("No Ollama models found. Please pull a model first.") return [] except Exception as e: logger.warning(f"Could not fetch Ollama models: {e}. Using static list.") return OLLAMA_MODELS.copy() def get_default_ollama_model() -> str: """Get default Ollama model.""" return DEFAULT_OLLAMA_MODEL def detect_default_provider() -> str: """Detect default provider based on environment. Loads .env file first to ensure environment variables are available. On Hugging Face Spaces → transformers (ZeroGPU). Locally → Groq if key present, otherwise Ollama (Streamlit UX). """ from pathlib import Path from dotenv import load_dotenv # Load .env file to ensure GROQ_API_KEY is available env_path = Path(".env") if env_path.exists(): load_dotenv(env_path) # SPACE_ID / SYSTEM=spaces are set on Hugging Face Spaces if os.getenv("SPACE_ID") or os.getenv("SYSTEM") == "spaces": return "transformers" if os.getenv("GROQ_API_KEY"): return "groq" return "ollama"