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
19 changes: 19 additions & 0 deletions mcp_server/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -251,6 +251,25 @@ The `config.yaml` file supports environment variable expansion using `${VAR_NAME

You can set these variables in a `.env` file in the project directory.

### Search tuning (env vars, prefix `SEARCH__`)

These knobs tune the `search_memory_facts` and `search_nodes` tools. They are
nested config, so the env prefix is `SEARCH__` (double underscore).

| Var | Default | Purpose |
|---|---|---|
| `SEARCH__RERANKER` | `rrf` | `rrf` (reciprocal-rank fusion; relevance-ordered and diverse — the reliable default), `mmr` (diversity-aware but miscalibrated for low-magnitude embedding similarities — validate before use), or `cross_encoder`. |
| `SEARCH__MMR_LAMBDA` | `0.5` | MMR relevance/diversity tradeoff; 1.0 = no diversity. Only used when reranker is `mmr`. |
| `SEARCH__RERANKER_MIN_SCORE` | `0.0` | Min score to keep a result. Only meaningful for `cross_encoder`; ignored for `mmr` (which must not filter) and left at 0 for `rrf`. |
| `SEARCH__MAX_FACTS` | `6` | Default result count for `search_memory_facts`. |
| `SEARCH__MAX_NODES` | `6` | Default result count for `search_nodes`. |
| `SEARCH__EXCLUDE_INVALIDATED` | `true` | Exclude superseded/expired facts by default. Per-call override: `include_invalidated=true`. |

**Cross-encoder caveat:** `SEARCH__RERANKER=cross_encoder` gives a true relevance
floor but requires a cross-encoder client configured on the Graphiti instance and
adds a model call per search (latency + another provider dependency). It is **not**
the default; validate reliability before enabling.

## Running the Server

### Default Setup (FalkorDB Combined Container)
Expand Down
35 changes: 34 additions & 1 deletion mcp_server/src/config/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

import os
from pathlib import Path
from typing import Any
from typing import Any, Literal

import yaml
from pydantic import BaseModel, Field
Expand Down Expand Up @@ -246,6 +246,38 @@ class EdgeTypeMapEntry(BaseModel):
)


class SearchTuningConfig(BaseModel):
"""Tuning knobs for the MCP search tools (search_nodes / search_memory_facts)."""

reranker: Literal['rrf', 'mmr', 'cross_encoder'] = Field(
default='rrf',
description="Reranker for hybrid search. 'rrf' is plain reciprocal-rank "
"fusion (the reliable default: relevance-ordered, diverse). 'mmr' penalizes "
'redundancy but is miscalibrated for low-magnitude embedding similarities '
'(it can rank outliers over relevant clustered facts); validate before use. '
"'cross_encoder' scores true relevance but requires a configured "
'cross-encoder client and adds a model call per search.',
)
mmr_lambda: float = Field(
default=0.5,
description='MMR relevance/diversity tradeoff (1.0 = pure relevance, no '
'diversity; lower = more diversity). Only used when reranker == mmr.',
)
reranker_min_score: float = Field(
default=0.0,
description='Minimum reranker score to keep a result. Meaningful only for '
'cross_encoder (0-1 calibrated); ignored for mmr (which must not filter) '
'and left at 0 for rrf.',
)
max_facts: int = Field(default=6, description='Default max facts for search_memory_facts')
max_nodes: int = Field(default=6, description='Default max nodes for search_nodes')
exclude_invalidated: bool = Field(
default=True,
description='Exclude superseded/expired edges from fact search by default. '
'Clients can opt back in per call with include_invalidated=True.',
)


class GraphitiAppConfig(BaseModel):
"""Graphiti-specific configuration."""

Expand All @@ -270,6 +302,7 @@ class GraphitiConfig(BaseSettings):
embedder: EmbedderConfig = Field(default_factory=EmbedderConfig)
database: DatabaseConfig = Field(default_factory=DatabaseConfig)
graphiti: GraphitiAppConfig = Field(default_factory=GraphitiAppConfig)
search: SearchTuningConfig = Field(default_factory=SearchTuningConfig)

# Additional server options
destroy_graph: bool = Field(default=False, description='Clear graph on startup')
Expand Down
88 changes: 60 additions & 28 deletions mcp_server/src/graphiti_mcp_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,11 @@
from services.factories import DatabaseDriverFactory, EmbedderFactory, LLMClientFactory
from services.queue_service import QueueService
from utils.formatting import format_fact_result, to_edge_result, to_node_result
from utils.search_tuning import (
apply_liveness_filter,
build_edge_search_config,
build_node_search_config,
)
from utils.type_config import (
build_edge_type_map,
build_edge_types,
Expand Down Expand Up @@ -478,7 +483,7 @@ async def add_memory(
async def search_nodes(
query: str,
group_ids: str | list[str] | None = None,
max_nodes: int = 10,
max_nodes: int | None = None,
entity_types: list[str] | None = None,
center_node_uuid: str | None = None,
) -> NodeSearchResponse | ErrorResponse:
Expand All @@ -488,7 +493,8 @@ async def search_nodes(
query: The search query
group_ids: Optional group ID, or list of group IDs, to filter results (a single
string is accepted and treated as a one-element list)
max_nodes: Maximum number of nodes to return (default: 10)
max_nodes: Maximum number of nodes to return (defaults to the server's
search.max_nodes setting)
entity_types: Optional list of entity type names (node labels) to filter by
center_node_uuid: Optional UUID of a node to center the search around. Results
closer to this node in the graph are ranked higher.
Expand All @@ -499,6 +505,8 @@ async def search_nodes(
return ErrorResponse(error='Graphiti service not initialized')

try:
max_nodes = max_nodes if max_nodes is not None else config.search.max_nodes

client = await graphiti_service.get_client()

# Accept a scalar group_id or a list; fall back to the default when omitted.
Expand All @@ -511,22 +519,16 @@ async def search_nodes(
else []
)

# Create search filters
search_filters = SearchFilters(
node_labels=entity_types,
)
search_filters = SearchFilters(node_labels=entity_types)

# center_node_uuid is only honored by the node_distance reranker, so select
# that recipe when a center node is given (mirroring core's Graphiti.search);
# otherwise use RRF.
from graphiti_core.search.search_config_recipes import (
NODE_HYBRID_SEARCH_NODE_DISTANCE,
NODE_HYBRID_SEARCH_RRF,
node_config = build_node_search_config(
reranker=config.search.reranker,
mmr_lambda=config.search.mmr_lambda,
limit=max_nodes,
min_score=config.search.reranker_min_score,
center_node_uuid=center_node_uuid,
)

node_config = (
NODE_HYBRID_SEARCH_NODE_DISTANCE if center_node_uuid else NODE_HYBRID_SEARCH_RRF
)
results = await client.search_(
query=query,
config=node_config,
Expand All @@ -535,15 +537,16 @@ async def search_nodes(
search_filter=search_filters,
)

# Extract nodes from results
nodes = results.nodes[:max_nodes] if results.nodes else []
scores = results.node_reranker_scores[:max_nodes] if results.node_reranker_scores else []

if not nodes:
return NodeSearchResponse(message='No relevant nodes found', nodes=[])

# Format the results (embeddings stripped by to_node_result)
node_results = [to_node_result(node) for node in nodes]

node_results = [
to_node_result(node, score=(scores[i] if i < len(scores) else None))
for i, node in enumerate(nodes)
]
return NodeSearchResponse(message='Nodes retrieved successfully', nodes=node_results)
except Exception as e:
error_msg = str(e)
Expand All @@ -555,36 +558,48 @@ async def search_nodes(
async def search_memory_facts(
query: str,
group_ids: str | list[str] | None = None,
max_facts: int = 10,
max_facts: int | None = None,
center_node_uuid: str | None = None,
edge_types: list[str] | None = None,
valid_at_after: str | None = None,
valid_at_before: str | None = None,
invalid_at_after: str | None = None,
invalid_at_before: str | None = None,
include_invalidated: bool | None = None,
) -> FactSearchResponse | ErrorResponse:
"""Search the graph memory for relevant facts (entity edges).

Args:
query: The search query
group_ids: Optional group ID, or list of group IDs, to filter results (a single
string is accepted and treated as a one-element list)
max_facts: Maximum number of facts to return (default: 10)
max_facts: Maximum number of facts to return (defaults to the server's
search.max_facts setting)
center_node_uuid: Optional UUID of a node to center the search around
edge_types: Optional list of edge (fact) type names to filter by
valid_at_after: Optional ISO-8601 lower bound; only facts whose valid_at is at or
after this time are returned (timezone-naive is treated as UTC)
valid_at_before: Optional ISO-8601 upper bound on a fact's valid_at
invalid_at_after: Optional ISO-8601 lower bound on a fact's invalid_at
invalid_at_before: Optional ISO-8601 upper bound on a fact's invalid_at
include_invalidated: When True, include superseded/expired facts in
results (for historical/temporal reasoning). Defaults to the
server's search.exclude_invalidated setting (excluded by default).
"""
global graphiti_service

if graphiti_service is None:
return ErrorResponse(error='Graphiti service not initialized')

try:
# Validate max_facts parameter
# Resolve defaults from config.
max_facts = max_facts if max_facts is not None else config.search.max_facts
include_invalidated = (
include_invalidated
if include_invalidated is not None
else (not config.search.exclude_invalidated)
)

if max_facts <= 0:
return ErrorResponse(error='max_facts must be a positive integer')

Expand All @@ -600,6 +615,9 @@ async def search_memory_facts(
except ValueError as e:
return ErrorResponse(error=f'Invalid date filter: {e}')

if not include_invalidated:
search_filter = apply_liveness_filter(search_filter)

client = await graphiti_service.get_client()

# Accept a scalar group_id or a list; fall back to the default when omitted.
Expand All @@ -612,18 +630,32 @@ async def search_memory_facts(
else []
)

relevant_edges = await client.search(
group_ids=effective_group_ids,
edge_config = build_edge_search_config(
reranker=config.search.reranker,
mmr_lambda=config.search.mmr_lambda,
limit=max_facts,
min_score=config.search.reranker_min_score,
center_node_uuid=center_node_uuid,
)

results = await client.search_(
query=query,
num_results=max_facts,
config=edge_config,
group_ids=effective_group_ids,
center_node_uuid=center_node_uuid,
search_filter=search_filter,
search_filter=search_filter if search_filter is not None else SearchFilters(),
)

if not relevant_edges:
edges = results.edges[:max_facts] if results.edges else []
scores = results.edge_reranker_scores[:max_facts] if results.edge_reranker_scores else []

if not edges:
return FactSearchResponse(message='No relevant facts found', facts=[])

facts = [format_fact_result(edge) for edge in relevant_edges]
facts = [
format_fact_result(edge, score=(scores[i] if i < len(scores) else None))
for i, edge in enumerate(edges)
]
return FactSearchResponse(message='Facts retrieved successfully', facts=facts)
except Exception as e:
error_msg = str(e)
Expand Down
2 changes: 2 additions & 0 deletions mcp_server/src/models/response_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ class NodeResult(TypedDict):
summary: str | None
group_id: str
attributes: dict[str, Any]
score: float | None


class NodeSearchResponse(TypedDict):
Expand Down Expand Up @@ -74,6 +75,7 @@ class EdgeResult(TypedDict):
created_at: str | None
valid_at: str | None
invalid_at: str | None
score: float | None


class TripletResponse(TypedDict):
Expand Down
12 changes: 8 additions & 4 deletions mcp_server/src/utils/formatting.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
from models.response_types import EdgeResult, NodeResult


def to_node_result(node: EntityNode) -> NodeResult:
def to_node_result(node: EntityNode, score: float | None = None) -> NodeResult:
"""Build a NodeResult TypedDict from an EntityNode, dropping embeddings."""
attrs = node.attributes if node.attributes else {}
attrs = {k: v for k, v in attrs.items() if 'embedding' not in k.lower()}
Expand All @@ -20,10 +20,11 @@ def to_node_result(node: EntityNode) -> NodeResult:
summary=node.summary,
group_id=node.group_id,
attributes=attrs,
score=score,
)


def to_edge_result(edge: EntityEdge) -> EdgeResult:
def to_edge_result(edge: EntityEdge, score: float | None = None) -> EdgeResult:
"""Build an EdgeResult TypedDict from an EntityEdge."""
return EdgeResult(
uuid=edge.uuid,
Expand All @@ -35,6 +36,7 @@ def to_edge_result(edge: EntityEdge) -> EdgeResult:
created_at=edge.created_at.isoformat() if edge.created_at else None,
valid_at=edge.valid_at.isoformat() if edge.valid_at else None,
invalid_at=edge.invalid_at.isoformat() if edge.invalid_at else None,
score=score,
)


Expand All @@ -61,13 +63,14 @@ def format_node_result(node: EntityNode) -> dict[str, Any]:
return result


def format_fact_result(edge: EntityEdge) -> dict[str, Any]:
"""Format an entity edge into a readable result.
def format_fact_result(edge: EntityEdge, score: float | None = None) -> dict[str, Any]:
"""Format an entity edge into a readable result, including its rerank score.

Since EntityEdge is a Pydantic BaseModel, we can use its built-in serialization capabilities.

Args:
edge: The EntityEdge to format
score: Optional reranker score for this edge

Returns:
A dictionary representation of the edge with serialized dates and excluded embeddings
Expand All @@ -79,4 +82,5 @@ def format_fact_result(edge: EntityEdge) -> dict[str, Any]:
},
)
result.get('attributes', {}).pop('fact_embedding', None)
result['score'] = score
return result
Loading
Loading