Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
151 changes: 148 additions & 3 deletions openlrc/chatbot.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,8 @@
return GPTBot, chatbot_model
elif chatbot_type == "anthropic":
return ClaudeBot, chatbot_model
elif chatbot_type == "litellm":
return LiteLLMBot, chatbot_model
else:
raise ValueError(f"Invalid chatbot type {chatbot_type}.")

Expand Down Expand Up @@ -260,7 +262,6 @@
api_key: str | None = None,
extra_body: dict | None = None,
):

# clamp temperature to 0-2
temperature = max(0, min(2, temperature))

Expand Down Expand Up @@ -459,7 +460,6 @@
api_key: str | None = None,
extra_body: dict | None = None,
):

# clamp temperature to 0-1
temperature = max(0, min(1, temperature))

Expand Down Expand Up @@ -774,4 +774,149 @@
self.client.close()


provider2chatbot = {ModelProvider.OPENAI: GPTBot, ModelProvider.ANTHROPIC: ClaudeBot, ModelProvider.GOOGLE: GeminiBot}
class LiteLLMBot(ChatBot):
"""ChatBot backed by the LiteLLM SDK, routing to 100+ LLM providers."""

def __init__(
self,
model_name: str = "openai/gpt-4o",
temperature: float = 1,
top_p: float = 1,
retry: int = 8,
max_async: int = 16,
fee_limit: float = 0.8,
proxy: str | None = None,
base_url_config: dict | None = None,
api_key: str | None = None,
extra_body: dict | None = None,
):
temperature = max(0, min(2, temperature))
super().__init__(model_name, temperature, top_p, retry, max_async, fee_limit)
self.api_key = api_key
self.api_base = (base_url_config or {}).get("litellm")
self.extra_body = dict(extra_body) if extra_body else {}
if proxy:
logger.warning(
"LiteLLMBot does not support per-instance proxy. "
"Set HTTP_PROXY/HTTPS_PROXY environment variables instead."
)

def update_fee(self, response):
usage = getattr(response, "usage", None)
if usage is None:
return
prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0
completion_tokens = getattr(usage, "completion_tokens", 0) or 0
self.api_fees[-1] += (
prompt_tokens * self.model_info.input_price + completion_tokens * self.model_info.output_price
) / 1000000

def get_content(self, response):
content = response.choices[0].message.content
return content if content is not None else ""

def _create_chat(
self,
messages: list[dict],
stop_sequences: list[str] | None = None,
output_checker: Callable = lambda user_input, generated_content: True,
temperature: float | None = None,
top_p: float | None = None,
):
try:
import litellm
except ImportError:
raise ImportError(
"litellm is required for the litellm: provider. Install with: pip install 'openlrc[litellm]'"
)

effective_temperature = temperature if temperature is not None else self.temperature
effective_top_p = top_p if top_p is not None else self.top_p

completion_kwargs: dict = {
"model": self.model_name,
"messages": messages,
"temperature": effective_temperature,
"max_tokens": self._compute_max_tokens(messages),
"drop_params": True,
}
completion_kwargs["top_p"] = effective_top_p
if stop_sequences:
completion_kwargs["stop"] = stop_sequences
if self.api_key:
completion_kwargs["api_key"] = self.api_key
if self.api_base:
completion_kwargs["api_base"] = self.api_base
completion_kwargs.update(self.extra_body)

response = None
validated = False
for i in range(self.retry):
try:
response = litellm.completion(**completion_kwargs)
self.update_fee(response)

if response.choices[0].finish_reason == "length":
usage = getattr(response, "usage", None)
raise LengthExceedException(
prompt_tokens=getattr(usage, "prompt_tokens", -1) if usage else -1,
completion_tokens=getattr(usage, "completion_tokens", -1) if usage else -1,
total_tokens=getattr(usage, "total_tokens", -1) if usage else -1,
)

response_text = remove_stop(self.get_content(response), stop_sequences)

if not output_checker(messages[-1]["content"], response_text):
logger.warning(f"Invalid response format. Retry num: {i + 1}.")
continue

validated = True
break
except litellm.AuthenticationError as e:
raise ChatBotException(f"Authentication failed: {e}") from e
except (
litellm.BadRequestError,
litellm.NotFoundError,
litellm.PermissionDeniedError,
litellm.UnprocessableEntityError,
) as e:
raise ChatBotException(f"Client error: {e}") from e
except (
litellm.RateLimitError,
litellm.APIConnectionError,
litellm.InternalServerError,
litellm.ServiceUnavailableError,
litellm.Timeout,
json.decoder.JSONDecodeError,
) as e:
sleep_time = self._get_sleep_time(e)
logger.warning(f"{type(e).__name__}: {e}. Wait {sleep_time}s before retry. Retry num: {i + 1}.")
time.sleep(sleep_time)

