Skip to content

Commit b23d5cb

Browse files
authored
Merge pull request #134 from RheagalFire/feat/add-litellm-provider
feat: add LiteLLM as unified LLM provider gateway
2 parents d915929 + d8b9df7 commit b23d5cb

4 files changed

Lines changed: 299 additions & 6 deletions

File tree

openlrc/chatbot.py

Lines changed: 148 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -67,6 +67,8 @@ def route_chatbot(model: str) -> tuple[type, str]:
6767
return GPTBot, chatbot_model
6868
elif chatbot_type == "anthropic":
6969
return ClaudeBot, chatbot_model
70+
elif chatbot_type == "litellm":
71+
return LiteLLMBot, chatbot_model
7072
else:
7173
raise ValueError(f"Invalid chatbot type {chatbot_type}.")
7274

@@ -260,7 +262,6 @@ def __init__(
260262
api_key: str | None = None,
261263
extra_body: dict | None = None,
262264
):
263-
264265
# clamp temperature to 0-2
265266
temperature = max(0, min(2, temperature))
266267

@@ -459,7 +460,6 @@ def __init__(
459460
api_key: str | None = None,
460461
extra_body: dict | None = None,
461462
):
462-
463463
# clamp temperature to 0-1
464464
temperature = max(0, min(1, temperature))
465465

@@ -774,4 +774,149 @@ def close(self):
774774
self.client.close()
775775

776776

