dmChatbotBackend / src /core /model_manager.py
github-actions
Auto deploy from GitHub
cb505ff
Raw History Blame Contribute Delete
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
# ------------------------------------------------------------------
@staticmethod
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()
@staticmethod
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")
@staticmethod
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()