fix(mcp): make multi-project search opt-in (#807)

Signed-off-by: phernandez <paul@basicmachines.co>
This commit is contained in:
Paul Hernandez
2026-05-08 12:54:20 -05:00
committed by GitHub
parent 26381aeed1
commit 831dc1ecdc
12 changed files with 739 additions and 55 deletions
+5 -13
View File
@@ -439,19 +439,16 @@ def _canonical_memory_path_for_workspace(
workspace_type: str,
project_permalink: str,
remainder: str,
include_project: bool,
) -> str:
"""Return the stored canonical path for a workspace-qualified memory URL."""
normalized_remainder = remainder.strip("/")
if workspace_type == "organization":
prefix = f"{generate_permalink(workspace_slug)}/{project_permalink}"
elif workspace_type == "personal":
prefix = project_permalink if include_project else ""
else:
if workspace_type not in {"organization", "personal"}:
raise ValueError(f"Unsupported workspace_type for memory URL routing: {workspace_type}")
if not prefix:
return normalized_remainder
# Trigger: a caller supplied a workspace-qualified memory URL.
# Why: the first two path segments are the global route, even for Personal.
# Outcome: lookups preserve the complete workspace/project canonical permalink.
prefix = f"{generate_permalink(workspace_slug)}/{project_permalink}"
if not normalized_remainder:
return prefix
return f"{prefix}/{normalized_remainder}"
@@ -483,7 +480,6 @@ def _canonical_memory_path_for_active_route(
workspace_type=workspace_context.workspace_type,
project_permalink=active_project.permalink,
remainder=workspace_remainder,
include_project=include_project,
)
if cached_workspace is not None:
@@ -492,7 +488,6 @@ def _canonical_memory_path_for_active_route(
workspace_type=cached_workspace.workspace_type,
project_permalink=active_project.permalink,
remainder=workspace_remainder,
include_project=include_project,
)
if not include_project:
@@ -582,7 +577,6 @@ async def resolve_workspace_qualified_memory_url(
workspace_type=entry.workspace.workspace_type,
project_permalink=entry.project.permalink,
remainder=remainder,
include_project=ConfigManager().config.permalinks_include_project,
)
return WorkspaceMemoryUrlResolution(entry=entry, canonical_path=canonical_path)
@@ -1130,7 +1124,6 @@ async def resolve_project_and_path(
workspace_type=cached_workspace.workspace_type,
project_permalink=cached_project.permalink,
remainder=remainder,
include_project=bool(include_project),
)
return cached_project, resolved_path, True
@@ -1153,7 +1146,6 @@ async def resolve_project_and_path(
workspace_type=workspace_context.workspace_type,
project_permalink=project_permalink,
remainder=remainder,
include_project=bool(include_project),
)
return active_project, resolved_path, True
+29 -18
View File
@@ -6,7 +6,7 @@ a list containing a single `{"type": "text", "text": "{...json...}"}` item.
"""
import json
from typing import Any, Dict, List, Optional
from typing import Any, Dict, List, Optional, cast
from fastmcp import Context
from loguru import logger
@@ -17,18 +17,30 @@ from basic_memory.mcp.tools.search import search_notes
from basic_memory.schemas.search import SearchResponse, SearchResult
def _identifier_for_read_note(identifier: str) -> str:
"""Convert ChatGPT result ids into routable Basic Memory identifiers."""
stripped = identifier.strip()
if stripped.startswith("memory://") or "/" not in stripped:
return identifier
return f"memory://{stripped}"
def _format_search_results_for_chatgpt(
results: SearchResponse | list[SearchResult] | list[dict[str, Any]] | dict[str, Any],
results: SearchResponse | list[SearchResult | dict[str, Any]] | dict[str, Any],
) -> List[Dict[str, Any]]:
"""Format search results according to ChatGPT's expected schema.
Returns a list of result objects with id, title, and url fields.
"""
if isinstance(results, SearchResponse):
raw_results: list[SearchResult] | list[dict[str, Any]] = results.results
raw_results: list[SearchResult | dict[str, Any]] = list(results.results)
elif isinstance(results, dict):
nested_results = results.get("results")
raw_results = nested_results if isinstance(nested_results, list) else []
raw_results = (
cast(list[SearchResult | dict[str, Any]], nested_results)
if isinstance(nested_results, list)
else []
)
else:
raw_results = results
@@ -113,8 +125,7 @@ async def search(
logger.info(f"ChatGPT search request: query='{query}'")
try:
# Let search_notes resolve the default project via get_project_client(),
# which works in both local mode (ConfigManager) and cloud mode (database).
# Keep this adapter tiny: the real search behavior lives in search_notes.
results = await search_notes(
query=query,
page=1,
@@ -123,7 +134,6 @@ async def search(
context=context,
)
# Handle string error responses from search_notes
if isinstance(results, str):
logger.warning(f"Search failed with error: {results[:100]}...")
search_results = {
@@ -131,16 +141,17 @@ async def search(
"error": "Search failed",
"error_details": results[:500], # Truncate long error messages
}
else:
# Format successful results for ChatGPT
raw_results = results.get("results", []) if isinstance(results, dict) else []
formatted_results = _format_search_results_for_chatgpt(raw_results)
search_results = {
"results": formatted_results,
"total_count": len(raw_results), # Use actual count from results
"query": query,
}
logger.info(f"Search completed: {len(formatted_results)} results returned")
return [{"type": "text", "text": json.dumps(search_results, ensure_ascii=False)}]
raw_results = results.get("results", []) if isinstance(results, dict) else []
formatted_results = _format_search_results_for_chatgpt(raw_results)
search_results = {
"results": formatted_results,
"total_count": len(raw_results), # Use actual count from results
"query": query,
}
logger.info(f"Search completed: {len(formatted_results)} results returned")
# Return in MCP content array format as required by OpenAI
return [{"type": "text", "text": json.dumps(search_results, ensure_ascii=False)}]
@@ -180,7 +191,7 @@ async def fetch(
# which works in both local mode (ConfigManager) and cloud mode (database).
content = str(
await read_note(
identifier=id,
identifier=_identifier_for_read_note(id),
context=context,
)
)
+44 -16
View File
@@ -18,7 +18,8 @@ from basic_memory.mcp.project_context import (
from basic_memory.mcp.server import mcp
from basic_memory.mcp.tools.search import search_notes
from basic_memory.schemas.memory import memory_url_path
from basic_memory.utils import validate_project_path
from basic_memory.utils import generate_permalink, validate_project_path
from basic_memory.workspace_context import current_workspace_permalink_context
def _is_exact_title_match(identifier: str, title: str) -> bool:
@@ -277,24 +278,51 @@ async def read_note(
value = item.get("file_path")
return str(value) if value else None
try:
# Try to resolve identifier to entity ID
entity_id = await knowledge_client.resolve_entity(entity_path, strict=True)
def _legacy_workspace_unqualified_path(path: str) -> str | None:
workspace_context = current_workspace_permalink_context()
if workspace_context is None:
return None
# Fetch content using entity ID
response = await resource_client.read(entity_id)
workspace_prefix = generate_permalink(workspace_context.workspace_slug)
project_prefix = active_project.permalink
qualified_prefix = f"{workspace_prefix}/{project_prefix}"
normalized_path = path.strip("/")
if normalized_path == qualified_prefix:
return project_prefix
if normalized_path.startswith(f"{qualified_prefix}/"):
return f"{project_prefix}/{normalized_path.removeprefix(f'{qualified_prefix}/')}"
return None
# If successful, return the content
if response.status_code == 200:
logger.info(
"Returning read_note result from resource: {path}", path=entity_path
direct_lookup_paths = [entity_path]
legacy_path = _legacy_workspace_unqualified_path(entity_path)
if legacy_path and legacy_path not in direct_lookup_paths:
# Trigger: existing cloud rows may still use project-prefixed permalinks.
# Why: new workspace-qualified IDs should read old notes without a re-sync.
# Outcome: try the legacy path after the canonical workspace path misses.
direct_lookup_paths.append(legacy_path)
for direct_lookup_path in direct_lookup_paths:
try:
# Try to resolve identifier to entity ID
entity_id = await knowledge_client.resolve_entity(
direct_lookup_path, strict=True
)
if output_format == "json":
return await _read_json_payload(entity_id)
return response.text
except Exception as e: # pragma: no cover
logger.info(f"Direct lookup failed for '{entity_path}': {e}")
# Continue to fallback methods
# Fetch content using entity ID
response = await resource_client.read(entity_id)
# If successful, return the content
if response.status_code == 200:
logger.info(
"Returning read_note result from resource: {path}",
path=direct_lookup_path,
)
if output_format == "json":
return await _read_json_payload(entity_id)
return response.text
except Exception as e: # pragma: no cover
logger.info(f"Direct lookup failed for '{direct_lookup_path}': {e}")
# Continue to alternate direct lookup paths, then fallback methods
# Fallback 1: Try title search via API
logger.info(f"Search title for: {identifier}")
+278 -1
View File
@@ -2,7 +2,8 @@
import re
from textwrap import dedent
from typing import Annotated, List, Optional, Dict, Any, Literal
from typing import Annotated, List, Optional, Dict, Any, Literal, cast
from uuid import UUID
import logfire
from loguru import logger
@@ -22,6 +23,7 @@ from basic_memory.schemas.search import (
SearchItemType,
SearchQuery,
SearchResponse,
SearchResult,
SearchRetrievalMode,
)
@@ -294,6 +296,247 @@ def _format_search_markdown(result: SearchResponse, project: str, query: str | N
return "\n".join(parts)
def _valid_project_id(value: object) -> str | None:
"""Return a UUID project id string when one is present."""
if not isinstance(value, str) or not value.strip():
return None
try:
return str(UUID(value))
except ValueError:
return None
def _matches_constrained_project(project: dict[str, Any], constrained_project: object) -> bool:
"""Return True when a project list row satisfies BASIC_MEMORY_MCP_PROJECT."""
if not isinstance(constrained_project, str) or not constrained_project.strip():
return True
candidates = {
value
for value in (
project.get("name"),
project.get("qualified_name"),
project.get("external_id"),
)
if isinstance(value, str)
}
return constrained_project in candidates
def _search_project_refs(projects_payload: object) -> list[dict[str, str | None]]:
"""Extract project routing refs for optional account-scoped search."""
if not isinstance(projects_payload, dict):
return []
payload = cast(dict[str, Any], projects_payload)
projects = payload.get("projects")
if not isinstance(projects, list):
return []
refs: list[dict[str, str | None]] = []
seen: set[tuple[str | None, str | None]] = set()
constrained_project = payload.get("constrained_project")
for item in projects:
if not isinstance(item, dict) or not _matches_constrained_project(
item, constrained_project
):
continue
project = item.get("qualified_name") or item.get("name")
project_name = project if isinstance(project, str) and project.strip() else None
project_id = _valid_project_id(item.get("external_id"))
if project_name is None and project_id is None:
continue
key = (project_name, project_id)
if key in seen:
continue
seen.add(key)
refs.append({"project": project_name, "project_id": project_id})
return refs
async def _load_search_project_refs(context: Context | None = None) -> list[dict[str, str | None]]:
"""Load accessible projects for search_all_projects without coupling the wrapper tool."""
from basic_memory.mcp.tools.project_management import list_memory_projects
return _search_project_refs(await list_memory_projects(output_format="json", context=context))
def _raw_results_from_search_payload(
results: SearchResponse | list[SearchResult | dict[str, Any]] | dict[str, Any],
) -> list[SearchResult | dict[str, Any]]:
"""Return the result list from any search_notes JSON-compatible payload."""
if isinstance(results, SearchResponse):
return list(results.results)
if isinstance(results, dict):
nested_results = results.get("results")
return (
cast(list[SearchResult | dict[str, Any]], nested_results)
if isinstance(nested_results, list)
else []
)
return list(results)
def _result_score(result: SearchResult | dict[str, Any]) -> float:
"""Return a comparable search score for merged project results."""
if isinstance(result, SearchResult):
return result.score
score = result.get("score")
return float(score) if isinstance(score, int | float) else 0.0
def _qualify_permalink_for_project(permalink: object, project: str | None) -> object:
"""Return a workspace-qualified permalink when the project ref supplies one."""
if not isinstance(permalink, str) or not permalink.strip():
return permalink
if not isinstance(project, str) or "/" not in project.strip("/"):
return permalink
normalized_permalink = permalink.strip("/")
qualified_project = project.strip("/")
if normalized_permalink == qualified_project or normalized_permalink.startswith(
f"{qualified_project}/"
):
return normalized_permalink
workspace_slug, project_permalink = qualified_project.split("/", 1)
if normalized_permalink == project_permalink or normalized_permalink.startswith(
f"{project_permalink}/"
):
return f"{workspace_slug}/{normalized_permalink}"
return f"{qualified_project}/{normalized_permalink}"
def _qualify_results_for_project(
results: list[SearchResult | dict[str, Any]],
project_ref: dict[str, str | None],
) -> list[dict[str, Any]]:
"""Attach the searched workspace/project prefix to each result permalink."""
qualified: list[dict[str, Any]] = []
for result in results:
if isinstance(result, SearchResult):
result_data = result.model_dump()
else:
result_data = dict(result)
result_data["permalink"] = _qualify_permalink_for_project(
result_data.get("permalink"),
project_ref.get("project"),
)
qualified.append(result_data)
return qualified
def _result_total(results: dict[str, Any], raw_results: list[SearchResult | dict[str, Any]]) -> int:
"""Return the best available total for a per-project search payload."""
total = results.get("total")
if isinstance(total, int) and total > 0:
return total
return len(raw_results) + (1 if results.get("has_more") is True else 0)
def _project_ref_label(project_ref: dict[str, str | None]) -> str:
"""Return a stable log label for a project search ref."""
return project_ref.get("project") or project_ref.get("project_id") or "<unknown project>"
async def _search_all_projects(
*,
query: str | None,
page: int,
page_size: int,
search_type: str | None,
output_format: Literal["text", "json"],
note_types: list[str],
entity_types: list[str],
after_date: str | None,
metadata_filters: dict[str, Any] | None,
tags: list[str] | None,
status: str | None,
min_similarity: float | None,
context: Context | None,
) -> dict | str:
"""Search every accessible project when the caller explicitly opts in."""
requested_page = max(page, 1)
requested_page_size = max(page_size, 1)
project_refs = await _load_search_project_refs(context=context)
if not project_refs:
response = SearchResponse(
results=[],
current_page=requested_page,
page_size=requested_page_size,
total=0,
has_more=False,
)
if output_format == "json":
return response.model_dump(mode="json", exclude_none=True)
return _format_search_markdown(response, "all projects", query)
per_project_page_size = requested_page * requested_page_size
merged_results: list[dict[str, Any]] = []
total = 0
any_project_has_more = False
for project_ref in project_refs:
try:
results = await search_notes(
query=query,
project=project_ref["project"],
project_id=project_ref["project_id"],
page=1,
page_size=per_project_page_size,
search_type=search_type,
output_format="json",
note_types=note_types or None,
entity_types=entity_types or None,
after_date=after_date,
metadata_filters=metadata_filters,
tags=tags,
status=status,
min_similarity=min_similarity,
search_all_projects=False,
context=context,
)
except Exception as exc:
logger.warning(
f"Multi-project search failed for project {_project_ref_label(project_ref)}: {exc}"
)
continue
if isinstance(results, str):
if not results.startswith("# Search Failed"):
return results
logger.warning(
"Multi-project search failed for project "
f"{_project_ref_label(project_ref)}: {results}"
)
continue
raw_results = _raw_results_from_search_payload(results)
total += _result_total(results, raw_results)
any_project_has_more = any_project_has_more or results.get("has_more") is True
merged_results.extend(_qualify_results_for_project(raw_results, project_ref))
sorted_results = sorted(merged_results, key=_result_score, reverse=True)
start = (requested_page - 1) * requested_page_size
end = start + requested_page_size
paged_results = sorted_results[start:end]
response = SearchResponse.model_validate(
{
"results": paged_results,
"current_page": requested_page,
"page_size": requested_page_size,
"total": total,
"has_more": any_project_has_more or total > end or len(sorted_results) > end,
}
)
if output_format == "json":
return response.model_dump(mode="json", exclude_none=True)
return _format_search_markdown(response, "all projects", query)
@mcp.tool(
description="Search across all content in the knowledge base with advanced syntax support.",
# TODO: re-enable once MCP client rendering is working
@@ -309,6 +552,13 @@ async def search_notes(
] = None,
project: Optional[str] = None,
project_id: Optional[str] = None,
search_all_projects: Annotated[
bool,
Field(
default=False,
validation_alias=AliasChoices("search_all_projects", "all_projects"),
),
] = False,
# `offset` is intentionally NOT aliased to `page`: offset is item-indexed
# (skip N items) while page is 1-indexed page-number. Direct aliasing would
# silently return the wrong slice.
@@ -373,6 +623,8 @@ async def search_notes(
Project Resolution:
Server resolves projects in this order: Single Project Mode → project parameter → default project.
If project unknown, use list_memory_projects() or recent_activity() first.
Set search_all_projects=True to search every accessible project; this is opt-in because it
performs one search per project.
## Search Syntax Examples
@@ -448,6 +700,8 @@ async def search_notes(
project_id: Project external_id (UUID). Prefer this over `project` when known —
it routes to the exact project regardless of name collisions across cloud
workspaces. Takes precedence over `project`. Get from list_memory_projects().
search_all_projects: Optional opt-in to search every accessible project. Ignored when
`project` or `project_id` is supplied.
page: The page number of results to return (default 1)
page_size: The number of results to return per page (default 10)
search_type: Type of search to perform, one of:
@@ -562,12 +816,35 @@ async def search_notes(
if detected:
project = detected
# Trigger: caller explicitly requests account/workspace-wide search and did not
# already provide a concrete project route.
# Why: multi-project fan-out can be slow, so default search remains project-scoped.
# Outcome: run one normal search per accessible project and merge ranked results.
if search_all_projects and project is None and project_id is None:
all_projects_result = await _search_all_projects(
query=query,
page=page,
page_size=page_size,
search_type=search_type,
output_format=output_format,
note_types=note_types,
entity_types=entity_types,
after_date=after_date,
metadata_filters=metadata_filters,
tags=tags,
status=status,
min_similarity=min_similarity,
context=context,
)
return all_projects_result
with logfire.span(
"mcp.tool.search_notes",
entrypoint="mcp",
tool_name="search_notes",
requested_project=project,
requested_project_id=project_id,
search_all_projects=search_all_projects,
search_type=search_type or "default",
output_format=output_format,
page=page,
+2 -2
View File
@@ -14,14 +14,14 @@ _WORKSPACE_TYPES = {"personal", "organization"}
@dataclass(frozen=True)
class WorkspacePermalinkContext:
"""Workspace metadata needed to build canonical organization permalinks."""
"""Workspace metadata needed to build canonical workspace permalinks."""
workspace_slug: str
workspace_type: str
@property
def should_prefix_permalinks(self) -> bool:
return self.workspace_type == "organization" and bool(self.workspace_slug)
return bool(self.workspace_slug)
_workspace_permalink_context: ContextVar[WorkspacePermalinkContext | None] = ContextVar(