Nandly_backend / backend /core /models.py
Anand Sharma
Production Hardening: Health check, Bug fixes, and Key fallback
1207478
Raw
History Blame Contribute Delete
3.14 kB
import os
from typing import Optional, Union, Any, List
from langchain_google_genai import ChatGoogleGenerativeAI
from langchain_openai import ChatOpenAI
from langchain_anthropic import ChatAnthropic
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage, AIMessage
def get_llm(
provider: str,
model: str,
api_key: str,
temperature: float = 0.7,
streaming: bool = True
):
"""
Factory function to get a unified LLM instance based on provider.
Supports DEV_MODE if api_key is missing.
"""
dev_mode = os.getenv("DEV_MODE", "false").lower() == "true"
# Fallback to server key if user key is missing (Prod & Dev)
target_api_key = api_key
if not api_key:
if provider == "google":
target_api_key = os.getenv("GOOGLE_API_KEY")
elif provider == "openai":
target_api_key = os.getenv("OPENAI_API_KEY")
elif provider == "anthropic":
target_api_key = os.getenv("ANTHROPIC_API_KEY")
elif provider == "nvidia":
target_api_key = os.getenv("NVIDIA_API_KEY")
if not target_api_key:
raise ValueError(f"Missing API Key for {provider}. Please ensure server secrets are set or provide a key in settings.")
if provider == "google":
return ChatGoogleGenerativeAI(
model=model,
google_api_key=target_api_key,
temperature=temperature,
streaming=streaming,
# Suppress common safety filters using the dictionary format
safety_settings={
"HARM_CATEGORY_HARASSMENT": "BLOCK_NONE",
"HARM_CATEGORY_HATE_SPEECH": "BLOCK_NONE",
"HARM_CATEGORY_SEXUALLY_EXPLICIT": "BLOCK_NONE",
"HARM_CATEGORY_DANGEROUS_CONTENT": "BLOCK_NONE",
}
)
elif provider == "openai":
return ChatOpenAI(
model=model,
api_key=target_api_key,
temperature=temperature,
streaming=streaming
)
elif provider == "anthropic":
return ChatAnthropic(
model=model,
api_key=target_api_key,
temperature=temperature,
streaming=streaming
)
elif provider == "nvidia":
return ChatOpenAI(
model=model,
api_key=target_api_key,
base_url=os.getenv("NVIDIA_BASE_URL", "https://integrate.api.nvidia.com/v1"),
temperature=temperature,
streaming=streaming
)
else:
raise ValueError(f"Unsupported provider: {provider}")
def convert_history_to_langchain(history: List[dict]):
"""Convert standard role/content dicts to LangChain message objects."""
messages = []
for msg in history:
role = msg.get("role", "user")
content = msg.get("content", "")
if role == "user":
messages.append(HumanMessage(content=content))
elif role == "assistant":
messages.append(AIMessage(content=content))
elif role == "system":
messages.append(SystemMessage(content=content))
return messages