| import os |
| import time |
| import json |
| import asyncio |
| import requests |
| 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, |
| } |
| |
| 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 |
|
|
| 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): |
| yield self._make_request(messages, config, **kwargs) |
|
|
| async def astream(self, messages, config=None, **kwargs): |
| loop = asyncio.get_event_loop() |
| response = await loop.run_in_executor(None, lambda: self._make_request(messages, config, **kwargs)) |
| yield response |
|
|
| class RateLimitFallbackWrapper: |
| def __init__(self, main_llm, fallback_llms): |
| self.main_llm = main_llm |
| self.fallback_llms = fallback_llms |
| self.bound_tools = None |
|
|
| 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 |
|
|
| async def ainvoke(self, messages, config=None, **kwargs): |
| try: |
| return await self.main_llm.ainvoke(messages, config=config, **kwargs) |
| except Exception as e: |
| logger.warning(f"LLM Error with main model: {e}. Attempting fallbacks immediately.") |
|
|
| for idx, fb_llm in enumerate(self.fallback_llms): |
| try: |
| logger.info(f"Trying fallback model {idx+1} [Model: {fb_llm.model}]") |
| return await fb_llm.ainvoke(messages, config=config, **kwargs) |
| except Exception as fb_e: |
| logger.warning(f"Fallback {idx+1} failed: {fb_e}") |
|
|
| raise RuntimeError("All models (main and fallbacks) failed.") from None |
|
|
| def invoke(self, messages, config=None, **kwargs): |
| try: |
| return self.main_llm.invoke(messages, config=config, **kwargs) |
| except Exception as e: |
| logger.warning(f"LLM Error with main model: {e}. Attempting fallbacks immediately.") |
|
|
| for idx, fb_llm in enumerate(self.fallback_llms): |
| try: |
| logger.info(f"Trying fallback model {idx+1}") |
| return fb_llm.invoke(messages, config=config, **kwargs) |
| except Exception as fb_e: |
| logger.warning(f"Fallback {idx+1} failed: {fb_e}") |
|
|
| raise RuntimeError("All models failed synchronously.") from None |
|
|
| def stream(self, messages, config=None, **kwargs): |
| try: |
| yield from self.main_llm.stream(messages, config=config, **kwargs) |
| return |
| except Exception as e: |
| logger.warning(f"LLM stream error with main model: {e}. Attempting fallbacks immediately.") |
|
|
| for idx, fb_llm in enumerate(self.fallback_llms): |
| try: |
| logger.info(f"Trying fallback model {idx+1} for streaming") |
| yield from fb_llm.stream(messages, config=config, **kwargs) |
| return |
| except Exception as fb_e: |
| logger.warning(f"Fallback {idx+1} streaming failed: {fb_e}") |
|
|
| raise RuntimeError("All models failed while streaming.") |
|
|
| async def astream(self, messages, config=None, **kwargs): |
| try: |
| async for chunk in self.main_llm.astream(messages, config=config, **kwargs): |
| yield chunk |
| return |
| except Exception as e: |
| logger.warning(f"LLM async stream error with main model: {e}. Attempting fallbacks immediately.") |
|
|
| for idx, fb_llm in enumerate(self.fallback_llms): |
| try: |
| logger.info(f"Trying fallback model {idx+1} for async streaming") |
| async for chunk in fb_llm.astream(messages, config=config, **kwargs): |
| yield chunk |
| return |
| except Exception as fb_e: |
| logger.warning(f"Fallback {idx+1} async streaming failed: {fb_e}") |
|
|
| raise RuntimeError("All models failed while async streaming.") |
|
|
| class ModelManager: |
| def __init__(self, model_name: str = "google/gemma-4-26b-a4b-it:free"): |
| self.provider = os.getenv("MODEL_PROVIDER", "openrouter").lower() |
| self.model_name = os.getenv("OPENROUTER_MODEL_NAME", model_name) |
|
|
| 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_llm(self, temperature: float = 0, model_name: str = None): |
| model = model_name or self.model_name |
| logger.info(f"Initializing LLM: Provider={self.provider}, Model={model}") |
| |
| primary_api_key, secondary_api_key = self._get_openrouter_api_keys() |
| base_url = "https://openrouter.ai/api/v1" |
|
|
| main_llm = ReqModel( |
| model=model, |
| temperature=temperature, |
| base_url=base_url, |
| api_key=primary_api_key, |
| headers={ |
| "HTTP-Referer": "https://github.com/Sudharshan-3904/dmChatbot", |
| "X-Title": "Medical AI Chatbot" |
| } |
| ) |
|
|
| fallback_llms = [] |
| if secondary_api_key: |
| fallback_llms.append( |
| ReqModel( |
| model=model, |
| temperature=temperature, |
| base_url=base_url, |
| api_key=secondary_api_key, |
| headers={ |
| "HTTP-Referer": "https://github.com/Sudharshan-3904/dmChatbot", |
| "X-Title": "Medical AI Chatbot" |
| } |
| ) |
| ) |
|
|
| fallback_models = [ |
| "google/gemma-4-26b-a4b-it:free", |
| "google/gemma-4-31b-it:free", |
| "openai/gpt-oss-20b:free" |
| ] |
| |
| fallback_llms.extend([ |
| ReqModel( |
| model=m, |
| temperature=temperature, |
| base_url=base_url, |
| api_key=primary_api_key, |
| headers={ |
| "HTTP-Referer": "https://github.com/Sudharshan-3904/dmChatbot", |
| "X-Title": "Medical AI Chatbot" |
| } |
| ) for m in fallback_models if m != model |
| ]) |
| |
| return RateLimitFallbackWrapper(main_llm, fallback_llms) |
|
|
| model_manager = ModelManager() |
|
|