From 348f44b6327ed77dcb8386b6f36d23e0fdd23480 Mon Sep 17 00:00:00 2001 From: Kartik Date: Thu, 19 Mar 2026 15:26:00 +0530 Subject: [PATCH] fix(reranker): support nested llm config in LLMReranker for non-OpenAI providers (#4405) --- mem0/configs/rerankers/llm.py | 8 +- mem0/reranker/llm_reranker.py | 39 +++-- tests/rerankers/conftest.py | 11 ++ tests/rerankers/test_llm_reranker_config.py | 63 +++++++ .../test_llm_reranker_nested_config.py | 154 ++++++++++++++++++ tests/rerankers/test_llm_reranker_rerank.py | 125 ++++++++++++++ 6 files changed, 385 insertions(+), 15 deletions(-) create mode 100644 tests/rerankers/conftest.py create mode 100644 tests/rerankers/test_llm_reranker_config.py create mode 100644 tests/rerankers/test_llm_reranker_nested_config.py create mode 100644 tests/rerankers/test_llm_reranker_rerank.py diff --git a/mem0/configs/rerankers/llm.py b/mem0/configs/rerankers/llm.py index e1475645c..de64818b5 100644 --- a/mem0/configs/rerankers/llm.py +++ b/mem0/configs/rerankers/llm.py @@ -1,4 +1,5 @@ -from typing import Optional +from typing import Any, Dict, Optional + from pydantic import Field from mem0.configs.rerankers.base import BaseRerankerConfig @@ -46,3 +47,8 @@ class LLMRerankerConfig(BaseRerankerConfig): default=None, description="Custom prompt template for scoring documents" ) + llm: Optional[Dict[str, Any]] = Field( + default=None, + description="Nested LLM configuration with 'provider' and 'config' keys. " + "Overrides top-level provider/model/api_key when provided.", + ) diff --git a/mem0/reranker/llm_reranker.py b/mem0/reranker/llm_reranker.py index d53f3c5fa..a474ea2d8 100644 --- a/mem0/reranker/llm_reranker.py +++ b/mem0/reranker/llm_reranker.py @@ -1,10 +1,10 @@ import re -from typing import List, Dict, Any, Union +from typing import Any, Dict, List, Union -from mem0.reranker.base import BaseReranker -from mem0.utils.factory import LlmFactory from mem0.configs.rerankers.base import BaseRerankerConfig from mem0.configs.rerankers.llm import LLMRerankerConfig +from mem0.reranker.base import BaseReranker +from mem0.utils.factory import LlmFactory class LLMReranker(BaseReranker): @@ -33,19 +33,30 @@ class LLMReranker(BaseReranker): self.config = config - # Create LLM configuration for the factory - llm_config = { - "model": self.config.model, - "temperature": self.config.temperature, - "max_tokens": self.config.max_tokens, - } - - # Add API key if provided - if self.config.api_key: - llm_config["api_key"] = self.config.api_key + # If a nested ``llm`` dict is provided (e.g. for non-OpenAI providers + # like Ollama that need provider-specific fields such as + # ``ollama_base_url``), use it to configure the LLM factory. + if self.config.llm: + nested = self.config.llm + llm_provider = nested.get("provider", self.config.provider) + llm_config: dict = dict(nested.get("config") or {}) + llm_config.setdefault("model", self.config.model) + llm_config.setdefault("temperature", self.config.temperature) + llm_config.setdefault("max_tokens", self.config.max_tokens) + if self.config.api_key: + llm_config.setdefault("api_key", self.config.api_key) + else: + llm_provider = self.config.provider + llm_config = { + "model": self.config.model, + "temperature": self.config.temperature, + "max_tokens": self.config.max_tokens, + } + if self.config.api_key: + llm_config["api_key"] = self.config.api_key # Initialize LLM using the factory - self.llm = LlmFactory.create(self.config.provider, llm_config) + self.llm = LlmFactory.create(llm_provider, llm_config) # Default scoring prompt self.scoring_prompt = getattr(self.config, 'scoring_prompt', None) or self._get_default_prompt() diff --git a/tests/rerankers/conftest.py b/tests/rerankers/conftest.py new file mode 100644 index 000000000..258d959cb --- /dev/null +++ b/tests/rerankers/conftest.py @@ -0,0 +1,11 @@ +from unittest.mock import MagicMock, patch + +import pytest + + +@pytest.fixture +def mock_llm(): + with patch("mem0.reranker.llm_reranker.LlmFactory") as mock_factory: + mock_llm_instance = MagicMock() + mock_factory.create.return_value = mock_llm_instance + yield mock_factory, mock_llm_instance diff --git a/tests/rerankers/test_llm_reranker_config.py b/tests/rerankers/test_llm_reranker_config.py new file mode 100644 index 000000000..d3577a49b --- /dev/null +++ b/tests/rerankers/test_llm_reranker_config.py @@ -0,0 +1,63 @@ +from mem0.configs.rerankers.base import BaseRerankerConfig +from mem0.configs.rerankers.llm import LLMRerankerConfig +from mem0.reranker.llm_reranker import LLMReranker + + +class TestLLMRerankerConfig: + def test_default_config(self): + config = LLMRerankerConfig() + assert config.model == "gpt-4o-mini" + assert config.provider == "openai" + assert config.temperature == 0.0 + assert config.max_tokens == 100 + assert config.llm is None + assert config.scoring_prompt is None + assert config.top_k is None + + def test_nested_llm_field_accepted(self): + config = LLMRerankerConfig( + llm={"provider": "ollama", "config": {"ollama_base_url": "http://localhost:11434"}} + ) + assert config.llm["provider"] == "ollama" + assert config.llm["config"]["ollama_base_url"] == "http://localhost:11434" + + +class TestLLMRerankerInit: + def test_init_with_dict_config(self, mock_llm): + mock_factory, _ = mock_llm + reranker = LLMReranker({"provider": "openai", "model": "gpt-4o", "api_key": "sk-test"}) + + assert reranker.config.provider == "openai" + assert reranker.config.model == "gpt-4o" + mock_factory.create.assert_called_once_with( + "openai", + {"model": "gpt-4o", "temperature": 0.0, "max_tokens": 100, "api_key": "sk-test"}, + ) + + def test_init_with_llm_reranker_config(self, mock_llm): + mock_factory, _ = mock_llm + config = LLMRerankerConfig(provider="anthropic", model="claude-3-haiku", api_key="sk-ant") + reranker = LLMReranker(config) + + assert reranker.config.provider == "anthropic" + mock_factory.create.assert_called_once_with( + "anthropic", + {"model": "claude-3-haiku", "temperature": 0.0, "max_tokens": 100, "api_key": "sk-ant"}, + ) + + def test_init_converts_base_reranker_config(self, mock_llm): + mock_factory, _ = mock_llm + base_config = BaseRerankerConfig(provider="openai", model="gpt-4o-mini") + reranker = LLMReranker(base_config) + + assert isinstance(reranker.config, LLMRerankerConfig) + assert reranker.config.temperature == 0.0 + assert reranker.config.max_tokens == 100 + + def test_init_without_api_key(self, mock_llm): + mock_factory, _ = mock_llm + LLMReranker({"provider": "openai", "model": "gpt-4o-mini"}) + + call_args = mock_factory.create.call_args + llm_config = call_args[0][1] + assert "api_key" not in llm_config diff --git a/tests/rerankers/test_llm_reranker_nested_config.py b/tests/rerankers/test_llm_reranker_nested_config.py new file mode 100644 index 000000000..7017780ef --- /dev/null +++ b/tests/rerankers/test_llm_reranker_nested_config.py @@ -0,0 +1,154 @@ +from mem0.reranker.llm_reranker import LLMReranker + + +class TestNestedLLMConfig: + def test_nested_llm_overrides_provider(self, mock_llm): + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "openai", + "model": "gpt-4o-mini", + "llm": { + "provider": "ollama", + "config": {"model": "llama3", "ollama_base_url": "http://localhost:11434"}, + }, + }) + + call_args = mock_factory.create.call_args + assert call_args[0][0] == "ollama" + + def test_nested_llm_passes_provider_specific_config(self, mock_llm): + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "openai", + "llm": { + "provider": "ollama", + "config": { + "model": "llama3", + "ollama_base_url": "http://localhost:11434", + }, + }, + }) + + call_args = mock_factory.create.call_args + llm_config = call_args[0][1] + assert llm_config["ollama_base_url"] == "http://localhost:11434" + assert llm_config["model"] == "llama3" + + def test_nested_llm_inherits_top_level_defaults(self, mock_llm): + """Nested config should inherit temperature/max_tokens from top-level if not overridden.""" + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "openai", + "temperature": 0.0, + "max_tokens": 100, + "llm": { + "provider": "ollama", + "config": {"model": "llama3"}, + }, + }) + + call_args = mock_factory.create.call_args + llm_config = call_args[0][1] + assert llm_config["temperature"] == 0.0 + assert llm_config["max_tokens"] == 100 + + def test_nested_llm_config_values_take_precedence(self, mock_llm): + """Values explicitly set in nested config should not be overridden by top-level defaults.""" + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "openai", + "model": "gpt-4o-mini", + "temperature": 0.0, + "max_tokens": 100, + "llm": { + "provider": "ollama", + "config": { + "model": "custom-model", + "temperature": 0.5, + "max_tokens": 200, + }, + }, + }) + + call_args = mock_factory.create.call_args + llm_config = call_args[0][1] + assert llm_config["model"] == "custom-model" + assert llm_config["temperature"] == 0.5 + assert llm_config["max_tokens"] == 200 + + def test_nested_llm_falls_back_to_top_level_provider(self, mock_llm): + """If nested llm dict has no 'provider', use top-level provider.""" + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "anthropic", + "model": "claude-3-haiku", + "llm": { + "config": {"model": "claude-3-sonnet"}, + }, + }) + + call_args = mock_factory.create.call_args + assert call_args[0][0] == "anthropic" + assert call_args[0][1]["model"] == "claude-3-sonnet" + + def test_nested_llm_with_empty_config(self, mock_llm): + """Nested llm with no config dict should still work, using top-level defaults.""" + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "openai", + "model": "gpt-4o-mini", + "llm": {"provider": "ollama"}, + }) + + call_args = mock_factory.create.call_args + assert call_args[0][0] == "ollama" + llm_config = call_args[0][1] + assert llm_config["model"] == "gpt-4o-mini" + assert llm_config["temperature"] == 0.0 + assert llm_config["max_tokens"] == 100 + + def test_nested_llm_with_none_config(self, mock_llm): + """Nested llm with config: None should still work, using top-level defaults.""" + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "openai", + "model": "gpt-4o-mini", + "llm": {"provider": "ollama", "config": None}, + }) + + call_args = mock_factory.create.call_args + assert call_args[0][0] == "ollama" + llm_config = call_args[0][1] + assert llm_config["model"] == "gpt-4o-mini" + + def test_nested_llm_inherits_top_level_api_key(self, mock_llm): + """Top-level api_key should be inherited by nested config if not already set.""" + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "openai", + "api_key": "sk-top-level", + "llm": { + "provider": "openai", + "config": {"model": "gpt-4o"}, + }, + }) + + call_args = mock_factory.create.call_args + llm_config = call_args[0][1] + assert llm_config["api_key"] == "sk-top-level" + + def test_nested_llm_config_api_key_not_overridden(self, mock_llm): + """If nested config already has api_key, top-level api_key should not override it.""" + mock_factory, _ = mock_llm + LLMReranker({ + "provider": "openai", + "api_key": "sk-top-level", + "llm": { + "provider": "openai", + "config": {"model": "gpt-4o", "api_key": "sk-nested"}, + }, + }) + + call_args = mock_factory.create.call_args + llm_config = call_args[0][1] + assert llm_config["api_key"] == "sk-nested" diff --git a/tests/rerankers/test_llm_reranker_rerank.py b/tests/rerankers/test_llm_reranker_rerank.py new file mode 100644 index 000000000..9e16f149a --- /dev/null +++ b/tests/rerankers/test_llm_reranker_rerank.py @@ -0,0 +1,125 @@ +import pytest + +from mem0.reranker.llm_reranker import LLMReranker + + +class TestExtractScore: + @pytest.fixture + def reranker(self, mock_llm): + return LLMReranker({"provider": "openai"}) + + @pytest.mark.parametrize( + "text,expected", + [ + ("0.85", 0.85), + ("0.0", 0.0), + ("1.0", 1.0), + ("The score is 0.72.", 0.72), + ("Score: 0.9 out of 1.0", 0.9), + ], + ) + def test_valid_scores(self, reranker, text, expected): + assert reranker._extract_score(text) == expected + + def test_no_score_returns_fallback(self, reranker): + assert reranker._extract_score("no numbers here") == 0.5 + + def test_clamps_to_1(self, reranker): + assert reranker._extract_score("1.0") == 1.0 + + +class TestRerank: + def test_empty_documents(self, mock_llm): + reranker = LLMReranker({"provider": "openai"}) + result = reranker.rerank("query", []) + assert result == [] + + def test_documents_sorted_by_score_descending(self, mock_llm): + _, mock_llm_instance = mock_llm + mock_llm_instance.generate_response.side_effect = ["0.3", "0.9", "0.6"] + + reranker = LLMReranker({"provider": "openai"}) + docs = [ + {"memory": "low relevance"}, + {"memory": "high relevance"}, + {"memory": "mid relevance"}, + ] + + result = reranker.rerank("test query", docs) + + assert len(result) == 3 + assert result[0]["rerank_score"] == 0.9 + assert result[1]["rerank_score"] == 0.6 + assert result[2]["rerank_score"] == 0.3 + + def test_top_k_limits_results(self, mock_llm): + _, mock_llm_instance = mock_llm + mock_llm_instance.generate_response.side_effect = ["0.9", "0.5", "0.1"] + + reranker = LLMReranker({"provider": "openai"}) + docs = [{"memory": f"doc{i}"} for i in range(3)] + + result = reranker.rerank("query", docs, top_k=2) + assert len(result) == 2 + + def test_config_top_k_used_when_arg_not_provided(self, mock_llm): + _, mock_llm_instance = mock_llm + mock_llm_instance.generate_response.side_effect = ["0.9", "0.5", "0.1"] + + reranker = LLMReranker({"provider": "openai", "top_k": 1}) + docs = [{"memory": f"doc{i}"} for i in range(3)] + + result = reranker.rerank("query", docs) + assert len(result) == 1 + + def test_text_field_extraction(self, mock_llm): + _, mock_llm_instance = mock_llm + mock_llm_instance.generate_response.return_value = "0.8" + + reranker = LLMReranker({"provider": "openai"}) + reranker.rerank("query", [{"text": "some text"}]) + + prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"] + assert "some text" in prompt_sent + + def test_content_field_extraction(self, mock_llm): + _, mock_llm_instance = mock_llm + mock_llm_instance.generate_response.return_value = "0.8" + + reranker = LLMReranker({"provider": "openai"}) + reranker.rerank("query", [{"content": "some content"}]) + + prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"] + assert "some content" in prompt_sent + + def test_fallback_score_on_llm_error(self, mock_llm): + _, mock_llm_instance = mock_llm + mock_llm_instance.generate_response.side_effect = RuntimeError("API error") + + reranker = LLMReranker({"provider": "openai"}) + result = reranker.rerank("query", [{"memory": "doc"}]) + + assert len(result) == 1 + assert result[0]["rerank_score"] == 0.5 + + def test_custom_scoring_prompt(self, mock_llm): + _, mock_llm_instance = mock_llm + mock_llm_instance.generate_response.return_value = "0.7" + + custom_prompt = "Rate this: query={query} doc={document}" + reranker = LLMReranker({"provider": "openai", "scoring_prompt": custom_prompt}) + reranker.rerank("my query", [{"memory": "my doc"}]) + + prompt_sent = mock_llm_instance.generate_response.call_args[1]["messages"][0]["content"] + assert prompt_sent == "Rate this: query=my query doc=my doc" + + def test_original_doc_not_mutated(self, mock_llm): + _, mock_llm_instance = mock_llm + mock_llm_instance.generate_response.return_value = "0.8" + + reranker = LLMReranker({"provider": "openai"}) + original_doc = {"memory": "test", "id": "123"} + result = reranker.rerank("query", [original_doc]) + + assert "rerank_score" not in original_doc + assert "rerank_score" in result[0]