diff --git a/mcp_server/src/graphiti_mcp_server.py b/mcp_server/src/graphiti_mcp_server.py index efd9536de3..5261c68e87 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,11 @@ async def initialize(self) -> None: except Exception as e: logger.warning(f'Failed to create embedder client: {e}') + # Create cross-encoder (reranker) client. Without this, Graphiti defaults to + # OpenAIRerankerClient, which needs an OpenAI API key even on non-OpenAI setups. + # Reranker setup errors must remain fatal rather than silently restoring that default. + cross_encoder_client = CrossEncoderFactory.create(self.config.llm, self.config.embedder) + # Get database configuration db_config = DatabaseDriverFactory.create_config(self.config.database) @@ -240,6 +250,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 +261,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..6119a2a95d 100644 --- a/mcp_server/src/services/factories.py +++ b/mcp_server/src/services/factories.py @@ -1,5 +1,6 @@ """Factory classes for creating LLM, Embedder, and Database clients.""" +from graphiti_core.cross_encoder.client import CrossEncoderClient from graphiti_core.embedder import EmbedderClient, OpenAIEmbedder from graphiti_core.llm_client import LLMClient, OpenAIClient from graphiti_core.llm_client.config import LLMConfig as GraphitiLLMConfig @@ -405,6 +406,107 @@ def create(config: EmbedderConfig) -> EmbedderClient: raise ValueError(f'Unsupported Embedder provider: {provider}') +class CrossEncoderFactory: + """Factory for creating cross-encoder (reranker) clients based on configuration. + + Graphiti defaults the cross_encoder to OpenAIRerankerClient, which needs an OpenAI API key. + To keep the server usable on non-OpenAI setups, pick a reranker from the LLM provider, then + the embedder provider, and fall back to the local BGE reranker. + """ + + @staticmethod + def create(llm_config: LLMConfig, embedder_config: EmbedderConfig) -> CrossEncoderClient: + """Create a cross-encoder client based on the configured providers.""" + import logging + + logger = logging.getLogger(__name__) + + # Try the LLM provider first, then the embedder, before falling back to a local model. + for source, config in (('LLM', llm_config), ('embedder', embedder_config)): + reranker = CrossEncoderFactory._reranker_for_provider(source, config, logger) + if reranker is not None: + return reranker + + # No provider reranker available (e.g. Anthropic LLM + Voyage embedder), so use the + # local BGE cross-encoder, which needs no API key. + logger.warning( + 'No provider reranker available, using local BGERerankerClient ' + '(downloads BAAI/bge-reranker-v2-m3, ~2.3 GB, on first run)' + ) + try: + from graphiti_core.cross_encoder.bge_reranker_client import BGERerankerClient + except ImportError as e: + raise ValueError( + 'No provider reranker is available for this configuration, and the local ' + 'BGE fallback requires the optional sentence-transformers dependency. ' + "Install the MCP server's 'providers' extra (uv sync --extra providers), " + "install graphiti-core's 'sentence-transformers' extra " + "(pip install 'graphiti-core[sentence-transformers]'), or configure the " + 'LLM or embedder provider as OpenAI or Gemini with a valid API key.' + ) from e + + return BGERerankerClient() + + @staticmethod + def _reranker_for_provider( + source: str, config: LLMConfig | EmbedderConfig, logger + ) -> CrossEncoderClient | None: + """Return a reranker for this provider, or None if it has no native one.""" + provider = config.provider.lower() + + match provider: + case 'openai': + if not config.providers.openai: + return None + from graphiti_core.cross_encoder.openai_reranker_client import ( + OpenAIRerankerClient, + ) + + logger.info(f'Using OpenAIRerankerClient from {source} provider') + return OpenAIRerankerClient( + config=GraphitiLLMConfig( + api_key=config.providers.openai.api_key, + base_url=config.providers.openai.api_url, + ) + ) + + case 'azure_openai': + azure_config = config.providers.azure_openai + if not azure_config or not azure_config.api_url: + return None + from openai import AsyncOpenAI + + base_url = azure_config.api_url + if not base_url.endswith('/'): + base_url += '/' + if not base_url.endswith('openai/v1/'): + base_url += 'openai/v1/' + from graphiti_core.cross_encoder.openai_reranker_client import ( + OpenAIRerankerClient, + ) + + logger.info(f'Using OpenAIRerankerClient (Azure) from {source} provider') + return OpenAIRerankerClient( + client=AsyncOpenAI(base_url=base_url, api_key=azure_config.api_key) + ) + + case 'gemini': + if not config.providers.gemini: + return None + from graphiti_core.cross_encoder.gemini_reranker_client import ( + GeminiRerankerClient, + ) + + logger.info(f'Using GeminiRerankerClient from {source} provider') + return GeminiRerankerClient( + config=GraphitiLLMConfig(api_key=config.providers.gemini.api_key) + ) + + case _: + # anthropic, groq, voyage etc. have no native reranker in graphiti-core. + return None + + class DatabaseDriverFactory: """Factory for creating Database drivers based on configuration. diff --git a/mcp_server/tests/test_cross_encoder_factory.py b/mcp_server/tests/test_cross_encoder_factory.py new file mode 100644 index 0000000000..71c30a1b21 --- /dev/null +++ b/mcp_server/tests/test_cross_encoder_factory.py @@ -0,0 +1,102 @@ +#!/usr/bin/env python3 +"""Unit tests for CrossEncoderFactory reranker selection.""" + +import builtins +import logging +import sys +from pathlib import Path +from unittest.mock import AsyncMock, Mock + +import pytest + +# Add the src directory to the path (mirrors the other factory tests) +sys.path.insert(0, str(Path(__file__).parent.parent / 'src')) + +from graphiti_core.cross_encoder.gemini_reranker_client import GeminiRerankerClient +from graphiti_core.cross_encoder.openai_reranker_client import OpenAIRerankerClient + +import graphiti_mcp_server +from config.schema import ( + AnthropicProviderConfig, + DatabaseConfig, + EmbedderConfig, + EmbedderProvidersConfig, + GeminiProviderConfig, + GraphitiConfig, + LLMConfig, + LLMProvidersConfig, + OpenAIProviderConfig, + VoyageProviderConfig, +) +from services.factories import CrossEncoderFactory + + +class TestCrossEncoderFactory: + """The reranker is inferred from the providers, so a non-OpenAI setup does not need OPENAI_API_KEY.""" + + def test_openai_llm_uses_openai_reranker(self): + llm = LLMConfig( + provider='openai', + providers=LLMProvidersConfig(openai=OpenAIProviderConfig(api_key='test-key')), + ) + embedder = EmbedderConfig( + provider='openai', + providers=EmbedderProvidersConfig(openai=OpenAIProviderConfig(api_key='test-key')), + ) + assert isinstance(CrossEncoderFactory.create(llm, embedder), OpenAIRerankerClient) + + def test_anthropic_llm_falls_back_to_gemini_embedder(self): + # Anthropic has no native reranker, so the factory should pick up the Gemini embedder's + # key instead of defaulting to OpenAIRerankerClient (which would need OPENAI_API_KEY). + llm = LLMConfig( + provider='anthropic', + providers=LLMProvidersConfig(anthropic=AnthropicProviderConfig(api_key='test-key')), + ) + embedder = EmbedderConfig( + provider='gemini', + providers=EmbedderProvidersConfig(gemini=GeminiProviderConfig(api_key='test-key')), + ) + assert isinstance(CrossEncoderFactory.create(llm, embedder), GeminiRerankerClient) + + def test_missing_local_reranker_dependency_is_actionable(self, monkeypatch, caplog): + llm = LLMConfig( + provider='anthropic', + providers=LLMProvidersConfig(anthropic=AnthropicProviderConfig(api_key='test-key')), + ) + embedder = EmbedderConfig( + provider='voyage', + providers=EmbedderProvidersConfig(voyage=VoyageProviderConfig(api_key='test-key')), + ) + real_import = builtins.__import__ + + def import_without_bge(name, *args, **kwargs): + if name == 'graphiti_core.cross_encoder.bge_reranker_client': + raise ImportError('sentence-transformers is not installed') + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, '__import__', import_without_bge) + caplog.set_level(logging.INFO) + + with pytest.raises(ValueError, match="MCP server's 'providers' extra"): + CrossEncoderFactory.create(llm, embedder) + + assert '~2.3 GB' in caplog.text + + +@pytest.mark.asyncio +async def test_graphiti_service_does_not_swallow_reranker_configuration_error(monkeypatch): + error = ValueError('reranker setup failed') + + def fail_reranker_setup(*_args): + raise error + + fake_client = Mock() + fake_client.build_indices_and_constraints = AsyncMock() + monkeypatch.setattr(CrossEncoderFactory, 'create', fail_reranker_setup) + monkeypatch.setattr(graphiti_mcp_server, 'Graphiti', Mock(return_value=fake_client)) + service = graphiti_mcp_server.GraphitiService( + GraphitiConfig(database=DatabaseConfig(provider='neo4j')) + ) + + with pytest.raises(ValueError, match='reranker setup failed'): + await service.initialize()