From bd06a1950019a9fa6f3a7249c6c024976823052f Mon Sep 17 00:00:00 2001 From: RheagalFire Date: Sat, 13 Jun 2026 23:04:30 +0530 Subject: [PATCH 1/2] feat: add LiteLLM as AI gateway provider --- src/memu/app/settings.py | 5 ++ src/memu/llm/backends/__init__.py | 10 +++- src/memu/llm/backends/litellm.py | 9 ++++ src/memu/llm/http_client.py | 3 ++ tests/llm/test_litellm_provider.py | 80 ++++++++++++++++++++++++++++++ 5 files changed, 106 insertions(+), 1 deletion(-) create mode 100644 src/memu/llm/backends/litellm.py create mode 100644 tests/llm/test_litellm_provider.py diff --git a/src/memu/app/settings.py b/src/memu/app/settings.py index adcb4f16..e1edc5fd 100644 --- a/src/memu/app/settings.py +++ b/src/memu/app/settings.py @@ -135,6 +135,11 @@ def set_provider_defaults(self) -> "LLMConfig": self.api_key = "XAI_API_KEY" if self.chat_model == "gpt-4o-mini": self.chat_model = "grok-2-latest" + elif self.provider == "litellm": + if self.base_url == "https://api.openai.com/v1": + self.base_url = "http://localhost:4000" + if self.api_key == "OPENAI_API_KEY": + self.api_key = "LITELLM_API_KEY" return self diff --git a/src/memu/llm/backends/__init__.py b/src/memu/llm/backends/__init__.py index 5350e7b2..fe988708 100644 --- a/src/memu/llm/backends/__init__.py +++ b/src/memu/llm/backends/__init__.py @@ -1,7 +1,15 @@ from memu.llm.backends.base import LLMBackend from memu.llm.backends.doubao import DoubaoLLMBackend from memu.llm.backends.grok import GrokBackend +from memu.llm.backends.litellm import LiteLLMBackend from memu.llm.backends.openai import OpenAILLMBackend from memu.llm.backends.openrouter import OpenRouterLLMBackend -__all__ = ["DoubaoLLMBackend", "GrokBackend", "LLMBackend", "OpenAILLMBackend", "OpenRouterLLMBackend"] +__all__ = [ + "DoubaoLLMBackend", + "GrokBackend", + "LLMBackend", + "LiteLLMBackend", + "OpenAILLMBackend", + "OpenRouterLLMBackend", +] diff --git a/src/memu/llm/backends/litellm.py b/src/memu/llm/backends/litellm.py new file mode 100644 index 00000000..ed4a2a70 --- /dev/null +++ b/src/memu/llm/backends/litellm.py @@ -0,0 +1,9 @@ +from __future__ import annotations + +from memu.llm.backends.openai import OpenAILLMBackend + + +class LiteLLMBackend(OpenAILLMBackend): + """Backend for LiteLLM AI gateway proxy (OpenAI-compatible).""" + + name = "litellm" diff --git a/src/memu/llm/http_client.py b/src/memu/llm/http_client.py index ba84b05b..29a6dc8d 100644 --- a/src/memu/llm/http_client.py +++ b/src/memu/llm/http_client.py @@ -12,6 +12,7 @@ from memu.llm.backends.base import LLMBackend from memu.llm.backends.doubao import DoubaoLLMBackend from memu.llm.backends.grok import GrokBackend +from memu.llm.backends.litellm import LiteLLMBackend from memu.llm.backends.openai import OpenAILLMBackend from memu.llm.backends.openrouter import OpenRouterLLMBackend @@ -73,6 +74,7 @@ def parse_embedding_response(self, data: dict[str, Any]) -> list[list[float]]: OpenAILLMBackend.name: OpenAILLMBackend, DoubaoLLMBackend.name: DoubaoLLMBackend, GrokBackend.name: GrokBackend, + LiteLLMBackend.name: LiteLLMBackend, OpenRouterLLMBackend.name: OpenRouterLLMBackend, } @@ -291,6 +293,7 @@ def _load_embedding_backend(self, provider: str) -> _EmbeddingBackend: _OpenAIEmbeddingBackend.name: _OpenAIEmbeddingBackend, _DoubaoEmbeddingBackend.name: _DoubaoEmbeddingBackend, "grok": _OpenAIEmbeddingBackend, + "litellm": _OpenAIEmbeddingBackend, _OpenRouterEmbeddingBackend.name: _OpenRouterEmbeddingBackend, } factory = backends.get(provider) diff --git a/tests/llm/test_litellm_provider.py b/tests/llm/test_litellm_provider.py new file mode 100644 index 00000000..1116df05 --- /dev/null +++ b/tests/llm/test_litellm_provider.py @@ -0,0 +1,80 @@ +import unittest + +from memu.app.settings import LLMConfig +from memu.llm.backends.litellm import LiteLLMBackend + + +class TestLiteLLMProvider(unittest.IsolatedAsyncioTestCase): + def test_settings_defaults(self): + """Test that setting provider='litellm' sets the correct defaults.""" + config = LLMConfig(provider="litellm") + self.assertEqual(config.base_url, "http://localhost:4000") + self.assertEqual(config.api_key, "LITELLM_API_KEY") + self.assertEqual(config.chat_model, "gpt-4o-mini") + + def test_settings_preserves_custom_values(self): + """Test that custom values are not overridden by litellm defaults.""" + config = LLMConfig( + provider="litellm", + base_url="http://my-proxy:8000", + api_key="sk-my-key", + chat_model="anthropic/claude-sonnet-4-6", + ) + self.assertEqual(config.base_url, "http://my-proxy:8000") + self.assertEqual(config.api_key, "sk-my-key") + self.assertEqual(config.chat_model, "anthropic/claude-sonnet-4-6") + + def test_backend_payload_parsing(self): + """Test that LiteLLMBackend parses responses correctly (inherited from OpenAI).""" + backend = LiteLLMBackend() + + dummy_response = {"choices": [{"message": {"content": "LiteLLM response content", "role": "assistant"}}]} + + result = backend.parse_summary_response(dummy_response) + self.assertEqual(result, "LiteLLM response content") + + def test_backend_summary_payload(self): + """Test that LiteLLMBackend builds correct summary payloads.""" + backend = LiteLLMBackend() + + payload = backend.build_summary_payload( + text="Hello world", + system_prompt="Summarize this.", + chat_model="anthropic/claude-sonnet-4-6", + max_tokens=100, + ) + + self.assertEqual(payload["model"], "anthropic/claude-sonnet-4-6") + self.assertEqual(len(payload["messages"]), 2) + self.assertEqual(payload["messages"][0]["role"], "system") + self.assertEqual(payload["messages"][1]["role"], "user") + self.assertEqual(payload["max_tokens"], 100) + + def test_backend_vision_payload(self): + """Test that LiteLLMBackend builds correct vision payloads.""" + backend = LiteLLMBackend() + + payload = backend.build_vision_payload( + prompt="Describe this image", + base64_image="abc123", + mime_type="image/png", + system_prompt=None, + chat_model="openai/gpt-4o", + max_tokens=200, + ) + + self.assertEqual(payload["model"], "openai/gpt-4o") + self.assertEqual(payload["messages"][0]["role"], "user") + content = payload["messages"][0]["content"] + self.assertEqual(content[0]["type"], "text") + self.assertEqual(content[1]["type"], "image_url") + + def test_backend_name(self): + """Test that the backend name is 'litellm'.""" + backend = LiteLLMBackend() + self.assertEqual(backend.name, "litellm") + + def test_backend_endpoint(self): + """Test that the backend uses OpenAI-compatible endpoint.""" + backend = LiteLLMBackend() + self.assertEqual(backend.summary_endpoint, "/chat/completions") From d075bd7babba44e503953bad3ad91f40a3152c08 Mon Sep 17 00:00:00 2001 From: RheagalFire Date: Mon, 15 Jun 2026 23:50:13 +0530 Subject: [PATCH 2/2] feat: add LiteLLM SDK client backend with litellm optional dependency --- pyproject.toml | 3 +- src/memu/app/service.py | 10 ++ src/memu/app/settings.py | 2 + src/memu/llm/litellm_sdk.py | 154 ++++++++++++++++++++++++++ tests/llm/test_litellm_provider.py | 167 ++++++++++++++++++++++------- 5 files changed, 296 insertions(+), 40 deletions(-) create mode 100644 src/memu/llm/litellm_sdk.py diff --git a/pyproject.toml b/pyproject.toml index c7858ec9..38ea4872 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -70,6 +70,7 @@ postgres = ["pgvector>=0.3.4", "sqlalchemy[postgresql-psycopgbinary]>=2.0.36"] langgraph = ["langgraph>=0.0.10", "langchain-core>=0.1.0"] claude = ["claude-agent-sdk>=0.1.24"] lazyllm = ["lazyllm>=0.7.3"] +litellm = ["litellm>=1.80.0,<1.87.0"] [project.urls] "Homepage" = "https://github.com/NevaMind-AI/MemU" @@ -78,7 +79,7 @@ lazyllm = ["lazyllm>=0.7.3"] [tool.deptry.per_rule_ignores] # Optional dependencies used in examples/ -DEP002 = ["claude-agent-sdk"] +DEP002 = ["claude-agent-sdk", "litellm"] [tool.mypy] files = ["src", "tests"] diff --git a/src/memu/app/service.py b/src/memu/app/service.py index 4e2dea04..cedf9838 100644 --- a/src/memu/app/service.py +++ b/src/memu/app/service.py @@ -117,6 +117,16 @@ def _init_llm_client(self, config: LLMConfig | None = None) -> Any: endpoint_overrides=cfg.endpoint_overrides, embed_model=cfg.embed_model, ) + elif backend == "litellm": + from memu.llm.litellm_sdk import LiteLLMSDKClient + + return LiteLLMSDKClient( + chat_model=cfg.chat_model, + embed_model=cfg.embed_model, + api_key=cfg.api_key, + api_base=cfg.base_url if cfg.base_url != "http://localhost:4000" else None, + embed_batch_size=cfg.embed_batch_size, + ) elif backend == "lazyllm_backend": from memu.llm.lazyllm_client import LazyLLMClient diff --git a/src/memu/app/settings.py b/src/memu/app/settings.py index e1edc5fd..0b2c4410 100644 --- a/src/memu/app/settings.py +++ b/src/memu/app/settings.py @@ -136,6 +136,8 @@ def set_provider_defaults(self) -> "LLMConfig": if self.chat_model == "gpt-4o-mini": self.chat_model = "grok-2-latest" elif self.provider == "litellm": + if self.client_backend == "sdk": + self.client_backend = "litellm" if self.base_url == "https://api.openai.com/v1": self.base_url = "http://localhost:4000" if self.api_key == "OPENAI_API_KEY": diff --git a/src/memu/llm/litellm_sdk.py b/src/memu/llm/litellm_sdk.py new file mode 100644 index 00000000..b2cb0114 --- /dev/null +++ b/src/memu/llm/litellm_sdk.py @@ -0,0 +1,154 @@ +from __future__ import annotations + +import base64 +import logging +from pathlib import Path +from typing import Any, cast + +logger = logging.getLogger(__name__) + + +class LiteLLMSDKClient: + """LLM client using the LiteLLM Python SDK for 100+ provider support.""" + + def __init__( + self, + *, + chat_model: str, + embed_model: str, + api_key: str | None = None, + api_base: str | None = None, + embed_batch_size: int = 1, + ): + self.chat_model = chat_model + self.embed_model = embed_model + self.api_key = api_key or None + self.api_base = api_base or None + self.embed_batch_size = embed_batch_size + + async def chat( + self, + prompt: str, + *, + max_tokens: int | None = None, + system_prompt: str | None = None, + temperature: float = 0.2, + ) -> tuple[str, dict[str, Any]]: + import litellm + + messages: list[dict[str, str]] = [] + if system_prompt is not None: + messages.append({"role": "system", "content": system_prompt}) + messages.append({"role": "user", "content": prompt}) + + kwargs: dict[str, Any] = { + "model": self.chat_model, + "messages": messages, + "temperature": temperature, + "drop_params": True, + } + if max_tokens is not None: + kwargs["max_tokens"] = max_tokens + if self.api_key: + kwargs["api_key"] = self.api_key + if self.api_base: + kwargs["api_base"] = self.api_base + + response = await litellm.acompletion(**kwargs) + data = response.model_dump() + content = data["choices"][0]["message"]["content"] or "" + logger.debug("LiteLLM chat response: %s", data) + return content, data + + async def summarize( + self, + text: str, + *, + max_tokens: int | None = None, + system_prompt: str | None = None, + ) -> tuple[str, dict[str, Any]]: + prompt = system_prompt or "Summarize the text in one short paragraph." + return await self.chat( + text, + max_tokens=max_tokens, + system_prompt=prompt, + temperature=0.2, + ) + + async def vision( + self, + prompt: str, + image_path: str, + *, + max_tokens: int | None = None, + system_prompt: str | None = None, + ) -> tuple[str, dict[str, Any]]: + import litellm + + image_data = Path(image_path).read_bytes() + base64_image = base64.b64encode(image_data).decode("utf-8") + + suffix = Path(image_path).suffix.lower() + mime_type = { + ".jpg": "image/jpeg", + ".jpeg": "image/jpeg", + ".png": "image/png", + ".gif": "image/gif", + ".webp": "image/webp", + }.get(suffix, "image/jpeg") + + messages: list[dict[str, Any]] = [] + if system_prompt: + messages.append({"role": "system", "content": system_prompt}) + + messages.append({ + "role": "user", + "content": [ + {"type": "text", "text": prompt}, + {"type": "image_url", "image_url": {"url": f"data:{mime_type};base64,{base64_image}"}}, + ], + }) + + kwargs: dict[str, Any] = { + "model": self.chat_model, + "messages": messages, + "temperature": 0.2, + "drop_params": True, + } + if max_tokens is not None: + kwargs["max_tokens"] = max_tokens + if self.api_key: + kwargs["api_key"] = self.api_key + if self.api_base: + kwargs["api_base"] = self.api_base + + response = await litellm.acompletion(**kwargs) + data = response.model_dump() + content = data["choices"][0]["message"]["content"] or "" + logger.debug("LiteLLM vision response: %s", data) + return content, data + + async def embed(self, inputs: list[str]) -> tuple[list[list[float]], dict[str, Any] | None]: + import litellm + + kwargs: dict[str, Any] = {"model": self.embed_model, "drop_params": True} + if self.api_key: + kwargs["api_key"] = self.api_key + if self.api_base: + kwargs["api_base"] = self.api_base + + if len(inputs) <= self.embed_batch_size: + response = await litellm.aembedding(input=inputs, **kwargs) + data = response.model_dump() + return [cast(list[float], d["embedding"]) for d in data["data"]], data + + all_embeddings: list[list[float]] = [] + last_data: dict[str, Any] | None = None + for idx in range(0, len(inputs), self.embed_batch_size): + batch = inputs[idx : idx + self.embed_batch_size] + response = await litellm.aembedding(input=batch, **kwargs) + data = response.model_dump() + all_embeddings.extend([cast(list[float], d["embedding"]) for d in data["data"]]) + last_data = data + + return all_embeddings, last_data diff --git a/tests/llm/test_litellm_provider.py b/tests/llm/test_litellm_provider.py index 1116df05..0c48ed54 100644 --- a/tests/llm/test_litellm_provider.py +++ b/tests/llm/test_litellm_provider.py @@ -1,59 +1,50 @@ +import sys +import types import unittest +from unittest import mock from memu.app.settings import LLMConfig from memu.llm.backends.litellm import LiteLLMBackend -class TestLiteLLMProvider(unittest.IsolatedAsyncioTestCase): - def test_settings_defaults(self): - """Test that setting provider='litellm' sets the correct defaults.""" - config = LLMConfig(provider="litellm") - self.assertEqual(config.base_url, "http://localhost:4000") - self.assertEqual(config.api_key, "LITELLM_API_KEY") - self.assertEqual(config.chat_model, "gpt-4o-mini") +def _install_litellm_stub(): + """Install a fake litellm module so tests run without the real package.""" + fake = types.ModuleType("litellm") + fake.acompletion = mock.AsyncMock(name="litellm.acompletion") + fake.aembedding = mock.AsyncMock(name="litellm.aembedding") + sys.modules["litellm"] = fake + return fake - def test_settings_preserves_custom_values(self): - """Test that custom values are not overridden by litellm defaults.""" - config = LLMConfig( - provider="litellm", - base_url="http://my-proxy:8000", - api_key="sk-my-key", - chat_model="anthropic/claude-sonnet-4-6", - ) - self.assertEqual(config.base_url, "http://my-proxy:8000") - self.assertEqual(config.api_key, "sk-my-key") - self.assertEqual(config.chat_model, "anthropic/claude-sonnet-4-6") - def test_backend_payload_parsing(self): - """Test that LiteLLMBackend parses responses correctly (inherited from OpenAI).""" +class TestLiteLLMBackend(unittest.TestCase): + def test_backend_name(self): backend = LiteLLMBackend() + self.assertEqual(backend.name, "litellm") - dummy_response = {"choices": [{"message": {"content": "LiteLLM response content", "role": "assistant"}}]} + def test_backend_endpoint(self): + backend = LiteLLMBackend() + self.assertEqual(backend.summary_endpoint, "/chat/completions") + def test_backend_payload_parsing(self): + backend = LiteLLMBackend() + dummy_response = {"choices": [{"message": {"content": "LiteLLM response", "role": "assistant"}}]} result = backend.parse_summary_response(dummy_response) - self.assertEqual(result, "LiteLLM response content") + self.assertEqual(result, "LiteLLM response") def test_backend_summary_payload(self): - """Test that LiteLLMBackend builds correct summary payloads.""" backend = LiteLLMBackend() - payload = backend.build_summary_payload( text="Hello world", system_prompt="Summarize this.", chat_model="anthropic/claude-sonnet-4-6", max_tokens=100, ) - self.assertEqual(payload["model"], "anthropic/claude-sonnet-4-6") self.assertEqual(len(payload["messages"]), 2) - self.assertEqual(payload["messages"][0]["role"], "system") - self.assertEqual(payload["messages"][1]["role"], "user") self.assertEqual(payload["max_tokens"], 100) def test_backend_vision_payload(self): - """Test that LiteLLMBackend builds correct vision payloads.""" backend = LiteLLMBackend() - payload = backend.build_vision_payload( prompt="Describe this image", base64_image="abc123", @@ -62,19 +53,117 @@ def test_backend_vision_payload(self): chat_model="openai/gpt-4o", max_tokens=200, ) - self.assertEqual(payload["model"], "openai/gpt-4o") - self.assertEqual(payload["messages"][0]["role"], "user") content = payload["messages"][0]["content"] self.assertEqual(content[0]["type"], "text") self.assertEqual(content[1]["type"], "image_url") - def test_backend_name(self): - """Test that the backend name is 'litellm'.""" - backend = LiteLLMBackend() - self.assertEqual(backend.name, "litellm") - def test_backend_endpoint(self): - """Test that the backend uses OpenAI-compatible endpoint.""" - backend = LiteLLMBackend() - self.assertEqual(backend.summary_endpoint, "/chat/completions") +class TestLiteLLMSettings(unittest.TestCase): + def test_defaults(self): + config = LLMConfig(provider="litellm") + self.assertEqual(config.base_url, "http://localhost:4000") + self.assertEqual(config.api_key, "LITELLM_API_KEY") + self.assertEqual(config.client_backend, "litellm") + + def test_preserves_custom_values(self): + config = LLMConfig( + provider="litellm", + base_url="http://my-proxy:8000", + api_key="sk-my-key", + chat_model="anthropic/claude-sonnet-4-6", + ) + self.assertEqual(config.base_url, "http://my-proxy:8000") + self.assertEqual(config.api_key, "sk-my-key") + self.assertEqual(config.chat_model, "anthropic/claude-sonnet-4-6") + + def test_httpx_backend_preserved(self): + config = LLMConfig(provider="litellm", client_backend="httpx") + self.assertEqual(config.client_backend, "httpx") + + +class TestLiteLLMSDKClient(unittest.IsolatedAsyncioTestCase): + def setUp(self): + self.fake_litellm = _install_litellm_stub() + + def tearDown(self): + sys.modules.pop("litellm", None) + + async def test_chat_calls_acompletion(self): + from types import SimpleNamespace + + mock_response = SimpleNamespace( + model_dump=lambda: { + "choices": [{"message": {"content": "4", "role": "assistant"}}], + "usage": {"prompt_tokens": 10, "completion_tokens": 1, "total_tokens": 11}, + } + ) + self.fake_litellm.acompletion.return_value = mock_response + + from memu.llm.litellm_sdk import LiteLLMSDKClient + + client = LiteLLMSDKClient( + chat_model="anthropic/claude-sonnet-4-6", + embed_model="text-embedding-3-small", + api_key="sk-test", + ) + + text, _data = await client.chat("What is 2+2?", max_tokens=10) + + self.assertEqual(text, "4") + self.fake_litellm.acompletion.assert_called_once() + call_kwargs = self.fake_litellm.acompletion.call_args[1] + self.assertEqual(call_kwargs["model"], "anthropic/claude-sonnet-4-6") + self.assertTrue(call_kwargs["drop_params"]) + self.assertEqual(call_kwargs["api_key"], "sk-test") + self.assertEqual(call_kwargs["max_tokens"], 10) + + async def test_chat_omits_api_key_when_none(self): + from types import SimpleNamespace + + mock_response = SimpleNamespace( + model_dump=lambda: { + "choices": [{"message": {"content": "ok", "role": "assistant"}}], + } + ) + self.fake_litellm.acompletion.return_value = mock_response + + from memu.llm.litellm_sdk import LiteLLMSDKClient + + client = LiteLLMSDKClient( + chat_model="openai/gpt-4o-mini", + embed_model="text-embedding-3-small", + ) + + await client.chat("test") + + call_kwargs = self.fake_litellm.acompletion.call_args[1] + self.assertNotIn("api_key", call_kwargs) + self.assertNotIn("api_base", call_kwargs) + + async def test_embed_calls_aembedding(self): + from types import SimpleNamespace + + mock_response = SimpleNamespace( + model_dump=lambda: { + "data": [{"embedding": [0.1, 0.2, 0.3]}], + "usage": {"total_tokens": 5}, + } + ) + self.fake_litellm.aembedding.return_value = mock_response + + from memu.llm.litellm_sdk import LiteLLMSDKClient + + client = LiteLLMSDKClient( + chat_model="openai/gpt-4o-mini", + embed_model="text-embedding-3-small", + ) + + embeddings, _data = await client.embed(["hello"]) + + self.assertEqual(len(embeddings), 1) + self.assertEqual(embeddings[0], [0.1, 0.2, 0.3]) + self.fake_litellm.aembedding.assert_called_once() + call_kwargs = self.fake_litellm.aembedding.call_args[1] + self.assertEqual(call_kwargs["model"], "text-embedding-3-small") + self.assertTrue(call_kwargs["drop_params"])