777-
provider2chatbot = {ModelProvider.OPENAI: GPTBot, ModelProvider.ANTHROPIC: ClaudeBot, ModelProvider.GOOGLE: GeminiBot}
777+
class LiteLLMBot(ChatBot):
778+
"""ChatBot backed by the LiteLLM SDK, routing to 100+ LLM providers."""
779+
780+
def __init__(
781+
self,
782+
model_name: str = "openai/gpt-4o",
783+
temperature: float = 1,
784+
top_p: float = 1,
785+
retry: int = 8,
786+
max_async: int = 16,
787+
fee_limit: float = 0.8,
788+
proxy: str | None = None,
789+
base_url_config: dict | None = None,
790+
api_key: str | None = None,
791+
extra_body: dict | None = None,
792+
):
793+
temperature = max(0, min(2, temperature))
794+
super().__init__(model_name, temperature, top_p, retry, max_async, fee_limit)
795+
self.api_key = api_key
796+
self.api_base = (base_url_config or {}).get("litellm")
797+
self.extra_body = dict(extra_body) if extra_body else {}
798+
if proxy:
799+
logger.warning(
800+
"LiteLLMBot does not support per-instance proxy. "
801+
"Set HTTP_PROXY/HTTPS_PROXY environment variables instead."
802+
)
803+
804+
def update_fee(self, response):
805+
usage = getattr(response, "usage", None)
806+
if usage is None:
807+
return
808+
prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0
809+
completion_tokens = getattr(usage, "completion_tokens", 0) or 0
810+
self.api_fees[-1] += (
811+
prompt_tokens * self.model_info.input_price + completion_tokens * self.model_info.output_price
812+
) / 1000000
813+
814+
def get_content(self, response):
815+
content = response.choices[0].message.content
816+
return content if content is not None else ""
817+
818+
def _create_chat(
819+
self,
820+
messages: list[dict],
821+
stop_sequences: list[str] | None = None,
822+
output_checker: Callable = lambda user_input, generated_content: True,
823+
temperature: float | None = None,
824+
top_p: float | None = None,
825+
):
826+
try:
827+
import litellm
828+
except ImportError:
829+
raise ImportError(
830+
"litellm is required for the litellm: provider. Install with: pip install 'openlrc[litellm]'"
831+
)
832+
833+
effective_temperature = temperature if temperature is not None else self.temperature
834+
effective_top_p = top_p if top_p is not None else self.top_p
835+
836+
completion_kwargs: dict = {
837+
"model": self.model_name,
838+
"messages": messages,
839+
"temperature": effective_temperature,
840+
"max_tokens": self._compute_max_tokens(messages),
841+
"drop_params": True,
842+
}
843+
completion_kwargs["top_p"] = effective_top_p
844+
if stop_sequences:
845+
completion_kwargs["stop"] = stop_sequences
846+
if self.api_key:
847+
completion_kwargs["api_key"] = self.api_key
848+
if self.api_base:
849+
completion_kwargs["api_base"] = self.api_base
850+
completion_kwargs.update(self.extra_body)
851+
852+
response = None
853+
validated = False
854+
for i in range(self.retry):
855+
try:
856+
response = litellm.completion(**completion_kwargs)
857+
self.update_fee(response)
858+
859+
if response.choices[0].finish_reason == "length":
860+
usage = getattr(response, "usage", None)
861+
raise LengthExceedException(
862+
prompt_tokens=getattr(usage, "prompt_tokens", -1) if usage else -1,
863+
completion_tokens=getattr(usage, "completion_tokens", -1) if usage else -1,
864+
total_tokens=getattr(usage, "total_tokens", -1) if usage else -1,
865+
)
866+
867+
response_text = remove_stop(self.get_content(response), stop_sequences)
868+
869+
if not output_checker(messages[-1]["content"], response_text):
870+
logger.warning(f"Invalid response format. Retry num: {i + 1}.")
871+
continue
872+
873+
validated = True
874+
break
875+
except litellm.AuthenticationError as e:
876+
raise ChatBotException(f"Authentication failed: {e}") from e
877+
except (
878+
litellm.BadRequestError,
879+
litellm.NotFoundError,
880+
litellm.PermissionDeniedError,
881+
litellm.UnprocessableEntityError,
882+
) as e:
883+
raise ChatBotException(f"Client error: {e}") from e
884+
except (
885+
litellm.RateLimitError,
886+
litellm.APIConnectionError,
887+
litellm.InternalServerError,
888+
litellm.ServiceUnavailableError,
889+
litellm.Timeout,
890+
json.decoder.JSONDecodeError,
891+
) as e:
892+
sleep_time = self._get_sleep_time(e)
893+
logger.warning(f"{type(e).__name__}: {e}. Wait {sleep_time}s before retry. Retry num: {i + 1}.")
894+
time.sleep(sleep_time)
895+
896+
if not response:
897+
raise ChatBotException("Failed to create a chat.")
898+
899+
if not validated:
900+
logger.warning("Response format validation failed after all retries, returning best-effort response.")
901+
902+
return response
903+
904+
@staticmethod
905+
def _get_sleep_time(error):
906+
qualname = type(error).__name__
907+
if qualname == "RateLimitError":
908+
return random.randint(30, 60)
909+
elif qualname in ("Timeout", "APIConnectionError"):
910+
return 3
911+
elif qualname == "JSONDecodeError":
912+
return 1
913+
else:
914+
return 15
915+
916+
917+
provider2chatbot = {
918+
ModelProvider.OPENAI: GPTBot,
919+
ModelProvider.ANTHROPIC: ClaudeBot,
920+
ModelProvider.GOOGLE: GeminiBot,
921+
ModelProvider.LITELLM: LiteLLMBot,
922+
}

openlrc/models.py

Lines changed: 33 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ class ModelProvider(Enum):
99
ANTHROPIC = "anthropic"
1010
OPENAI = "openai"
1111
GOOGLE = "google"
12+
LITELLM = "litellm"
1213
THIRD_PARTY = "third_party"
1314

1415

@@ -430,6 +431,22 @@ def __init__(self, model_name: str):
430431
latest_alias=None,
431432
)
432433

434+
class DefaultLiteLLMModelInfo(ModelInfo):
435+
"""Default configuration for LiteLLM-routed models."""
436+
437+
def __init__(self, model_name: str):
438+
super().__init__(
439+
name=model_name,
440+
provider=ModelProvider.LITELLM,
441+
input_price=1.0,
442+
output_price=1.0,
443+
max_tokens=8192,
444+
context_window=128000,
445+
vision_support=False,
446+
knowledge_cutoff=None,
447+
latest_alias=None,
448+
)
449+
433450
class DefaultThirdPartyModelInfo(ModelInfo):
434451
"""Default configuration for unrecognized third-party models."""
435452

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

