Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 13 additions & 1 deletion mcp_server/src/graphiti_mcp_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down
102 changes: 102 additions & 0 deletions mcp_server/src/services/factories.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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.

Expand Down
102 changes: 102 additions & 0 deletions mcp_server/tests/test_cross_encoder_factory.py
Original file line number Diff line number Diff line change
@@ -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()
Loading