diff --git a/mcp_server/src/graphiti_mcp_server.py b/mcp_server/src/graphiti_mcp_server.py index efd9536de3..39d1d77c56 100644 --- a/mcp_server/src/graphiti_mcp_server.py +++ b/mcp_server/src/graphiti_mcp_server.py @@ -37,7 +37,12 @@ SuccessResponse, TripletResponse, ) -from services.factories import DatabaseDriverFactory, EmbedderFactory, LLMClientFactory +from services.factories import ( + CrossEncoderFactory, + DatabaseDriverFactory, + EmbedderFactory, + LLMClientFactory, +) from services.queue_service import QueueService from utils.formatting import format_fact_result, to_edge_result, to_node_result from utils.type_config import ( @@ -213,6 +218,15 @@ async def initialize(self) -> None: except Exception as e: logger.warning(f'Failed to create embedder client: {e}') + # Create cross-encoder (reranker) client matched to the configured provider. + # Without this, graphiti-core defaults to OpenAIRerankerClient(), which requires + # OPENAI_API_KEY even when the LLM and embedder are non-OpenAI providers. + cross_encoder_client = None + try: + cross_encoder_client = CrossEncoderFactory.create(self.config.llm) + except Exception as e: + logger.warning(f'Failed to create cross-encoder client: {e}') + # Get database configuration db_config = DatabaseDriverFactory.create_config(self.config.database) @@ -240,6 +254,7 @@ async def initialize(self) -> None: graph_driver=falkor_driver, llm_client=llm_client, embedder=embedder_client, + cross_encoder=cross_encoder_client, max_coroutines=self.semaphore_limit, ) else: @@ -250,6 +265,7 @@ async def initialize(self) -> None: password=db_config['password'], llm_client=llm_client, embedder=embedder_client, + cross_encoder=cross_encoder_client, max_coroutines=self.semaphore_limit, ) except Exception as db_error: diff --git a/mcp_server/src/services/factories.py b/mcp_server/src/services/factories.py index c3f60cd3b3..b58a4f84bf 100644 --- a/mcp_server/src/services/factories.py +++ b/mcp_server/src/services/factories.py @@ -65,6 +65,13 @@ except ImportError: HAS_GROQ = False +try: + from graphiti_core.cross_encoder.gemini_reranker_client import GeminiRerankerClient + + HAS_GEMINI_RERANKER = True +except ImportError: + HAS_GEMINI_RERANKER = False + def _validate_api_key(provider_name: str, api_key: str | None, logger) -> str: """Validate API key is present. @@ -293,6 +300,50 @@ def create(config: LLMConfig) -> LLMClient: raise ValueError(f'Unsupported LLM provider: {provider}') +class CrossEncoderFactory: + """Factory for creating CrossEncoder (reranker) clients based on the LLM provider. + + When no ``cross_encoder`` is passed to ``Graphiti()``, graphiti-core falls back to + ``OpenAIRerankerClient()``, which requires ``OPENAI_API_KEY`` — even when the LLM and + embedder are both configured for another provider. This factory returns a reranker that + matches the configured provider so a non-OpenAI stack (e.g. all-Gemini) does not pull in a + hard OpenAI dependency. Providers without a dedicated reranker return ``None``, which + preserves the existing default behavior. + """ + + @staticmethod + def create(config: LLMConfig): + """Create a cross-encoder client for the configured LLM provider, or None.""" + import logging + + logger = logging.getLogger(__name__) + + provider = config.provider.lower() + + match provider: + case 'gemini': + if not HAS_GEMINI_RERANKER: + logger.warning( + 'Gemini reranker not available in current graphiti-core version; ' + 'falling back to the default cross-encoder' + ) + return None + if not config.providers.gemini: + raise ValueError('Gemini provider configuration not found') + + api_key = config.providers.gemini.api_key + _validate_api_key('Gemini reranker', api_key, logger) + + # Reuse the fast/cheap default reranker model from graphiti-core. + reranker_config = GraphitiLLMConfig(api_key=api_key) + return GeminiRerankerClient(config=reranker_config) + + case _: + # No dedicated reranker for this provider — return None so graphiti-core + # applies its existing default (OpenAIRerankerClient). + return None + + class EmbedderFactory: """Factory for creating Embedder clients based on configuration."""