475492
# If no exact match found, try to infer provider from model name
476493
# Note the model_provider is not vital for next processing
477-
if any(name in model_name.lower() for name in ["gpt", "openai", "davinci", "text-", "curie"]):
494+
# LiteLLM uses provider/model format -- infer from prefix
495+
lower_name = model_name.lower()
496+
if lower_name.startswith("openai/"):
497+
default_model = cls.DefaultOpenAIModelInfo(model_name)
498+
elif lower_name.startswith("anthropic/"):
499+
default_model = cls.DefaultAnthropicModelInfo(model_name)
500+
elif lower_name.startswith(("gemini/", "google/")):
501+
default_model = cls.DefaultGeminiModelInfo(model_name)
502+
elif "/" in model_name and lower_name.split("/")[0] in (
503+
"groq", "together_ai", "deepseek", "mistral", "bedrock",
504+
"vertex_ai", "azure", "cohere", "fireworks", "replicate",
505+
):
506+
default_model = cls.DefaultThirdPartyModelInfo(model_name)
507+
elif any(name in lower_name for name in ["gpt", "openai", "davinci", "text-", "curie"]):
478508
default_model = cls.DefaultOpenAIModelInfo(model_name)
479-
elif any(name in model_name.lower() for name in ["claude", "anthropic"]):
509+
elif any(name in lower_name for name in ["claude", "anthropic"]):
480510
default_model = cls.DefaultAnthropicModelInfo(model_name)
481-
elif any(name in model_name.lower() for name in ["gemini", "google", "palm"]):
511+
elif any(name in lower_name for name in ["gemini", "google", "palm"]):
482512
default_model = cls.DefaultGeminiModelInfo(model_name)
483513
else:
484514
default_model = cls.DefaultThirdPartyModelInfo(model_name)

pyproject.toml

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -52,6 +52,9 @@ full = [
5252
"torchaudio>=2.0.0",
5353
"deepfilternet>=0.5.6,<0.6",
5454
]
55+
litellm = [
56+
"litellm>=1.60,<1.85",
57+
]
5558

5659
[project.urls]
5760
Homepage = "https://github.com/zh-plus/openlrc"

