fix(reranker): support nested llm config in LLMReranker for non-OpenAI providers (#4405)
This commit is contained in:
@@ -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.",
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"
|
||||
@@ -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]
|
||||
Reference in New Issue
Block a user