From 51541743423e0211cbd97b87fb6e1e5b203fec8a Mon Sep 17 00:00:00 2001 From: kartik-mem0 Date: Wed, 18 Mar 2026 22:12:04 +0530 Subject: [PATCH] chore: add nested llm config support to LLM reranker --- mem0/configs/rerankers/llm.py | 8 ++++++- mem0/reranker/llm_reranker.py | 39 ++++++++++++++++++++++------------- 2 files changed, 32 insertions(+), 15 deletions(-) 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()