Download src/core/model_manager.py from DiabetesCareChatbot/dmChatbotBackend: direct link, hf CLI and curl.
- Browser
- Download file 38.6 kB
-
https://huggingface.co/spaces/DiabetesCareChatbot/dmChatbotBackend/resolve/main/src/core/model_manager.py
- Command line
-
hf download hf://spaces/DiabetesCareChatbot/dmChatbotBackend/src/core/model_manager.py
-
curl -L -o model_manager.py https://huggingface.co/spaces/DiabetesCareChatbot/dmChatbotBackend/resolve/main/src/core/model_manager.py
38.6 kB
| import os | |
| import time | |
| import json | |
| import asyncio | |
| import uuid | |
| import requests | |
| import aiohttp | |
| from dotenv import load_dotenv | |
| from langchain_core.messages import AIMessage, SystemMessage, HumanMessage | |
| from src.utils.logger import setup_logger | |
| logger = setup_logger("ModelManager") | |
| load_dotenv() | |
| class ReqModel: | |
| def __init__(self, model: str, temperature: float, base_url: str, api_key: str, headers: dict): | |
| self.model = model | |
| self.temperature = temperature | |
| self.base_url = base_url | |
| self.api_key = api_key | |
| self.headers = headers | |
| self.bound_tools = None | |
| def bind_tools(self, tools, **kwargs): | |
| new_model = ReqModel(self.model, self.temperature, self.base_url, self.api_key, self.headers) | |
| new_model.bound_tools = tools | |
| return new_model | |
| def _convert_messages(self, messages): | |
| req_msgs = [] | |
| for m in messages: | |
| if isinstance(m, SystemMessage): | |
| req_msgs.append({"role": "system", "content": m.content}) | |
| elif isinstance(m, HumanMessage): | |
| req_msgs.append({"role": "user", "content": m.content}) | |
| elif isinstance(m, AIMessage): | |
| req_msgs.append({"role": "assistant", "content": m.content}) | |
| elif isinstance(m, dict) and "role" in m and "content" in m: | |
| req_msgs.append(m) | |
| else: | |
| req_msgs.append({"role": "user", "content": str(getattr(m, 'content', m))}) | |
| return req_msgs | |
| def _format_tools(self): | |
| if not self.bound_tools: | |
| return None | |
| tools_list = [] | |
| for tool in self.bound_tools: | |
| if hasattr(tool, "name") and hasattr(tool, "description") and hasattr(tool, "args_schema"): | |
| tools_list.append({ | |
| "type": "function", | |
| "function": { | |
| "name": tool.name, | |
| "description": tool.description, | |
| "parameters": tool.args_schema.schema() if tool.args_schema else {"type": "object", "properties": {}} | |
| } | |
| }) | |
| return tools_list | |
| def _make_request(self, messages, config=None, **kwargs): | |
| url = f"{self.base_url}/chat/completions" | |
| headers = { | |
| "Authorization": f"Bearer {self.api_key}", | |
| "Content-Type": "application/json" | |
| } | |
| headers.update(self.headers) | |
| payload = { | |
| "model": self.model, | |
| "messages": self._convert_messages(messages), | |
| "temperature": self.temperature, | |
| "max_tokens": 4096, | |
| } | |
| formatted_tools = self._format_tools() | |
| if formatted_tools: | |
| payload["tools"] = formatted_tools | |
| response = requests.post(url, headers=headers, json=payload) | |
| response.raise_for_status() | |
| data = response.json() | |
| message = data["choices"][0]["message"] | |
| content = message.get("content", "") | |
| ai_message = AIMessage(content=content if content else "") | |
| if "tool_calls" in message and message["tool_calls"]: | |
| tool_calls = [] | |
| for tc in message["tool_calls"]: | |
| try: | |
| args = json.loads(tc["function"]["arguments"]) | |
| except Exception: | |
| args = {} | |
| tool_calls.append({ | |
| "name": tc["function"]["name"], | |
| "args": args, | |
| "id": tc["id"] | |
| }) | |
| ai_message.additional_kwargs["tool_calls"] = message["tool_calls"] | |
| ai_message.tool_calls = tool_calls | |
| elif content and isinstance(content, str): | |
| stripped = content.strip() | |
| if stripped.startswith("{") and stripped.endswith("}"): | |
| try: | |
| parsed_tc = json.loads(stripped) | |
| if isinstance(parsed_tc, dict) and "name" in parsed_tc and ("parameters" in parsed_tc or "arguments" in parsed_tc): | |
| args = parsed_tc.get("parameters") or parsed_tc.get("arguments") or {} | |
| call_id = f"call_{uuid.uuid4().hex[:8]}" | |
| ai_message.tool_calls = [{ | |
| "name": parsed_tc["name"], | |
| "args": args, | |
| "id": call_id | |
| }] | |
| ai_message.additional_kwargs["tool_calls"] = [{ | |
| "id": call_id, | |
| "type": "function", | |
| "function": { | |
| "name": parsed_tc["name"], | |
| "arguments": json.dumps(args) | |
| } | |
| }] | |
| ai_message.content = "" | |
| except Exception: | |
| pass | |
| ai_message.response_metadata = {"token_usage": data.get("usage", {})} | |
| return ai_message | |
| def invoke(self, messages, config=None, **kwargs): | |
| return self._make_request(messages, config, **kwargs) | |
| async def ainvoke(self, messages, config=None, **kwargs): | |
| loop = asyncio.get_event_loop() | |
| return await loop.run_in_executor(None, lambda: self._make_request(messages, config, **kwargs)) | |
| def stream(self, messages, config=None, **kwargs): | |
| from langchain_core.messages import AIMessageChunk | |
| url = f"{self.base_url}/chat/completions" | |
| headers = { | |
| "Authorization": f"Bearer {self.api_key}", | |
| "Content-Type": "application/json" | |
| } | |
| headers.update(self.headers) | |
| payload = { | |
| "model": self.model, | |
| "messages": self._convert_messages(messages), | |
| "temperature": self.temperature, | |
| "stream": True, | |
| "max_tokens": 4096, | |
| } | |
| formatted_tools = self._format_tools() | |
| if formatted_tools: | |
| payload["tools"] = formatted_tools | |
| with requests.post(url, headers=headers, json=payload, stream=True) as response: | |
| response.raise_for_status() | |
| for line in response.iter_lines(): | |
| if line: | |
| text = line.decode('utf-8').strip() | |
| if text.startswith("data: "): | |
| data_str = text[6:] | |
| if data_str == "[DONE]": | |
| break | |
| try: | |
| data = json.loads(data_str) | |
| choices = data.get("choices", []) | |
| if choices: | |
| delta = choices[0].get("delta", {}) | |
| content = delta.get("content", "") | |
| if content: | |
| yield AIMessageChunk(content=content) | |
| except Exception: | |
| pass | |
| async def astream(self, messages, config=None, **kwargs): | |
| from langchain_core.messages import AIMessageChunk | |
| url = f"{self.base_url}/chat/completions" | |
| headers = { | |
| "Authorization": f"Bearer {self.api_key}", | |
| "Content-Type": "application/json" | |
| } | |
| headers.update(self.headers) | |
| payload = { | |
| "model": self.model, | |
| "messages": self._convert_messages(messages), | |
| "temperature": self.temperature, | |
| "stream": True, | |
| "max_tokens": 4096, | |
| } | |
| formatted_tools = self._format_tools() | |
| if formatted_tools: | |
| payload["tools"] = formatted_tools | |
| yielded_any = False | |
| try: | |
| async with aiohttp.ClientSession() as session: | |
| async with session.post(url, headers=headers, json=payload) as response: | |
| response.raise_for_status() | |
| async for line in response.content: | |
| if line: | |
| text = line.decode('utf-8').strip() | |
| if text.startswith("data: "): | |
| data_str = text[6:] | |
| if data_str == "[DONE]": | |
| break | |
| try: | |
| data = json.loads(data_str) | |
| choices = data.get("choices", []) | |
| if choices: | |
| delta = choices[0].get("delta", {}) | |
| content = delta.get("content", "") | |
| if content: | |
| yielded_any = True | |
| yield AIMessageChunk(content=content) | |
| except Exception: | |
| pass | |
| except Exception as e: | |
| logger.error(f"Async stream request error: {e}") | |
| if not yielded_any: | |
| raise | |
| yield AIMessageChunk(content="") | |
| class RateLimitFallbackWrapper: | |
| """Wraps a primary LLM and a list of fallback LLMs. When the primary | |
| returns a rate-limit or availability error (HTTP 429 / 503), the wrapper | |
| automatically retries with the next fallback provider — no local rate | |
| limiting needed. | |
| """ | |
| # HTTP status codes that indicate a rate-limit / temporary-unavailability | |
| _RATE_LIMIT_CODES = {429, 503} | |
| # Maximum seconds to honour a Retry-After header before giving up | |
| _MAX_BACKOFF = 5 | |
| def __init__(self, main_llm, fallback_llms): | |
| self.main_llm = main_llm | |
| self.fallback_llms = fallback_llms | |
| self.bound_tools = None | |
| # ------------------------------------------------------------------ | |
| # Helpers | |
| # ------------------------------------------------------------------ | |
| def _is_rate_limit_error(exc: Exception) -> bool: | |
| """Return True if *exc* looks like an HTTP 429 or 503.""" | |
| # requests.HTTPError carries a response object | |
| if hasattr(exc, "response") and hasattr(exc.response, "status_code"): | |
| return exc.response.status_code in RateLimitFallbackWrapper._RATE_LIMIT_CODES | |
| # Some libraries embed the status code in the message | |
| msg = str(exc) | |
| return "429" in msg or "503" in msg or "rate" in msg.lower() | |
| def _get_retry_after(exc: Exception) -> float: | |
| """Extract a Retry-After value (in seconds) from the exception, if present.""" | |
| if hasattr(exc, "response") and hasattr(exc.response, "headers"): | |
| retry = exc.response.headers.get("Retry-After") | |
| if retry: | |
| try: | |
| return min(float(retry), RateLimitFallbackWrapper._MAX_BACKOFF) | |
| except (ValueError, TypeError): | |
| pass | |
| return 0.0 | |
| def _unavailable_message(self): | |
| return AIMessage( | |
| content=( | |
| "I'm sorry — all configured AI model providers are temporarily unavailable right now. " | |
| "This may be due to rate-limit errors, invalid API keys, or provider outages. " | |
| "Please wait a minute and try again, or check your API key and provider settings in the .env file." | |
| ) | |
| ) | |
| def _log_error(self, label: str, exc: Exception): | |
| if self._is_rate_limit_error(exc): | |
| logger.warning(f"{label}: Rate-limit / unavailability error — {exc}") | |
| else: | |
| logger.warning(f"{label}: {exc}") | |
| def bind_tools(self, tools, **kwargs): | |
| new_main = self.main_llm.bind_tools(tools, **kwargs) | |
| new_falls = [llm.bind_tools(tools, **kwargs) for llm in self.fallback_llms] | |
| new_wrapper = RateLimitFallbackWrapper(new_main, new_falls) | |
| new_wrapper.bound_tools = tools | |
| return new_wrapper | |
| # ------------------------------------------------------------------ | |
| # invoke / ainvoke | |
| # ------------------------------------------------------------------ | |
| def invoke(self, messages, config=None, **kwargs): | |
| all_llms = [self.main_llm] + list(self.fallback_llms) | |
| for idx, llm in enumerate(all_llms): | |
| label = f"Provider {idx+1} ({getattr(llm, 'model', '?')})" | |
| try: | |
| logger.info(f"Trying {label}") | |
| return llm.invoke(messages, config=config, **kwargs) | |
| except Exception as exc: | |
| self._log_error(label, exc) | |
| # Brief backoff on rate-limit before trying the next provider | |
| backoff = self._get_retry_after(exc) | |
| if backoff > 0: | |
| logger.info(f" Backing off {backoff:.1f}s before next provider") | |
| time.sleep(backoff) | |
| logger.warning("All providers/models failed; returning a graceful fallback response.") | |
| return self._unavailable_message() | |
| async def ainvoke(self, messages, config=None, **kwargs): | |
| all_llms = [self.main_llm] + list(self.fallback_llms) | |
| for idx, llm in enumerate(all_llms): | |
| label = f"Provider {idx+1} ({getattr(llm, 'model', '?')})" | |
| try: | |
| logger.info(f"Trying {label}") | |
| return await llm.ainvoke(messages, config=config, **kwargs) | |
| except Exception as exc: | |
| self._log_error(label, exc) | |
| backoff = self._get_retry_after(exc) | |
| if backoff > 0: | |
| logger.info(f" Backing off {backoff:.1f}s before next provider") | |
| await asyncio.sleep(backoff) | |
| logger.warning("All providers/models failed; returning a graceful fallback response.") | |
| return self._unavailable_message() | |
| # ------------------------------------------------------------------ | |
| # stream / astream | |
| # ------------------------------------------------------------------ | |
| def stream(self, messages, config=None, **kwargs): | |
| all_llms = [self.main_llm] + list(self.fallback_llms) | |
| for idx, llm in enumerate(all_llms): | |
| label = f"Provider {idx+1} ({getattr(llm, 'model', '?')})" | |
| try: | |
| logger.info(f"Trying {label} for streaming") | |
| yield from llm.stream(messages, config=config, **kwargs) | |
| return | |
| except Exception as exc: | |
| self._log_error(label, exc) | |
| backoff = self._get_retry_after(exc) | |
| if backoff > 0: | |
| logger.info(f" Backing off {backoff:.1f}s before next provider") | |
| time.sleep(backoff) | |
| logger.warning("All providers/models failed while streaming; returning a graceful fallback message.") | |
| yield self._unavailable_message() | |
| async def astream(self, messages, config=None, **kwargs): | |
| all_llms = [self.main_llm] + list(self.fallback_llms) | |
| for idx, llm in enumerate(all_llms): | |
| label = f"Provider {idx+1} ({getattr(llm, 'model', '?')})" | |
| try: | |
| logger.info(f"Trying {label} for async streaming") | |
| received_chunk = False | |
| async for chunk in llm.astream(messages, config=config, **kwargs): | |
| if getattr(chunk, "content", None): | |
| received_chunk = True | |
| yield chunk | |
| if received_chunk: | |
| return | |
| except Exception as exc: | |
| self._log_error(label, exc) | |
| backoff = self._get_retry_after(exc) | |
| if backoff > 0: | |
| logger.info(f" Backing off {backoff:.1f}s before next provider") | |
| await asyncio.sleep(backoff) | |
| logger.warning("All providers/models failed while async streaming; returning a graceful fallback message.") | |
| yield self._unavailable_message() | |
| # Ordered list of cloud LLM providers to check for availability. | |
| # The first provider with a valid API key becomes the primary; the rest are fallbacks. | |
| CLOUD_PROVIDERS = ["truefoundry", "openrouter"] | |
| class ModelManager: | |
| _openrouter_accessible_cache = None | |
| _openrouter_cache_time = 0.0 | |
| _CACHE_TTL = 300.0 # 5 minutes | |
| def __init__(self, model_name: str = "openai-main/gpt-4o-mini"): | |
| self.provider = self._normalize_provider(os.getenv("MODEL_PROVIDER", "truefoundry")) | |
| self.selected_model = None | |
| # Keep provider-specific model settings separate. | |
| if self.provider in ("ollama", "local_ollama", "local", "both", "localboth"): | |
| self.model_name = os.getenv("OLLAMA_MODEL_NAME", model_name) | |
| elif self.provider == "lmstudio": | |
| self.model_name = os.getenv("LM_STUDIO_MODEL_NAME", model_name) | |
| elif self.provider == "truefoundry": | |
| self.model_name = os.getenv("TRUEFOUNDRY_MODEL_NAME", "openai-main/gpt-4o-mini") | |
| elif self.provider == "portkey": | |
| self.model_name = os.getenv("PORTKEY_MODEL_NAME", "gpt-4o-mini") | |
| else: | |
| self.model_name = os.getenv("OPENROUTER_MODEL_NAME", "nvidia/nemotron-3-ultra-550b-a55b:free") | |
| def _normalize_provider(provider: str) -> str: | |
| normalized = (provider or "").strip().lower().replace("-", "_") | |
| return { | |
| "lm_studio": "lmstudio", | |
| "local_lm_studio": "lmstudio", | |
| }.get(normalized, normalized) | |
| def ping_openrouter_accessible_models(self, api_key: str = None) -> list[str]: | |
| """Dynamically ping OpenRouter to discover models that have access. | |
| Checks key access/tier on /auth/key, queries available models on /models, | |
| and pings candidate models with a minimal test request to verify they are | |
| reachable, online, and not blocked by guardrails. Results are cached. | |
| """ | |
| key = api_key or os.getenv("OPENROUTER_API_KEY") or os.getenv("OPENROUTER_SECONDARY_API_KEY") | |
| if not key: | |
| return [] | |
| now = time.time() | |
| if self._openrouter_accessible_cache and (now - self._openrouter_cache_time < self._CACHE_TTL): | |
| return list(self._openrouter_accessible_cache) | |
| headers = { | |
| "Authorization": f"Bearer {key}", | |
| "HTTP-Referer": "https://github.com/Sudharshan-3904/dmChatbot", | |
| "X-Title": "Medical AI Chatbot", | |
| } | |
| # 1. Determine key tier (free vs paid) | |
| is_free_tier = True | |
| try: | |
| auth_res = requests.get("https://openrouter.ai/api/v1/auth/key", headers=headers, timeout=5) | |
| if auth_res.status_code == 200: | |
| key_info = auth_res.json().get("data", {}) | |
| is_free_tier = bool(key_info.get("is_free_tier", True)) | |
| except Exception as exc: | |
| logger.warning(f"Error checking OpenRouter key tier: {exc}") | |
| # 2. Fetch models catalog | |
| catalog = [] | |
| try: | |
| models_res = requests.get("https://openrouter.ai/api/v1/models", headers=headers, timeout=8) | |
| if models_res.status_code == 200: | |
| catalog = models_res.json().get("data", []) | |
| except Exception as exc: | |
| logger.warning(f"Error querying OpenRouter models catalog: {exc}") | |
| # 3. Filter candidate models by access | |
| candidates = [] | |
| for m in catalog: | |
| mid = m.get("id", "") | |
| if not mid: | |
| continue | |
| if is_free_tier: | |
| pricing = m.get("pricing", {}) | |
| try: | |
| p_cost = float(pricing.get("prompt", 1)) | |
| c_cost = float(pricing.get("completion", 1)) | |
| except (ValueError, TypeError): | |
| p_cost, c_cost = 1, 1 | |
| if (p_cost == 0 and c_cost == 0) or ":free" in mid or mid == "openrouter/free": | |
| candidates.append(mid) | |
| else: | |
| candidates.append(mid) | |
| # 4. Preferred order of known reliable models to test first | |
| preferred_order = [ | |
| "nvidia/nemotron-3-ultra-550b-a55b:free", | |
| "nvidia/nemotron-3-super-120b-a12b:free", | |
| "openrouter/free", | |
| "google/gemma-4-26b-a4b-it:free", | |
| "google/gemma-4-31b-it:free", | |
| ] | |
| ordered_candidates = [m for m in preferred_order if m in candidates] + [ | |
| m for m in candidates if m not in preferred_order | |
| ] | |
| # 5. Dynamically test/ping candidate models to verify accessibility | |
| verified = [] | |
| for mid in ordered_candidates: | |
| try: | |
| ping_res = requests.post( | |
| "https://openrouter.ai/api/v1/chat/completions", | |
| headers=headers, | |
| json={ | |
| "model": mid, | |
| "messages": [{"role": "user", "content": "ping"}], | |
| "max_tokens": 1, | |
| }, | |
| timeout=3, | |
| ) | |
| if ping_res.status_code == 200: | |
| verified.append(mid) | |
| logger.info(f"Dynamically verified OpenRouter model with access: {mid}") | |
| if len(verified) >= 3: | |
| break | |
| else: | |
| logger.debug(f"OpenRouter model ping failed for {mid} (status={ping_res.status_code})") | |
| except Exception as exc: | |
| logger.debug(f"OpenRouter model ping error for {mid}: {exc}") | |
| if not verified: | |
| verified = ordered_candidates[:3] if ordered_candidates else [ | |
| "nvidia/nemotron-3-ultra-550b-a55b:free", | |
| "openrouter/free", | |
| ] | |
| self._openrouter_accessible_cache = verified | |
| self._openrouter_cache_time = now | |
| logger.info(f"OpenRouter accessible models: {verified}") | |
| return verified | |
| def _build_local_llm(self, base_url: str, model_name: str, temperature: float): | |
| return ReqModel( | |
| model=model_name, | |
| temperature=temperature, | |
| base_url=base_url.rstrip("/"), | |
| api_key=os.getenv("LOCAL_LLM_API_KEY", "local-model"), | |
| headers={}, | |
| ) | |
| def _build_local_wrappers(self, temperature: float, model_name: str): | |
| provider = self.provider | |
| ollama_base = os.getenv("OLLAMA_BASE_URL", "http://localhost:11434/v1") | |
| lm_studio_base = os.getenv("LM_STUDIO_BASE_URL", "http://localhost:1234/v1") | |
| ollama_model = self.selected_model if provider in ("ollama", "local_ollama", "local", "both", "localboth") else os.getenv("OLLAMA_MODEL_NAME") | |
| if not ollama_model: | |
| if provider in ("ollama", "local_ollama", "local", "both", "localboth"): | |
| try: | |
| from src.core.ollama_manager import ollama_manager | |
| available_models = ollama_manager.get_model_names() | |
| if available_models: | |
| # Preference order of standard models | |
| preferred_models = [ | |
| "llama3.2:latest", "llama3.2", | |
| "qwen2.5:1.5b", "qwen2.5", | |
| "gemma4:12b", "gemma4", | |
| "ministral-3:latest", "ministral-3", | |
| "granite4.1:8b", "granite4.1", | |
| "ornith-1.5:9b", "ornith-1.5" | |
| ] | |
| for pref in preferred_models: | |
| if pref in available_models: | |
| ollama_model = pref | |
| break | |
| if not ollama_model: | |
| ollama_model = available_models[0] | |
| logger.info(f"Ollama local model name not configured; auto-selected: {ollama_model} from available models {available_models}") | |
| except Exception as e: | |
| logger.warning(f"Error listing local Ollama models: {e}") | |
| if not ollama_model: | |
| ollama_model = model_name | |
| lm_studio_model = self.selected_model if provider == "lmstudio" else os.getenv("LM_STUDIO_MODEL_NAME", model_name) | |
| if provider in ("ollama", "local_ollama", "local"): | |
| main_llm = self._build_local_llm(ollama_base, ollama_model, temperature) | |
| fallback_llms = [self._build_local_llm(lm_studio_base, lm_studio_model, temperature)] if provider == "local" else [] | |
| return RateLimitFallbackWrapper(main_llm, fallback_llms) | |
| if provider == "lmstudio": | |
| main_llm = self._build_local_llm(lm_studio_base, lm_studio_model, temperature) | |
| return RateLimitFallbackWrapper(main_llm, [self._build_local_llm(ollama_base, ollama_model, temperature)]) | |
| if provider in ("both", "localboth"): | |
| main_llm = self._build_local_llm(ollama_base, ollama_model, temperature) | |
| fallback_llms = [self._build_local_llm(lm_studio_base, lm_studio_model, temperature)] | |
| return RateLimitFallbackWrapper(main_llm, fallback_llms) | |
| return None | |
| def _get_openrouter_api_keys(self): | |
| primary_api_key = os.getenv("OPENROUTER_API_KEY") | |
| secondary_api_key = os.getenv("OPENROUTER_SECONDARY_API_KEY") | |
| if not primary_api_key and secondary_api_key: | |
| logger.warning("OPENROUTER_API_KEY missing; using OPENROUTER_SECONDARY_API_KEY as the active key.") | |
| primary_api_key = secondary_api_key | |
| secondary_api_key = None | |
| if not primary_api_key: | |
| raise EnvironmentError( | |
| "OpenRouter requires OPENROUTER_API_KEY or OPENROUTER_SECONDARY_API_KEY in the environment." | |
| ) | |
| return primary_api_key, secondary_api_key | |
| def _get_provider_config(self, provider: str): | |
| """Return (api_key, base_url, model_name, extra_headers) for a cloud provider. | |
| Returns ``None`` for api_key when the provider is not configured. | |
| """ | |
| if provider == "openrouter": | |
| api_key = os.getenv("OPENROUTER_API_KEY") or os.getenv("OPENROUTER_SECONDARY_API_KEY") | |
| base_url = "https://openrouter.ai/api/v1" | |
| model = os.getenv("OPENROUTER_MODEL_NAME", self.model_name) | |
| headers = { | |
| "HTTP-Referer": "https://github.com/Sudharshan-3904/dmChatbot", | |
| "X-Title": "Medical AI Chatbot", | |
| } | |
| return api_key, base_url, model, headers | |
| if provider == "truefoundry": | |
| api_key = os.getenv("TRUEFOUNDRY_API_KEY") | |
| base_url = os.getenv("TRUEFOUNDRY_BASE_URL", "https://llm-gateway.truefoundry.com/api/inference/openai").rstrip("/") | |
| model = os.getenv("TRUEFOUNDRY_MODEL_NAME", "openai-main/gpt-4o-mini") | |
| return api_key, base_url, model, {} | |
| if provider == "portkey": | |
| api_key = os.getenv("PORTKEY_API_KEY") | |
| base_url = os.getenv("PORTKEY_BASE_URL", "https://api.portkey.ai/v1").rstrip("/") | |
| model = os.getenv("PORTKEY_MODEL_NAME", "gpt-4o-mini") | |
| headers = {} | |
| virtual_key = os.getenv("PORTKEY_VIRTUAL_KEY") | |
| if virtual_key: | |
| headers["x-portkey-virtual-key"] = virtual_key | |
| return api_key, base_url, model, headers | |
| return None, None, None, {} | |
| def _build_cloud_provider_llm(self, provider: str, temperature: float, model_override: str = None): | |
| """Build a ``ReqModel`` for a cloud provider. Returns ``None`` if the | |
| provider has no API key configured.""" | |
| api_key, base_url, default_model, extra_headers = self._get_provider_config(provider) | |
| if not api_key: | |
| return None | |
| model = model_override or default_model | |
| return ReqModel( | |
| model=model, | |
| temperature=temperature, | |
| base_url=base_url, | |
| api_key=api_key, | |
| headers=extra_headers, | |
| ) | |
| def get_default_model_for_provider(self, provider: str) -> str: | |
| """Return the one specific model configured for a given provider.""" | |
| p = provider.strip().lower() | |
| if p in ("ollama", "local_ollama", "local", "both", "localboth"): | |
| return os.getenv("OLLAMA_MODEL_NAME", "llama3.2:latest") | |
| elif p == "lmstudio": | |
| return os.getenv("LM_STUDIO_MODEL_NAME", "local-model") | |
| elif p == "truefoundry": | |
| return os.getenv("TRUEFOUNDRY_MODEL_NAME", "openai-main/gpt-4o-mini") | |
| elif p == "portkey": | |
| return os.getenv("PORTKEY_MODEL_NAME", "gpt-4o-mini") | |
| else: | |
| return os.getenv("OPENROUTER_MODEL_NAME", "nvidia/nemotron-3-ultra-550b-a55b:free") | |
| def get_configured_providers(self): | |
| """Return an ordered list of cloud providers with their availability status.""" | |
| result = [] | |
| for p in CLOUD_PROVIDERS: | |
| api_key, base_url, model, _ = self._get_provider_config(p) | |
| result.append({ | |
| "provider": p, | |
| "configured": bool(api_key), | |
| "base_url": base_url, | |
| "model": model or self.get_default_model_for_provider(p), | |
| }) | |
| return result | |
| def get_available_providers(self): | |
| """Return all available providers with their specific model and availability status.""" | |
| providers = [] | |
| name_map = { | |
| "openrouter": "OpenRouter", | |
| "truefoundry": "TrueFoundry", | |
| "ollama": "Ollama", | |
| "lmstudio": "LM Studio", | |
| "portkey": "Portkey", | |
| } | |
| # Cloud providers | |
| for p in CLOUD_PROVIDERS: | |
| api_key, base_url, model, _ = self._get_provider_config(p) | |
| providers.append({ | |
| "provider": p, | |
| "name": name_map.get(p, p.title()), | |
| "configured": bool(api_key), | |
| "base_url": base_url, | |
| "model": model or self.get_default_model_for_provider(p), | |
| }) | |
| # Local providers | |
| for p in ("ollama", "lmstudio"): | |
| providers.append({ | |
| "provider": p, | |
| "name": name_map.get(p, p.title()), | |
| "configured": True, | |
| "base_url": os.getenv("OLLAMA_BASE_URL" if p == "ollama" else "LM_STUDIO_BASE_URL", ""), | |
| "model": self.get_default_model_for_provider(p), | |
| }) | |
| return providers | |
| def get_llm(self, temperature: float = 0, model_name: str = None): | |
| model = self.selected_model or model_name or self.model_name | |
| logger.info(f"Initializing LLM: Provider={self.provider}, Model={model}") | |
| # ---------- local providers (Ollama / LM Studio) ---------- | |
| local_wrapper = self._build_local_wrappers(temperature, model) | |
| if local_wrapper is not None: | |
| return local_wrapper | |
| # ---------- cloud providers (ordered fallback chain) ---------- | |
| # If the configured provider is one of the cloud providers, start | |
| # with it; otherwise use the default priority order. | |
| if self.provider in CLOUD_PROVIDERS: | |
| ordered = [self.provider] + [p for p in CLOUD_PROVIDERS if p != self.provider] | |
| elif self.provider == "portkey": | |
| ordered = ["portkey"] + list(CLOUD_PROVIDERS) | |
| else: | |
| ordered = list(CLOUD_PROVIDERS) | |
| cloud_llms: list[ReqModel] = [] | |
| for provider in ordered: | |
| llm = self._build_cloud_provider_llm(provider, temperature, model if provider == self.provider else None) | |
| if llm is not None: | |
| cloud_llms.append(llm) | |
| logger.info(f" ✓ {provider} configured (model={llm.model})") | |
| else: | |
| logger.info(f" ✗ {provider} skipped — no API key") | |
| # For OpenRouter specifically, also add the secondary-key and | |
| # fallback-model variants (preserving existing behaviour). | |
| primary_or_key = os.getenv("OPENROUTER_API_KEY") | |
| secondary_key = os.getenv("OPENROUTER_SECONDARY_API_KEY") | |
| if primary_or_key and secondary_key and secondary_key != primary_or_key: | |
| cloud_llms.append( | |
| ReqModel( | |
| model=model, | |
| temperature=temperature, | |
| base_url="https://openrouter.ai/api/v1", | |
| api_key=secondary_key, | |
| headers={ | |
| "HTTP-Referer": "https://github.com/Sudharshan-3904/dmChatbot", | |
| "X-Title": "Medical AI Chatbot", | |
| }, | |
| ) | |
| ) | |
| or_key = primary_or_key or secondary_key | |
| if or_key: | |
| accessible_models = self.ping_openrouter_accessible_models(or_key) | |
| fallback_models = [m for m in accessible_models if m != model] | |
| cloud_llms.extend( | |
| ReqModel( | |
| model=m, | |
| temperature=temperature, | |
| base_url="https://openrouter.ai/api/v1", | |
| api_key=or_key, | |
| headers={ | |
| "HTTP-Referer": "https://github.com/Sudharshan-3904/dmChatbot", | |
| "X-Title": "Medical AI Chatbot", | |
| }, | |
| ) | |
| for m in fallback_models | |
| ) | |
| if not cloud_llms: | |
| raise EnvironmentError( | |
| "No cloud LLM provider is configured. " | |
| "Set at least one of OPENROUTER_API_KEY, TRUEFOUNDRY_API_KEY, or PORTKEY_API_KEY in your .env file." | |
| ) | |
| main_llm = cloud_llms[0] | |
| fallback_llms = cloud_llms[1:] | |
| logger.info(f"Cloud LLM chain: primary={main_llm.model} + {len(fallback_llms)} fallback(s)") | |
| return RateLimitFallbackWrapper(main_llm, fallback_llms) | |
| def list_available_models(self, provider: str = None): | |
| """Return normalized model metadata from a provider's models endpoint.""" | |
| provider = (provider or self.provider).strip().lower() | |
| endpoints = { | |
| "openrouter": ("https://openrouter.ai/api/v1/models", os.getenv("OPENROUTER_API_KEY")), | |
| "truefoundry": ( | |
| f"{os.getenv('TRUEFOUNDRY_BASE_URL', 'https://llm-gateway.truefoundry.com/api/inference/openai').rstrip('/')}/models", | |
| os.getenv("TRUEFOUNDRY_API_KEY"), | |
| ), | |
| "portkey": ( | |
| f"{os.getenv('PORTKEY_BASE_URL', 'https://api.portkey.ai/v1').rstrip('/')}/models", | |
| os.getenv("PORTKEY_API_KEY"), | |
| ), | |
| "ollama": (f"{os.getenv('OLLAMA_BASE_URL', 'http://localhost:11434/v1').rstrip('/')}/models", os.getenv("LOCAL_LLM_API_KEY", "local-model")), | |
| "lmstudio": (f"{os.getenv('LM_STUDIO_BASE_URL', 'http://localhost:1234/v1').rstrip('/')}/models", os.getenv("LOCAL_LLM_API_KEY", "local-model")), | |
| } | |
| if provider not in endpoints: | |
| raise ValueError( | |
| f"Unsupported MODEL_PROVIDER '{provider}'. " | |
| f"Expected one of: {', '.join(sorted(endpoints.keys()))}." | |
| ) | |
| endpoint, api_key = endpoints[provider] | |
| headers = {"Accept": "application/json"} | |
| if api_key: | |
| headers["Authorization"] = f"Bearer {api_key}" | |
| # Portkey may require the virtual key header | |
| if provider == "portkey": | |
| virtual_key = os.getenv("PORTKEY_VIRTUAL_KEY") | |
| if virtual_key: | |
| headers["x-portkey-virtual-key"] = virtual_key | |
| try: | |
| response = requests.get(endpoint, headers=headers, timeout=8) | |
| response.raise_for_status() | |
| payload = response.json() | |
| raw_data = payload.get("data", []) | |
| if provider == "openrouter": | |
| accessible_ids = set(self.ping_openrouter_accessible_models(api_key)) | |
| is_free_tier = True | |
| try: | |
| auth_res = requests.get("https://openrouter.ai/api/v1/auth/key", headers=headers, timeout=5) | |
| if auth_res.status_code == 200: | |
| is_free_tier = bool(auth_res.json().get("data", {}).get("is_free_tier", True)) | |
| except Exception: | |
| pass | |
| filtered = [] | |
| for item in raw_data: | |
| mid = item.get("id") or item.get("name") | |
| if not mid: | |
| continue | |
| if is_free_tier: | |
| pricing = item.get("pricing", {}) | |
| try: | |
| p_cost = float(pricing.get("prompt", 1)) | |
| c_cost = float(pricing.get("completion", 1)) | |
| except (ValueError, TypeError): | |
| p_cost, c_cost = 1, 1 | |
| if (p_cost == 0 and c_cost == 0) or ":free" in mid or mid == "openrouter/free": | |
| filtered.append(item) | |
| else: | |
| filtered.append(item) | |
| filtered.sort(key=lambda x: (0 if (x.get("id") or x.get("name")) in accessible_ids else 1)) | |
| raw_data = filtered | |
| models = [] | |
| for item in raw_data: | |
| model_id = item.get("id") or item.get("name") | |
| if model_id: | |
| models.append({ | |
| "id": model_id, | |
| "name": item.get("name") or model_id, | |
| "owned_by": item.get("owned_by") or provider, | |
| }) | |
| return models | |
| except requests.exceptions.RequestException as exc: | |
| logger.warning("Unable to list %s models: %s", provider, exc) | |
| return [] | |
| def select_model(self, provider: str, model: str = None): | |
| """Set the provider and optional model used by newly created agents.""" | |
| normalized_provider = self._normalize_provider(provider) | |
| valid_providers = {"openrouter", "truefoundry", "portkey", "ollama", "lmstudio"} | |
| if normalized_provider not in valid_providers: | |
| raise ValueError(f"provider must be one of: {', '.join(sorted(valid_providers))}") | |
| self.provider = normalized_provider | |
| default_model = self.get_default_model_for_provider(normalized_provider) | |
| self.model_name = default_model | |
| self.selected_model = model.strip() if isinstance(model, str) and model.strip() else default_model | |
| model_manager = ModelManager() | |