fix(reranker): support nested llm config in LLMReranker for non-OpenAI providers (#4405)

This commit is contained in:
Kartik
2026-03-19 15:26:00 +05:30
committed by GitHub
parent 0c4d0290cb
commit 348f44b632
6 changed files with 385 additions and 15 deletions
+7 -1
View File
@@ -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.",
)
+25 -14
View File
@@ -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()
+11
View File
@@ -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
@@ -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
@@ -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"
+125
View File
@@ -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]