if not response:
raise ChatBotException("Failed to create a chat.")

if not validated:
logger.warning("Response format validation failed after all retries, returning best-effort response.")

return response

@staticmethod
def _get_sleep_time(error):
qualname = type(error).__name__
if qualname == "RateLimitError":
return random.randint(30, 60)

Check notice

Code scanning / Bandit

Standard pseudo-random generators are not suitable for security/cryptographic purposes. Note

Standard pseudo-random generators are not suitable for security/cryptographic purposes.
elif qualname in ("Timeout", "APIConnectionError"):
return 3
elif qualname == "JSONDecodeError":
return 1
else:
return 15


provider2chatbot = {
ModelProvider.OPENAI: GPTBot,
ModelProvider.ANTHROPIC: ClaudeBot,
ModelProvider.GOOGLE: GeminiBot,
ModelProvider.LITELLM: LiteLLMBot,
}
36 changes: 33 additions & 3 deletions openlrc/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ class ModelProvider(Enum):
ANTHROPIC = "anthropic"
OPENAI = "openai"
GOOGLE = "google"
LITELLM = "litellm"
THIRD_PARTY = "third_party"


Expand Down Expand Up @@ -430,6 +431,22 @@ def __init__(self, model_name: str):
latest_alias=None,
)

class DefaultLiteLLMModelInfo(ModelInfo):
"""Default configuration for LiteLLM-routed models."""

def __init__(self, model_name: str):
super().__init__(
name=model_name,
provider=ModelProvider.LITELLM,
input_price=1.0,
output_price=1.0,
max_tokens=8192,
context_window=128000,
vision_support=False,
knowledge_cutoff=None,
latest_alias=None,
)

class DefaultThirdPartyModelInfo(ModelInfo):
"""Default configuration for unrecognized third-party models."""

Expand Down Expand Up @@ -474,11 +491,24 @@ def get_model(cls, model_name: str, beta: bool = False) -> ModelInfo:

# If no exact match found, try to infer provider from model name
# Note the model_provider is not vital for next processing
if any(name in model_name.lower() for name in ["gpt", "openai", "davinci", "text-", "curie"]):
# LiteLLM uses provider/model format -- infer from prefix
lower_name = model_name.lower()
if lower_name.startswith("openai/"):
default_model = cls.DefaultOpenAIModelInfo(model_name)
elif lower_name.startswith("anthropic/"):
default_model = cls.DefaultAnthropicModelInfo(model_name)
elif lower_name.startswith(("gemini/", "google/")):
default_model = cls.DefaultGeminiModelInfo(model_name)
elif "/" in model_name and lower_name.split("/")[0] in (
"groq", "together_ai", "deepseek", "mistral", "bedrock",
"vertex_ai", "azure", "cohere", "fireworks", "replicate",
):
default_model = cls.DefaultThirdPartyModelInfo(model_name)
elif any(name in lower_name for name in ["gpt", "openai", "davinci", "text-", "curie"]):
default_model = cls.DefaultOpenAIModelInfo(model_name)
elif any(name in model_name.lower() for name in ["claude", "anthropic"]):
elif any(name in lower_name for name in ["claude", "anthropic"]):
default_model = cls.DefaultAnthropicModelInfo(model_name)
elif any(name in model_name.lower() for name in ["gemini", "google", "palm"]):
elif any(name in lower_name for name in ["gemini", "google", "palm"]):
default_model = cls.DefaultGeminiModelInfo(model_name)
else:
default_model = cls.DefaultThirdPartyModelInfo(model_name)
Expand Down
3 changes: 3 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,9 @@ full = [
"torchaudio>=2.0.0",
"deepfilternet>=0.5.6,<0.6",
]
litellm = [
"litellm>=1.60,<1.85",
]

[project.urls]
Homepage = "https://github.com/zh-plus/openlrc"
Expand Down
115 changes: 115 additions & 0 deletions tests/test_litellm_bot.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,115 @@
# Copyright (C) 2026. Hao Zheng
# All rights reserved.

import unittest
from unittest.mock import MagicMock, patch

from openlrc.chatbot import LiteLLMBot, route_chatbot
from openlrc.exceptions import ChatBotException


