fix tests, add SearchResult response type for search

This commit is contained in:
phernandez
2025-01-04 22:22:38 -06:00
parent e35dbab156
commit fa7e986d3b
10 changed files with 235 additions and 472 deletions
@@ -5,18 +5,19 @@ from typing import List
from loguru import logger
from basic_memory.services.search_service import SearchService
from basic_memory.schemas.search import SearchQuery, SearchResult
from basic_memory.schemas.search import SearchQuery, SearchResult, SearchResponse
from basic_memory.deps import get_search_service
router = APIRouter(prefix="/search", tags=["search"])
@router.post("/", response_model=List[SearchResult])
@router.post("/", response_model=SearchResponse)
async def search(
query: SearchQuery,
search_service: SearchService = Depends(get_search_service)
):
"""Search across all knowledge and documents."""
return await search_service.search(query)
results = await search_service.search(query)
return SearchResponse(results=results)
@router.post("/reindex")
async def reindex(
+2 -2
View File
@@ -24,7 +24,7 @@ from basic_memory.mcp.tools.knowledge import (
)
from basic_memory.mcp.tools.search import (
search_nodes,
search,
open_nodes,
)
@@ -56,7 +56,7 @@ __all__ = [
"delete_relations",
# Search tools
"search_nodes",
"search",
"get_entity",
"open_nodes",
+145 -241
View File
@@ -1,186 +1,151 @@
"""Search and query tools for Basic Memory MCP server."""
"""Search tools for Basic Memory MCP server."""
from typing import List, Optional
from datetime import datetime, timezone
import textwrap
from collections import defaultdict
from basic_memory.mcp.server import mcp
from basic_memory.schemas.request import SearchNodesRequest, OpenNodesRequest
from basic_memory.schemas.response import SearchNodesResponse, EntityListResponse
from basic_memory.schemas.search import SearchQuery, SearchResult, SearchItemType
from basic_memory.schemas.request import OpenNodesRequest
from basic_memory.schemas.response import EntityListResponse
from basic_memory.mcp.async_client import client
@mcp.tool(
category="search",
description="Search for entities across names, descriptions, observations, and relations",
description="Search across all content in basic-memory, including documents and entities",
examples=[
{
"name": "Technical Search",
"description": "Find implementation details and patterns",
"name": "Search with Metadata Analysis",
"description": "Search and analyze results by metadata",
"code": """
# Search for database-related components
results = await search_nodes(
request=SearchNodesRequest(
query="sqlite database implementation",
category="tech" # Focus on technical details
)
# Search for feature specs
results = await search(
text="implementation",
types=[SearchItemType.DOCUMENT]
)
# Analyze implementation patterns
for entity in results.matches:
print(f"\\n{entity.name} Implementation:")
# Technical details
tech_notes = [o.content for o in entity.observations
if o.category == "tech"]
if tech_notes:
print("Technical Notes:")
for note in tech_notes:
print(f"- {note}")
# Dependencies
deps = [r for r in entity.relations
if r.relation_type == "depends_on"]
if deps:
print("\\nDependencies:")
for dep in deps:
print(f"- {dep.to_id}")"""
# Group by category and status
by_category = defaultdict(list)
by_status = defaultdict(list)
for r in results:
meta = r.metadata
if 'category' in meta:
by_category[meta['category']].append(r)
if 'status' in meta:
by_status[meta['status']].append(r)
print("Results by Category:")
for category, items in by_category.items():
print(f"\\n{category.title()}:")
for item in items:
print(f"- {item.path_id} (score: {item.score:.2f})")
# Find high priority items
high_priority = [
r for r in results
if r.metadata.get('priority') in ['high', 'highest']
]
"""
},
{
"name": "Feature Context",
"description": "Build complete feature implementation context",
"name": "Recent Changes Analysis",
"description": "Search and analyze recent document changes",
"code": """
# Start with feature search
feature_results = await search_nodes(
request=SearchNodesRequest(
query="semantic search feature"
)
from datetime import datetime, timedelta
# Set cutoff date
cutoff = datetime.now(timezone.utc) - timedelta(days=7)
# Search for recent changes
results = await search(text="database", after_date=cutoff)
# Sort by update time
sorted_results = sorted(
results,
key=lambda x: x.metadata['updated_at'],
reverse=True
)
# Collect related entities for context
related_ids = set()
for entity in feature_results.matches:
# Add feature itself
related_ids.add(entity.path_id)
# Add related entities
for relation in entity.relations:
related_ids.add(relation.to_id)
print("Recent Changes:")
for r in sorted_results[:5]:
print(f"\\n{r.path_id}")
print(f"Updated: {r.metadata['updated_at']}")
if 'author' in r.metadata:
print(f"Author: {r.metadata['author']}")
print(f"Score: {r.score:.2f}")
"""
},
{
"name": "Entity Context Loading",
"description": "Search for entities and load their full context",
"code": """
# Find relevant components
results = await search(
text="knowledge graph",
types=[SearchItemType.ENTITY],
entity_types=["component"]
)
# Load complete context
if related_ids:
if results:
# Load full entity details
path_ids = [r.path_id for r in results]
context = await open_nodes(
request=OpenNodesRequest(
path_ids=list(related_ids)
)
request=OpenNodesRequest(path_ids=path_ids)
)
# Analyze implementation status
components = [e for e in context.entities
if e.entity_type == "component"]
tests = [e for e in context.entities
if e.entity_type == "test"]
specs = [e for e in context.entities
if e.entity_type == "specification"]
print("Implementation Status:")
print(f"- Components: {len(components)}")
print(f"- Tests: {len(tests)}")
print(f"- Specs: {len(specs)}")"""
},
{
"name": "Design Analysis",
"description": "Extract architectural decisions and patterns",
"code": """
# Search for design decisions
design_results = await search_nodes(
request=SearchNodesRequest(
query="architecture pattern",
category="design"
)
)
# Group decisions by component
from collections import defaultdict
decisions = defaultdict(list)
for entity in design_results.matches:
# Extract design observations
design_notes = [o for o in entity.observations
if o.category == "design"]
if design_notes:
decisions[entity.name].extend(design_notes)
# Show architectural decisions
for component, notes in decisions.items():
print(f"\\n{component} Architecture:")
for note in notes:
context = note.context or "Design Decision"
print(f"\\n{context}:")
print(f"- {note.content}")"""
},
{
"name": "Knowledge Chain",
"description": "Follow knowledge links to build deep context",
"code": """
# Start with initial concept
initial = await search_nodes(
request=SearchNodesRequest(query="semantic web")
)
# Build knowledge chain
seen_ids = set()
to_explore = set()
# Add initial matches
for entity in initial.matches:
seen_ids.add(entity.path_id)
for relation in entity.relations:
to_explore.add(relation.to_id)
# Explore up to 2 levels deep
knowledge_chain = initial.matches
for _ in range(2):
if not to_explore:
break
# Analyze implementation details
print("Implementation Components:")
for entity in context.entities:
print(f"\\n{entity.name}")
# Load next level
next_ids = list(to_explore - seen_ids)
if next_ids:
next_level = await open_nodes(
request=OpenNodesRequest(path_ids=next_ids)
)
# Update tracking
knowledge_chain.extend(next_level.entities)
seen_ids.update(next_ids)
to_explore.clear()
# Add new relations
for entity in next_level.entities:
for relation in entity.relations:
to_explore.add(relation.to_id)
# Analyze knowledge structure
print(f"Knowledge chain depth: {len(seen_ids)} entities")
type_counts = defaultdict(int)
for entity in knowledge_chain:
type_counts[entity.entity_type] += 1
print("\\nKnowledge composition:")
for type_, count in type_counts.items():
print(f"- {type_}: {count} entities")"""
# Show technical details
tech_notes = [
o.content for o in entity.observations
if o.category == 'tech'
]
if tech_notes:
print("Technical Notes:")
for note in tech_notes:
print(f"- {note}")
# Show dependencies
deps = [r for r in entity.relations if r.relation_type == 'depends_on']
if deps:
print("\\nDependencies:")
for dep in deps:
print(f"- {dep.to_id}")
"""
}
],
output_model=SearchNodesResponse,
]
)
async def search_nodes(request: SearchNodesRequest) -> SearchNodesResponse:
"""Search for entities in the knowledge graph.
async def search(
text: str,
types: Optional[List[SearchItemType]] = None,
entity_types: Optional[List[str]] = None,
after_date: Optional[datetime] = None
) -> List[SearchResult]:
"""Search across all content in basic-memory.
Args:
request: Search parameters including query text and optional category
text: Text to search for
types: Optional list of types to filter by (DOCUMENT, ENTITY)
entity_types: Optional list of entity types to filter by
after_date: Optional date to filter results after
Returns:
SearchNodesResponse containing matching entities and search metadata
List of SearchResult objects sorted by relevance
"""
url = "/knowledge/search"
response = await client.post(url, json=request.model_dump())
return SearchNodesResponse.model_validate(response.json())
query = SearchQuery(
text=text,
types=types,
entity_types=entity_types,
after_date=after_date
)
response = await client.post("/search/", json=query.model_dump())
return [SearchResult.model_validate(r) for r in response.json()]
@mcp.tool(
@@ -188,100 +153,39 @@ async def search_nodes(request: SearchNodesRequest) -> SearchNodesResponse:
description="Load multiple entities by their path_ids in a single request",
examples=[
{
"name": "Implementation Chain",
"description": "Load and analyze implementation dependencies",
"name": "Load Search Context",
"description": "Load full entity details from search results",
"code": """
# Load feature implementation chain
chain = await open_nodes(
request=OpenNodesRequest(
path_ids=[
"feature/semantic_search", # The feature
"component/search_service", # Core implementation
"component/index_service", # Supporting service
"test/search_integration", # Integration tests
"document/search_spec" # Documentation
]
)
# First search for entities
results = await search(
text="database implementation",
types=[SearchItemType.ENTITY]
)
def analyze_dependencies(entities):
deps = defaultdict(list)
for entity in entities:
# Direct dependencies
direct = [r.to_id for r in entity.relations
if r.relation_type == "depends_on"]
deps[entity.path_id].extend(direct)
# Implicit dependencies via observations
for obs in entity.observations:
if "requires" in obs.content.lower():
deps[entity.path_id].append(
f"Implicit: {obs.content}"
)
return deps
# Show implementation structure
deps = analyze_dependencies(chain.entities)
for path_id, dependencies in deps.items():
print(f"\\n{path_id} dependencies:")
for dep in dependencies:
print(f"- {dep}")"""
},
{
"name": "Technical Analysis",
"description": "Deep dive into technical implementation",
"code": """
# First find technical components
tech_results = await search_nodes(
request=SearchNodesRequest(
query="search implementation",
category="tech"
)
)
# Load full technical context
tech_ids = [e.path_id for e in tech_results.matches
if e.entity_type == "component"]
if tech_ids:
details = await open_nodes(
request=OpenNodesRequest(path_ids=tech_ids)
# Then load full context
if results:
path_ids = [r.path_id for r in results]
context = await open_nodes(
request=OpenNodesRequest(path_ids=path_ids)
)
# Analyze technical architecture
print("Technical Architecture:\\n")
for entity in details.entities:
print(f"{entity.name}:")
# Group by entity type
by_type = defaultdict(list)
for entity in context.entities:
by_type[entity.entity_type].append(entity)
# Core capabilities
tech_notes = [o.content for o in entity.observations
if o.category == "tech"]
if tech_notes:
print("\\nCapabilities:")
for note in tech_notes:
print(f"- {note}")
# Design decisions
design_notes = [o.content for o in entity.observations
if o.category == "design"]
if design_notes:
print("\\nDesign Decisions:")
for note in design_notes:
print(f"- {note}")
# Dependencies
deps = [r for r in entity.relations
if r.relation_type == "depends_on"]
if deps:
print("\\nDependencies:")
for dep in deps:
print(f"- {dep.to_id}")
print("\\n---")"""
# Show breakdown
for etype, entities in by_type.items():
print(f"\\n{etype.title()} Components:")
for entity in entities:
print(f"- {entity.name}")
if entity.observations:
print(f" {len(entity.observations)} observations")
if entity.relations:
print(f" {len(entity.relations)} relations")
"""
}
],
output_model=EntityListResponse,
]
)
async def open_nodes(request: OpenNodesRequest) -> EntityListResponse:
"""Load multiple entities by their path_ids.
@@ -54,7 +54,7 @@ class SearchRepository():
# Handle date filter
if query.after_date:
params["after_date"] = query.after_date.isoformat()
params["after_date"] = query.after_date
conditions.append(
"json_extract(metadata, '$.created_at') > :after_date"
)
+19 -4
View File
@@ -1,9 +1,9 @@
"""Search schemas for Basic Memory."""
from typing import Optional, List
from typing import Optional, List, Union
from datetime import datetime
from enum import Enum
from pydantic import BaseModel
from pydantic import BaseModel, field_validator
class SearchItemType(str, Enum):
@@ -17,7 +17,17 @@ class SearchQuery(BaseModel):
text: str
types: Optional[List[SearchItemType]] = None
entity_types: Optional[List[str]] = None
after_date: Optional[datetime] = None
after_date: Optional[Union[datetime, str]] = None
@field_validator('after_date')
@classmethod
def validate_date(cls, v: Optional[Union[datetime, str]]) -> Optional[str]:
"""Convert datetime to ISO format if needed."""
if v is None:
return None
if isinstance(v, datetime):
return v.isoformat()
return v # Assume it's already a string
class SearchResult(BaseModel):
@@ -26,4 +36,9 @@ class SearchResult(BaseModel):
file_path: str
type: SearchItemType
score: float
metadata: dict
metadata: dict
class SearchResponse(BaseModel):
"""Wrapper for search results list."""
results: List[SearchResult]
+14 -11
View File
@@ -6,7 +6,7 @@ import pytest
from httpx import AsyncClient
from basic_memory.config import ProjectConfig
from basic_memory.schemas.search import SearchItemType
from basic_memory.schemas.search import SearchItemType, SearchResponse
@pytest.mark.asyncio
@@ -28,10 +28,10 @@ async def test_document_indexing(client: AsyncClient, test_config):
json={"text": "unique searchable content", "types": [SearchItemType.DOCUMENT.value]},
)
assert search_response.status_code == 200
results = search_response.json()
assert len(results) == 1
assert results[0]["path_id"] == "test.md"
assert results[0]["type"] == SearchItemType.DOCUMENT.value
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 1
assert search_result.results[0].path_id == "test.md"
assert search_result.results[0].type == SearchItemType.DOCUMENT.value
@pytest.mark.asyncio
@@ -59,15 +59,16 @@ async def test_document_update_indexing(client: AsyncClient):
search_response = await client.post(
"/search/", json={"text": "sphinx", "types": [SearchItemType.DOCUMENT.value]}
)
results = search_response.json()
assert len(results) == 1
assert results[0]["path_id"] == "test.md"
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 1
assert search_result.results[0].path_id == "test.md"
# Original terms shouldn't be found
search_response = await client.post(
"/search/", json={"text": "without special", "types": [SearchItemType.DOCUMENT.value]}
)
assert len(search_response.json()) == 0
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 0
@pytest.mark.asyncio
@@ -86,7 +87,8 @@ async def test_document_delete_indexing(client: AsyncClient):
search_response = await client.post(
"/search/", json={"text": "should disappear", "types": [SearchItemType.DOCUMENT.value]}
)
assert len(search_response.json()) == 1
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 1
# Delete document
delete_response = await client.delete(f"/documents/{test_doc["path_id"]}")
@@ -96,7 +98,8 @@ async def test_document_delete_indexing(client: AsyncClient):
search_response = await client.post(
"/search/", json={"text": "should disappear", "types": [SearchItemType.DOCUMENT.value]}
)
assert len(search_response.json()) == 0
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 0
@pytest.mark.asyncio
+15 -13
View File
@@ -12,7 +12,7 @@ from basic_memory.schemas import (
ObservationResponse,
RelationResponse,
)
from basic_memory.schemas.search import SearchItemType
from basic_memory.schemas.search import SearchItemType, SearchResponse
async def create_entity(client) -> EntityResponse:
@@ -470,10 +470,10 @@ async def test_entity_indexing(client: AsyncClient):
"/search/", json={"text": "unique searchable", "types": [SearchItemType.ENTITY.value]}
)
assert search_response.status_code == 200
results = search_response.json()
assert len(results) == 1
assert results[0]["path_id"] == "test/search_test"
assert results[0]["type"] == SearchItemType.ENTITY.value
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 1
assert search_result.results[0].path_id == "test/search_test"
assert search_result.results[0].type == SearchItemType.ENTITY.value
@pytest.mark.asyncio
@@ -502,9 +502,9 @@ async def test_observation_update_indexing(client: AsyncClient):
search_response = await client.post(
"/search/", json={"text": "sphinx", "types": [SearchItemType.ENTITY.value]}
)
results = search_response.json()
assert len(results) == 1
assert results[0]["path_id"] == entity["path_id"]
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 1
assert search_result.results[0].path_id == entity["path_id"]
@pytest.mark.asyncio
@@ -525,7 +525,8 @@ async def test_entity_delete_indexing(client: AsyncClient):
search_response = await client.post(
"/search/", json={"text": "should be removed", "types": [SearchItemType.ENTITY.value]}
)
assert len(search_response.json()) == 1
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 1
# Delete entity
delete_response = await client.post(
@@ -537,7 +538,8 @@ async def test_entity_delete_indexing(client: AsyncClient):
search_response = await client.post(
"/search/", json={"text": "should be removed", "types": [SearchItemType.ENTITY.value]}
)
assert len(search_response.json()) == 0
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 0
@pytest.mark.asyncio
@@ -571,7 +573,7 @@ async def test_relation_indexing(client: AsyncClient):
search_response = await client.post(
"/search/", json={"text": "sphinx relation", "types": [SearchItemType.ENTITY.value]}
)
results = search_response.json()
assert len(results) == 2 # Both source and target entities
path_ids = {r["path_id"] for r in results}
search_result = SearchResponse.model_validate(search_response.json())
assert len(search_result.results) == 2 # Both source and target entities
path_ids = {r.path_id for r in search_result.results}
assert path_ids == {"test/source_test", "test/target_test"}
+33 -24
View File
@@ -6,7 +6,7 @@ import pytest
import pytest_asyncio
from sqlalchemy import text
from basic_memory import db
from basic_memory.schemas.search import SearchQuery, SearchItemType
from basic_memory.schemas.search import SearchQuery, SearchItemType, SearchResponse
@pytest.fixture
@@ -67,9 +67,9 @@ async def test_search_basic(client, indexed_entity):
}
)
assert response.status_code == 200
results = response.json()
assert len(results) == 1
assert results[0]["path_id"] == indexed_entity.path_id
search_results = SearchResponse.model_validate(response.json())
assert len(search_results.results) == 1
assert search_results.results[0].path_id == indexed_entity.path_id
@pytest.mark.asyncio
@@ -84,8 +84,8 @@ async def test_search_with_type_filter(client, indexed_entity):
}
)
assert response.status_code == 200
results = response.json()
assert len(results) == 1
search_results = SearchResponse.model_validate(response.json())
assert len(search_results.results) == 1
# Should not find with wrong type
response = await client.post(
@@ -96,7 +96,8 @@ async def test_search_with_type_filter(client, indexed_entity):
}
)
assert response.status_code == 200
assert len(response.json()) == 0
search_results = SearchResponse.model_validate(response.json())
assert len(search_results.results) == 0
@pytest.mark.asyncio
@@ -111,8 +112,8 @@ async def test_search_with_entity_type_filter(client, indexed_entity):
}
)
assert response.status_code == 200
results = response.json()
assert len(results) == 1
search_results = SearchResponse.model_validate(response.json())
assert len(search_results.results) == 1
# Should not find with wrong entity type
response = await client.post(
@@ -123,7 +124,8 @@ async def test_search_with_entity_type_filter(client, indexed_entity):
}
)
assert response.status_code == 200
assert len(response.json()) == 0
search_results = SearchResponse.model_validate(response.json())
assert len(search_results.results) == 0
@pytest.mark.asyncio
@@ -139,7 +141,8 @@ async def test_search_with_date_filter(client, indexed_entity):
}
)
assert response.status_code == 200
assert len(response.json()) == 1
search_results = SearchResponse.model_validate(response.json())
assert len(search_results.results) == 1
# Should not find with future date
future_date = datetime(2030, 1, 1, tzinfo=timezone.utc)
@@ -151,7 +154,8 @@ async def test_search_with_date_filter(client, indexed_entity):
}
)
assert response.status_code == 200
assert len(response.json()) == 0
search_results = SearchResponse.model_validate(response.json())
assert len(search_results.results) == 0
@pytest.mark.asyncio
@@ -171,9 +175,12 @@ async def test_search_scoring(client, indexed_entity):
assert exact_response.status_code == 200
assert partial_response.status_code == 200
exact_result = SearchResponse.model_validate(exact_response.json())
partial_result = SearchResponse.model_validate(partial_response.json())
exact_score = exact_response.json()[0]["score"]
partial_score = partial_response.json()[0]["score"]
exact_score = exact_result.results[0].score
partial_score = partial_result.results[0].score
assert exact_score > partial_score
@@ -186,7 +193,8 @@ async def test_search_empty(search_service, client):
json={"text": "nonexistent"}
)
assert response.status_code == 200
assert len(response.json()) == 0
search_result = SearchResponse.model_validate(response.json())
assert len(search_result.results) == 0
@pytest.mark.asyncio
@@ -218,7 +226,8 @@ async def test_reindex(
"/search/",
json={"text": "test"}
)
assert len(response.json()) == 0
search_results = SearchResponse.model_validate(response.json())
assert len(search_results.results) == 0
# Trigger reindex
reindex_response = await client.post("/search/reindex")
@@ -230,8 +239,8 @@ async def test_reindex(
"/search/",
json={"text": "test"}
)
results = search_response.json()
assert len(results) == 2 # Both entity and document should be found
search_results = SearchResponse.model_validate(search_response.json())
assert len(search_results.results) == 2 # Both entity and document should be found
@pytest.mark.asyncio
@@ -247,9 +256,9 @@ async def test_multiple_filters(client, indexed_entity):
}
)
assert response.status_code == 200
results = response.json()
assert len(results) == 1
result = results[0]
assert result["path_id"] == indexed_entity.path_id
assert result["type"] == SearchItemType.ENTITY.value
assert result["metadata"]["entity_type"] == "component"
search_result = SearchResponse.model_validate(response.json())
assert len(search_result.results) == 1
result = search_result.results[0]
assert result.path_id == indexed_entity.path_id
assert result.type == SearchItemType.ENTITY.value
assert result.metadata["entity_type"] == "component"
+2
View File
@@ -203,6 +203,7 @@ async def sync_service(
knowledge_sync_service: KnowledgeSyncService,
file_change_scanner: FileChangeScanner,
knowledge_parser: KnowledgeParser,
search_service: SearchService
) -> SyncService:
"""Create sync service for testing."""
return SyncService(
@@ -210,6 +211,7 @@ async def sync_service(
document_service=document_service,
knowledge_sync_service=knowledge_sync_service,
knowledge_parser=knowledge_parser,
search_service=search_service,
)
-173
View File
@@ -1,173 +0,0 @@
"""Tests for search_nodes MCP tool."""
import pytest
from basic_memory.mcp.tools.search import search_nodes
from basic_memory.mcp.tools.knowledge import create_entities
from basic_memory.schemas.base import Entity, ObservationCategory
from basic_memory.schemas.request import CreateEntityRequest, SearchNodesRequest, ObservationCreate
@pytest.mark.asyncio
async def test_basic_search(client):
"""Test basic text search."""
# Create some test entities
entity_request = CreateEntityRequest(
entities=[
Entity(
name="SearchComponent",
entity_type="component",
description="A searchable component",
observations=["This has some searchable text"]
),
Entity(
name="OtherComponent",
entity_type="component",
description="Another component",
observations=["This is unrelated"]
)
]
)
await create_entities(entity_request)
# Search for "searchable"
request = SearchNodesRequest(query="searchable")
result = await search_nodes(request)
# Should find one matching entity
assert len(result.matches) == 1
assert result.matches[0].name == "SearchComponent"
assert result.query == "searchable"
@pytest.mark.asyncio
async def test_search_with_category(client):
"""Test search with category filter."""
# Create an entity with different observation categories
obs_tech = ObservationCreate(
content="Technical detail about implementation",
category=ObservationCategory.TECH
)
obs_design = ObservationCreate(
content="Design decision about architecture",
category=ObservationCategory.DESIGN
)
entity_request = CreateEntityRequest(
entities=[
Entity(
name="TestEntity",
entity_type="test",
description="Test entity",
observations=[obs_tech.content, obs_design.content]
)
]
)
await create_entities(entity_request)
# Search for tech observations only
request = SearchNodesRequest(
query="implementation",
category=ObservationCategory.TECH
)
tech_result = await search_nodes(request)
assert len(tech_result.matches) == 1
# Search for design observations only
request = SearchNodesRequest(
query="architecture",
category=ObservationCategory.DESIGN
)
design_result = await search_nodes(request)
assert len(design_result.matches) == 1
@pytest.mark.asyncio
async def test_search_multiple_matches(client):
"""Test search returning multiple entities."""
# Create multiple entities with similar content
entity_request = CreateEntityRequest(
entities=[
Entity(
name="Component1",
entity_type="component",
description="Uses SQLite database",
observations=["Implements SQLite storage"]
),
Entity(
name="Component2",
entity_type="component",
description="Another SQLite component",
observations=["Also uses SQLite"]
)
]
)
await create_entities(entity_request)
# Search for SQLite
request = SearchNodesRequest(query="SQLite")
result = await search_nodes(request)
# Should find both entities
assert len(result.matches) == 2
names = {e.name for e in result.matches}
assert "Component1" in names
assert "Component2" in names
@pytest.mark.asyncio
async def test_search_no_matches(client):
"""Test search with no matching results."""
# Create an entity with unrelated content
entity_request = CreateEntityRequest(
entities=[
Entity(
name="UnrelatedEntity",
entity_type="test",
description="Something unrelated",
observations=["Nothing to see here"]
)
]
)
await create_entities(entity_request)
# Search for non-matching term
request = SearchNodesRequest(query="nonexistent")
result = await search_nodes(request)
# Should find no matches
assert len(result.matches) == 0
assert result.query == "nonexistent"
@pytest.mark.asyncio
async def test_search_case_insensitive(client):
"""Test that search is case insensitive."""
# Create entity with mixed case text
entity_request = CreateEntityRequest(
entities=[
Entity(
name="MixedCase",
entity_type="test",
description="Testing MIXED case text",
observations=["Some MiXeD cAsE content"]
)
]
)
await create_entities(entity_request)
# Search with different cases
lower_request = SearchNodesRequest(query="mixed")
upper_request = SearchNodesRequest(query="MIXED")
mixed_request = SearchNodesRequest(query="MiXeD")
# All should find the entity
lower_result = await search_nodes(lower_request)
upper_result = await search_nodes(upper_request)
mixed_result = await search_nodes(mixed_request)
assert len(lower_result.matches) == 1
assert len(upper_result.matches) == 1
assert len(mixed_result.matches) == 1
assert all(r.matches[0].name == "MixedCase"
for r in [lower_result, upper_result, mixed_result])