Skip to content
Closed
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
18 changes: 17 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,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)

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

Expand Down
Loading