def _make_response(content="translated text", finish_reason="stop", prompt_tokens=10, completion_tokens=20):
resp = MagicMock()
resp.choices = [MagicMock()]
resp.choices[0].message.content = content
resp.choices[0].finish_reason = finish_reason
resp.usage = MagicMock()
resp.usage.prompt_tokens = prompt_tokens
resp.usage.completion_tokens = completion_tokens
resp.usage.total_tokens = prompt_tokens + completion_tokens
return resp


def _bot(model="anthropic/claude-sonnet-4-6", **kw):
bot = LiteLLMBot(model_name=model, **kw)
bot.api_fees = [0]
return bot


class TestLiteLLMBot(unittest.TestCase):
@patch("litellm.completion", return_value=_make_response())
def test_create_chat_dispatches_to_litellm(self, mock_completion):
bot = _bot()
messages = [{"role": "system", "content": "Translate to French."}, {"role": "user", "content": "Hello world"}]
response = bot._create_chat(messages)
mock_completion.assert_called_once()
kw = mock_completion.call_args.kwargs
self.assertEqual(kw["model"], "anthropic/claude-sonnet-4-6")
self.assertTrue(kw["drop_params"])
self.assertEqual(bot.get_content(response), "translated text")

@patch("litellm.completion", return_value=_make_response())
def test_params_forwarded(self, mock_completion):
bot = _bot(model="openai/gpt-4o", temperature=0.3)
bot._create_chat([{"role": "user", "content": "test"}], temperature=0.5)
kw = mock_completion.call_args.kwargs
self.assertEqual(kw["temperature"], 0.5)

@patch("litellm.completion", return_value=_make_response())
def test_top_p_omitted_when_temperature_set(self, mock_completion):
bot = _bot(temperature=0.5)
bot._create_chat([{"role": "user", "content": "test"}])
kw = mock_completion.call_args.kwargs
self.assertNotIn("top_p", kw)

@patch("litellm.completion", return_value=_make_response())
def test_top_p_included_when_temperature_default(self, mock_completion):
bot = _bot(model="openai/gpt-4o", temperature=1.0, top_p=0.8)
bot._create_chat([{"role": "user", "content": "test"}])
kw = mock_completion.call_args.kwargs
self.assertEqual(kw["top_p"], 0.8)

@patch("litellm.completion", return_value=_make_response())
def test_fee_tracking(self, mock_completion):
bot = _bot(model="openai/gpt-4o")
bot._create_chat([{"role": "user", "content": "test"}])
self.assertGreater(bot.api_fees[-1], 0)

@patch("litellm.completion", return_value=_make_response(content=None))
def test_null_content_returns_empty(self, mock_completion):
bot = _bot(model="openai/gpt-4o")
response = bot._create_chat([{"role": "user", "content": "test"}])
self.assertEqual(bot.get_content(response), "")

@patch("litellm.completion", return_value=_make_response())
def test_stop_sequences_forwarded(self, mock_completion):
bot = _bot(model="openai/gpt-4o")
bot._create_chat([{"role": "user", "content": "test"}], stop_sequences=["END"])
kw = mock_completion.call_args.kwargs
self.assertEqual(kw["stop"], ["END"])

@patch("litellm.completion", return_value=_make_response())
def test_api_key_forwarded(self, mock_completion):
bot = _bot(model="openai/gpt-4o", api_key="sk-test-123")
bot._create_chat([{"role": "user", "content": "test"}])
kw = mock_completion.call_args.kwargs
self.assertEqual(kw["api_key"], "sk-test-123")

@patch("litellm.completion")
def test_auth_error_not_retried(self, mock_completion):
import litellm as _litellm

mock_completion.side_effect = _litellm.AuthenticationError(
message="invalid key", llm_provider="openai", model="openai/gpt-4o"
)
bot = _bot(model="openai/gpt-4o")
bot.retry = 3
with self.assertRaises(ChatBotException):
bot._create_chat([{"role": "user", "content": "test"}])
self.assertEqual(mock_completion.call_count, 1)

def test_route_chatbot_litellm(self):
cls, model = route_chatbot("litellm:openai/gpt-4o")
self.assertIs(cls, LiteLLMBot)
self.assertEqual(model, "openai/gpt-4o")

@patch("litellm.completion", return_value=_make_response())
def test_extra_body_forwarded(self, mock_completion):
bot = _bot(model="openai/gpt-4o", extra_body={"seed": 42})
bot._create_chat([{"role": "user", "content": "test"}])
kw = mock_completion.call_args.kwargs
self.assertEqual(kw["seed"], 42)


if __name__ == "__main__":
unittest.main()
Loading