diff --git a/pyproject.toml b/pyproject.toml index dcad6e9f..0da4c85a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -71,6 +71,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"] # Rich document ingestion (PDF, Word, PowerPoint, Excel, ...) via MarkItDown. document = ["markitdown[docx,pptx,xlsx,xls,pdf]>=0.1.0"] @@ -81,7 +82,7 @@ document = ["markitdown[docx,pptx,xlsx,xls,pdf]>=0.1.0"] [tool.deptry.per_rule_ignores] # Optional dependencies used in examples/ or imported lazily behind extras. -DEP002 = ["claude-agent-sdk", "markitdown"] +DEP002 = ["claude-agent-sdk", "litellm", "markitdown"] [tool.mypy] files = ["src", "tests"] diff --git a/src/memu/app/settings.py b/src/memu/app/settings.py index 6c15e7eb..0b45271a 100644 --- a/src/memu/app/settings.py +++ b/src/memu/app/settings.py @@ -146,6 +146,13 @@ def set_provider_defaults(self) -> "LLMConfig": self.api_key = api_key if self.chat_model == "gpt-5.4-mini": self.chat_model = chat_model + 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": + self.api_key = "LITELLM_API_KEY" return self diff --git a/src/memu/embedding/gateway.py b/src/memu/embedding/gateway.py index 528f57b7..5208d53b 100644 --- a/src/memu/embedding/gateway.py +++ b/src/memu/embedding/gateway.py @@ -47,6 +47,18 @@ def _build_lazyllm_client(cfg: EmbeddingConfig) -> Any: ) +def _build_litellm_client(cfg: EmbeddingConfig) -> Any: + from memu.llm.litellm_sdk import LiteLLMSDKClient + + return LiteLLMSDKClient( + chat_model="", + embed_model=cfg.embed_model, + api_key=cfg.api_key, + api_base=cfg.base_url if cfg.base_url != "https://api.openai.com/v1" else None, + embed_batch_size=cfg.embed_batch_size, + ) + + def _build_anthropic_client(cfg: EmbeddingConfig) -> Any: msg = ( "Anthropic does not provide an embeddings API. Configure an embedding " @@ -59,6 +71,7 @@ def _build_anthropic_client(cfg: EmbeddingConfig) -> Any: EMBEDDING_CLIENT_BUILDERS: dict[str, Callable[[EmbeddingConfig], Any]] = { "sdk": _build_sdk_client, "httpx": _build_httpx_client, + "litellm": _build_litellm_client, "lazyllm_backend": _build_lazyllm_client, "anthropic": _build_anthropic_client, } diff --git a/src/memu/llm/backends/__init__.py b/src/memu/llm/backends/__init__.py index 2a80b895..4ed7a6e3 100644 --- a/src/memu/llm/backends/__init__.py +++ b/src/memu/llm/backends/__init__.py @@ -4,6 +4,7 @@ from memu.llm.backends.doubao import DoubaoLLMBackend from memu.llm.backends.grok import GrokBackend from memu.llm.backends.kimi import KimiLLMBackend +from memu.llm.backends.litellm import LiteLLMBackend from memu.llm.backends.minimax import MiniMaxLLMBackend from memu.llm.backends.openai import OpenAILLMBackend from memu.llm.backends.openrouter import OpenRouterLLMBackend @@ -15,6 +16,7 @@ "GrokBackend", "KimiLLMBackend", "LLMBackend", + "LiteLLMBackend", "MiniMaxLLMBackend", "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/gateway.py b/src/memu/llm/gateway.py index 25fe0834..13dbfbcb 100644 --- a/src/memu/llm/gateway.py +++ b/src/memu/llm/gateway.py @@ -50,6 +50,18 @@ def _build_httpx_client(cfg: LLMConfig) -> Any: ) +def _build_litellm_client(cfg: LLMConfig) -> Any: + 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, + ) + + def _build_lazyllm_client(cfg: LLMConfig) -> Any: from memu.llm.lazyllm_client import LazyLLMClient @@ -72,6 +84,7 @@ def _build_lazyllm_client(cfg: LLMConfig) -> Any: "sdk": _build_sdk_client, "anthropic": _build_anthropic_client, "httpx": _build_httpx_client, + "litellm": _build_litellm_client, "lazyllm_backend": _build_lazyllm_client, } diff --git a/src/memu/llm/http_client.py b/src/memu/llm/http_client.py index c09bad29..fe196fd6 100644 --- a/src/memu/llm/http_client.py +++ b/src/memu/llm/http_client.py @@ -15,6 +15,7 @@ from memu.llm.backends.doubao import DoubaoLLMBackend from memu.llm.backends.grok import GrokBackend from memu.llm.backends.kimi import KimiLLMBackend +from memu.llm.backends.litellm import LiteLLMBackend from memu.llm.backends.minimax import MiniMaxLLMBackend from memu.llm.backends.openai import OpenAILLMBackend from memu.llm.backends.openrouter import OpenRouterLLMBackend @@ -30,6 +31,7 @@ def _load_proxy() -> str | None: OpenAILLMBackend.name: OpenAILLMBackend, ClaudeLLMBackend.name: ClaudeLLMBackend, GrokBackend.name: GrokBackend, + LiteLLMBackend.name: LiteLLMBackend, DeepSeekLLMBackend.name: DeepSeekLLMBackend, KimiLLMBackend.name: KimiLLMBackend, MiniMaxLLMBackend.name: MiniMaxLLMBackend, 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 new file mode 100644 index 00000000..0c48ed54 --- /dev/null +++ b/tests/llm/test_litellm_provider.py @@ -0,0 +1,169 @@ +import sys +import types +import unittest +from unittest import mock + +from memu.app.settings import LLMConfig +from memu.llm.backends.litellm import LiteLLMBackend + + +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 + + +class TestLiteLLMBackend(unittest.TestCase): + def test_backend_name(self): + backend = LiteLLMBackend() + self.assertEqual(backend.name, "litellm") + + 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") + + def test_backend_summary_payload(self): + 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["max_tokens"], 100) + + def test_backend_vision_payload(self): + 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") + content = payload["messages"][0]["content"] + self.assertEqual(content[0]["type"], "text") + self.assertEqual(content[1]["type"], "image_url") + + +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"])