tests/test_litellm_bot.py

Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,115 @@
1+
# Copyright (C) 2026. Hao Zheng
2+
# All rights reserved.
3+
4+
import unittest
5+
from unittest.mock import MagicMock, patch
6+
7+
from openlrc.chatbot import LiteLLMBot, route_chatbot
8+
from openlrc.exceptions import ChatBotException
9+
10+
11+
def _make_response(content="translated text", finish_reason="stop", prompt_tokens=10, completion_tokens=20):
12+
resp = MagicMock()
13+
resp.choices = [MagicMock()]
14+
resp.choices[0].message.content = content
15+
resp.choices[0].finish_reason = finish_reason
16+
resp.usage = MagicMock()
17+
resp.usage.prompt_tokens = prompt_tokens
18+
resp.usage.completion_tokens = completion_tokens
19+
resp.usage.total_tokens = prompt_tokens + completion_tokens
20+
return resp
21+
22+
23+
def _bot(model="anthropic/claude-sonnet-4-6", **kw):
24+
bot = LiteLLMBot(model_name=model, **kw)
25+
bot.api_fees = [0]
26+
return bot
27+
28+
29+
class TestLiteLLMBot(unittest.TestCase):
30+
@patch("litellm.completion", return_value=_make_response())
31+
def test_create_chat_dispatches_to_litellm(self, mock_completion):
32+
bot = _bot()
33+
messages = [{"role": "system", "content": "Translate to French."}, {"role": "user", "content": "Hello world"}]
34+
response = bot._create_chat(messages)
35+
mock_completion.assert_called_once()
36+
kw = mock_completion.call_args.kwargs
37+
self.assertEqual(kw["model"], "anthropic/claude-sonnet-4-6")
38+
self.assertTrue(kw["drop_params"])
39+
self.assertEqual(bot.get_content(response), "translated text")
40+
41+
@patch("litellm.completion", return_value=_make_response())
42+
def test_params_forwarded(self, mock_completion):
43+
bot = _bot(model="openai/gpt-4o", temperature=0.3)
44+
bot._create_chat([{"role": "user", "content": "test"}], temperature=0.5)
45+
kw = mock_completion.call_args.kwargs
46+
self.assertEqual(kw["temperature"], 0.5)
47+
48+
@patch("litellm.completion", return_value=_make_response())
49+
def test_top_p_omitted_when_temperature_set(self, mock_completion):
50+
bot = _bot(temperature=0.5)
51+
bot._create_chat([{"role": "user", "content": "test"}])
52+
kw = mock_completion.call_args.kwargs
53+
self.assertNotIn("top_p", kw)
54+
55+
@patch("litellm.completion", return_value=_make_response())
56+
def test_top_p_included_when_temperature_default(self, mock_completion):
57+
bot = _bot(model="openai/gpt-4o", temperature=1.0, top_p=0.8)
58+
bot._create_chat([{"role": "user", "content": "test"}])
59+
kw = mock_completion.call_args.kwargs
60+
self.assertEqual(kw["top_p"], 0.8)
61+
62+
@patch("litellm.completion", return_value=_make_response())
63+
def test_fee_tracking(self, mock_completion):
64+
bot = _bot(model="openai/gpt-4o")
65+
bot._create_chat([{"role": "user", "content": "test"}])
66+
self.assertGreater(bot.api_fees[-1], 0)
67+
68+
@patch("litellm.completion", return_value=_make_response(content=None))
69+
def test_null_content_returns_empty(self, mock_completion):
70+
bot = _bot(model="openai/gpt-4o")
71+
response = bot._create_chat([{"role": "user", "content": "test"}])
72+
self.assertEqual(bot.get_content(response), "")
73+
74+
@patch("litellm.completion", return_value=_make_response())
75+
def test_stop_sequences_forwarded(self, mock_completion):
76+
bot = _bot(model="openai/gpt-4o")
77+
bot._create_chat([{"role": "user", "content": "test"}], stop_sequences=["END"])
78+
kw = mock_completion.call_args.kwargs
79+
self.assertEqual(kw["stop"], ["END"])
80+
81+
@patch("litellm.completion", return_value=_make_response())
82+
def test_api_key_forwarded(self, mock_completion):
83+
bot = _bot(model="openai/gpt-4o", api_key="sk-test-123")
84+
bot._create_chat([{"role": "user", "content": "test"}])
85+
kw = mock_completion.call_args.kwargs
86+
self.assertEqual(kw["api_key"], "sk-test-123")
87+
88+
@patch("litellm.completion")
89+
def test_auth_error_not_retried(self, mock_completion):
90+
import litellm as _litellm
91+
92+
mock_completion.side_effect = _litellm.AuthenticationError(
93+
message="invalid key", llm_provider="openai", model="openai/gpt-4o"
94+
)
95+
bot = _bot(model="openai/gpt-4o")
96+
bot.retry = 3
97+
with self.assertRaises(ChatBotException):
98+
bot._create_chat([{"role": "user", "content": "test"}])
99+
self.assertEqual(mock_completion.call_count, 1)
100+
101+
def test_route_chatbot_litellm(self):
102+
cls, model = route_chatbot("litellm:openai/gpt-4o")
103+
self.assertIs(cls, LiteLLMBot)
104+
self.assertEqual(model, "openai/gpt-4o")
105+
106+
@patch("litellm.completion", return_value=_make_response())
107+
def test_extra_body_forwarded(self, mock_completion):
108+
bot = _bot(model="openai/gpt-4o", extra_body={"seed": 42})
109+
bot._create_chat([{"role": "user", "content": "test"}])
110+
kw = mock_completion.call_args.kwargs
111+
self.assertEqual(kw["seed"], 42)
112+
113+
114+
if __name__ == "__main__":
115+
unittest.main()

0 commit comments

Comments
 (0)