Compare commits

...
Author SHA1 Message Date
Nicolò Boschi 34bb86e1fe support tags 2026-01-13 18:18:01 +01:00
Nicolò Boschi b18b91588c support tags 2026-01-13 18:10:48 +01:00
Nicolò Boschi 7afbd1ef6c feat: add memory tags 2026-01-13 15:42:54 +01:00
Nicolò Boschi 48b19f5543 feat: add memory tags 2026-01-13 15:30:13 +01:00
51 changed files with 3969 additions and 354 deletions
+19
View File
@@ -174,6 +174,25 @@ This runs the same checks as the pre-commit hook (Ruff for Python, ESLint/Pretti
- Multi-bank queries are client responsibility to orchestrate
- Disposition traits only affect reflect, not recall
### Control Plane API Routes
When adding or modifying parameters in the dataplane API (hindsight-api), you must also update the control plane routes that proxy to it:
1. **API Routes** (`hindsight-control-plane/src/app/api/`):
- `recall/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/recall`
- `reflect/route.ts` - proxies to `/v1/default/banks/{bank_id}/reflect`
- `memories/retain/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/retain`
- Other routes follow the same pattern
2. **Client types** (`hindsight-control-plane/src/lib/api.ts`):
- Update the TypeScript type definitions for `recall()`, `reflect()`, `retain()` etc.
3. **Checklist when adding new API parameters**:
- Add parameter extraction in the route handler (destructure from `body`)
- Pass the parameter to the SDK call
- Update the client type definition in `lib/api.ts`
- Update any UI components that need to use the new parameter
### Python Style
- Python 3.11+, type hints required
- Async throughout (asyncpg, async FastAPI)
@@ -0,0 +1,48 @@
"""add_tags_column
Revision ID: g2a3b4c5d6e7
Revises: f1a2b3c4d5e6
Create Date: 2025-01-13
Add tags column to memory_units and documents tables for visibility scoping.
Tags enable filtering memories by scope (e.g., user IDs, session IDs) during recall/reflect.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "g2a3b4c5d6e7"
down_revision: str | Sequence[str] | None = "f1a2b3c4d5e6"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
"""Get schema prefix for table names (e.g., 'tenant_x.' or '' for public)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
"""Add tags column to memory_units and documents tables."""
schema = _get_schema_prefix()
# Add tags column to memory_units table
op.execute(f"ALTER TABLE {schema}memory_units ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
# Create GIN index for efficient array containment queries (tags && ARRAY['x'])
op.execute(f"CREATE INDEX IF NOT EXISTS idx_memory_units_tags ON {schema}memory_units USING GIN (tags)")
# Add tags column to documents table for document-level tags
op.execute(f"ALTER TABLE {schema}documents ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
def downgrade() -> None:
"""Remove tags columns and index."""
schema = _get_schema_prefix()
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_tags")
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS tags")
op.execute(f"ALTER TABLE {schema}documents DROP COLUMN IF EXISTS tags")
+167 -2
View File
@@ -37,6 +37,7 @@ from hindsight_api import MemoryEngine
from hindsight_api.engine.db_utils import acquire_with_retry
from hindsight_api.engine.memory_engine import Budget, fq_table
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES, TokenUsage
from hindsight_api.engine.search.tags import TagsMatch
from hindsight_api.extensions import HttpExtension, OperationValidationError, load_extension
from hindsight_api.metrics import create_metrics_collector, get_metrics_collector, initialize_metrics
from hindsight_api.models import RequestContext
@@ -81,6 +82,8 @@ class RecallRequest(BaseModel):
"trace": True,
"query_timestamp": "2023-05-30T23:40:00",
"include": {"entities": {"max_tokens": 500}},
"tags": ["user_a"],
"tags_match": "any",
}
}
)
@@ -99,6 +102,15 @@ class RecallRequest(BaseModel):
default_factory=IncludeOptions,
description="Options for including additional data (entities are included by default)",
)
tags: list[str] | None = Field(
default=None,
description="Filter memories by tags. If not specified, all memories are returned.",
)
tags_match: TagsMatch = Field(
default="any",
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
)
class RecallResult(BaseModel):
@@ -119,6 +131,7 @@ class RecallResult(BaseModel):
"document_id": "session_abc123",
"metadata": {"source": "slack"},
"chunk_id": "456e7890-e12b-34d5-a678-901234567890",
"tags": ["user_a", "user_b"],
}
},
}
@@ -134,6 +147,7 @@ class RecallResult(BaseModel):
document_id: str | None = None # Document this memory belongs to
metadata: dict[str, str] | None = None # User-defined metadata
chunk_id: str | None = None # Chunk this fact was extracted from
tags: list[str] | None = None # Visibility scope tags
class EntityObservationResponse(BaseModel):
@@ -306,6 +320,7 @@ class MemoryItem(BaseModel):
"metadata": {"source": "slack", "channel": "engineering"},
"document_id": "meeting_notes_2024_01_15",
"entities": [{"text": "Alice"}, {"text": "ML model", "type": "CONCEPT"}],
"tags": ["user_a", "user_b"],
}
},
)
@@ -319,6 +334,10 @@ class MemoryItem(BaseModel):
default=None,
description="Optional entities to combine with auto-extracted entities.",
)
tags: list[str] | None = Field(
default=None,
description="Optional tags for visibility scoping. Memories with tags can be filtered during recall.",
)
@field_validator("timestamp", mode="before")
@classmethod
@@ -353,6 +372,7 @@ class RetainRequest(BaseModel):
},
],
"async": False,
"document_tags": ["user_a", "user_b"],
}
}
)
@@ -363,6 +383,10 @@ class RetainRequest(BaseModel):
alias="async",
description="If true, process asynchronously in background. If false, wait for completion (default: false)",
)
document_tags: list[str] | None = Field(
default=None,
description="Tags applied to all items in this request. These are merged with any item-level tags.",
)
class RetainResponse(BaseModel):
@@ -431,6 +455,8 @@ class ReflectRequest(BaseModel):
},
"required": ["summary", "key_points"],
},
"tags": ["user_a"],
"tags_match": "any",
}
}
)
@@ -446,6 +472,15 @@ class ReflectRequest(BaseModel):
default=None,
description="Optional JSON Schema for structured output. When provided, the response will include a 'structured_output' field with the LLM response parsed according to this schema.",
)
tags: list[str] | None = Field(
default=None,
description="Filter memories by tags during reflection. If not specified, all memories are considered.",
)
tags_match: TagsMatch = Field(
default="any",
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
)
class OpinionItem(BaseModel):
@@ -728,6 +763,37 @@ class ListDocumentsResponse(BaseModel):
offset: int
class TagItem(BaseModel):
"""Single tag with usage count."""
tag: str = Field(description="The tag value")
count: int = Field(description="Number of memories with this tag")
class ListTagsResponse(BaseModel):
"""Response model for list tags endpoint."""
model_config = ConfigDict(
json_schema_extra={
"example": {
"items": [
{"tag": "user:alice", "count": 42},
{"tag": "user:bob", "count": 15},
{"tag": "session:abc123", "count": 8},
],
"total": 25,
"limit": 100,
"offset": 0,
}
}
)
items: list[TagItem]
total: int
limit: int
offset: int
class DocumentResponse(BaseModel):
"""Response model for get document endpoint."""
@@ -741,6 +807,7 @@ class DocumentResponse(BaseModel):
"created_at": "2024-01-15T10:30:00Z",
"updated_at": "2024-01-15T10:30:00Z",
"memory_unit_count": 15,
"tags": ["user_a", "session_123"],
}
}
)
@@ -752,6 +819,7 @@ class DocumentResponse(BaseModel):
created_at: str
updated_at: str
memory_unit_count: int
tags: list[str] = Field(default_factory=list, description="Tags associated with this document")
class DeleteDocumentResponse(BaseModel):
@@ -1179,6 +1247,37 @@ def _register_routes(app: FastAPI):
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/list: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.get(
"/v1/default/banks/{bank_id}/memories/{memory_id}",
summary="Get memory unit",
description="Get a single memory unit by ID with all its metadata including entities and tags.",
operation_id="get_memory",
tags=["Memory"],
)
async def api_get_memory(
bank_id: str,
memory_id: str,
request_context: RequestContext = Depends(get_request_context),
):
"""Get a single memory unit by ID."""
try:
data = await app.state.memory.get_memory_unit(
bank_id=bank_id,
memory_id=memory_id,
request_context=request_context,
)
if data is None:
raise HTTPException(status_code=404, detail=f"Memory unit '{memory_id}' not found")
return data
except (AuthenticationError, HTTPException):
raise
except Exception as e:
import traceback
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/{memory_id}: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.post(
"/v1/default/banks/{bank_id}/memories/recall",
response_model=RecallResponse,
@@ -1243,6 +1342,8 @@ def _register_routes(app: FastAPI):
include_chunks=include_chunks,
max_chunk_tokens=max_chunk_tokens,
request_context=request_context,
tags=request.tags,
tags_match=request.tags_match,
)
# Convert core MemoryFact objects to API RecallResult objects (excluding internal metrics)
@@ -1258,6 +1359,7 @@ def _register_routes(app: FastAPI):
mentioned_at=fact.mentioned_at,
document_id=fact.document_id,
chunk_id=fact.chunk_id,
tags=fact.tags,
)
for fact in core_result.results
]
@@ -1350,6 +1452,8 @@ def _register_routes(app: FastAPI):
max_tokens=request.max_tokens,
response_schema=request.response_schema,
request_context=request_context,
tags=request.tags,
tags_match=request.tags_match,
)
# Convert core MemoryFact objects to API ReflectFact objects if facts are requested
@@ -1734,6 +1838,59 @@ def _register_routes(app: FastAPI):
logger.error(f"Error in /v1/default/banks/{bank_id}/documents/{document_id}: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.get(
"/v1/default/banks/{bank_id}/tags",
response_model=ListTagsResponse,
summary="List tags",
description="List all unique tags in a memory bank with usage counts. "
"Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.",
operation_id="list_tags",
tags=["Memory"],
)
async def api_list_tags(
bank_id: str,
q: str | None = Query(
default=None,
description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). "
"Use '*' as wildcard. Case-insensitive.",
),
limit: int = Query(default=100, description="Maximum number of tags to return"),
offset: int = Query(default=0, description="Offset for pagination"),
request_context: RequestContext = Depends(get_request_context),
):
"""
List all unique tags in a memory bank.
Use this endpoint to discover available tags or expand wildcard patterns.
Supports '*' wildcards for flexible matching (case-insensitive):
- 'user:*' matches user:alice, user:bob
- '*-admin' matches role-admin, super-admin
- 'env*-prod' matches env-prod, environment-prod
Args:
bank_id: Memory Bank ID (from path)
q: Wildcard pattern to filter tags (use '*' as wildcard)
limit: Maximum number of tags to return (default: 100)
offset: Offset for pagination (default: 0)
"""
try:
data = await app.state.memory.list_tags(
bank_id=bank_id,
pattern=q,
limit=limit,
offset=offset,
request_context=request_context,
)
return data
except (AuthenticationError, HTTPException):
raise
except Exception as e:
import traceback
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
logger.error(f"Error in /v1/default/banks/{bank_id}/tags: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.get(
"/v1/default/chunks/{chunk_id:path}",
response_model=ChunkResponse,
@@ -2096,11 +2253,15 @@ def _register_routes(app: FastAPI):
content_dict["document_id"] = item.document_id
if item.entities:
content_dict["entities"] = [{"text": e.text, "type": e.type or "CONCEPT"} for e in item.entities]
if item.tags:
content_dict["tags"] = item.tags
contents.append(content_dict)
if request.async_:
# Async processing: queue task and return immediately
result = await app.state.memory.submit_async_retain(bank_id, contents, request_context=request_context)
result = await app.state.memory.submit_async_retain(
bank_id, contents, document_tags=request.document_tags, request_context=request_context
)
return RetainResponse.model_validate(
{
"success": True,
@@ -2114,7 +2275,11 @@ def _register_routes(app: FastAPI):
# Synchronous processing: wait for completion (record metrics)
with metrics.record_operation("retain", bank_id=bank_id, source="api"):
result, usage = await app.state.memory.retain_batch_async(
bank_id=bank_id, contents=contents, request_context=request_context, return_usage=True
bank_id=bank_id,
contents=contents,
document_tags=request.document_tags,
request_context=request_context,
return_usage=True,
)
return RetainResponse.model_validate(
@@ -151,6 +151,7 @@ from .retain import bank_utils, embedding_utils
from .retain.types import RetainContentDict
from .search import observation_utils, think_utils
from .search.reranking import CrossEncoderReranker
from .search.tags import TagsMatch
from .task_backend import AsyncIOQueueBackend, NoopTaskBackend, TaskBackend
@@ -1059,6 +1060,7 @@ class MemoryEngine(MemoryEngineInterface):
document_id: str | None = None,
fact_type_override: str | None = None,
confidence_score: float | None = None,
document_tags: list[str] | None = None,
return_usage: bool = False,
):
"""
@@ -1191,6 +1193,7 @@ class MemoryEngine(MemoryEngineInterface):
is_first_batch=i == 1, # Only upsert on first batch
fact_type_override=fact_type_override,
confidence_score=confidence_score,
document_tags=document_tags,
)
all_results.extend(sub_results)
total_usage = total_usage + sub_usage
@@ -1209,6 +1212,7 @@ class MemoryEngine(MemoryEngineInterface):
is_first_batch=True,
fact_type_override=fact_type_override,
confidence_score=confidence_score,
document_tags=document_tags,
)
# Call post-operation hook if validator is configured
@@ -1243,6 +1247,7 @@ class MemoryEngine(MemoryEngineInterface):
is_first_batch: bool = True,
fact_type_override: str | None = None,
confidence_score: float | None = None,
document_tags: list[str] | None = None,
) -> tuple[list[list[str]], "TokenUsage"]:
"""
Internal method for batch processing without chunking logic.
@@ -1259,6 +1264,7 @@ class MemoryEngine(MemoryEngineInterface):
is_first_batch: Whether this is the first batch (for chunked operations, only delete on first batch)
fact_type_override: Override fact type for all facts
confidence_score: Confidence score for opinions
document_tags: Tags applied to all items in this batch
Returns:
Tuple of (unit ID lists, token usage for fact extraction)
@@ -1283,6 +1289,7 @@ class MemoryEngine(MemoryEngineInterface):
is_first_batch=is_first_batch,
fact_type_override=fact_type_override,
confidence_score=confidence_score,
document_tags=document_tags,
)
def recall(
@@ -1341,6 +1348,8 @@ class MemoryEngine(MemoryEngineInterface):
include_chunks: bool = False,
max_chunk_tokens: int = 8192,
request_context: "RequestContext",
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> RecallResultModel:
"""
Recall memories using N*4-way parallel retrieval (N fact types × 4 retrieval methods).
@@ -1366,6 +1375,8 @@ class MemoryEngine(MemoryEngineInterface):
max_entity_tokens: Maximum tokens for entity observations (default 500)
include_chunks: Whether to include raw chunks in the response
max_chunk_tokens: Maximum tokens for chunks (default 8192)
tags: Optional list of tags for visibility filtering (OR matching - returns
memories that have at least one matching tag)
Returns:
RecallResultModel containing:
@@ -1438,6 +1449,8 @@ class MemoryEngine(MemoryEngineInterface):
max_chunk_tokens,
request_context,
semaphore_wait=semaphore_wait,
tags=tags,
tags_match=tags_match,
)
break # Success - exit retry loop
except Exception as e:
@@ -1556,6 +1569,8 @@ class MemoryEngine(MemoryEngineInterface):
max_chunk_tokens: int = 8192,
request_context: "RequestContext" = None,
semaphore_wait: float = 0.0,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> RecallResultModel:
"""
Search implementation with modular retrieval and reranking.
@@ -1585,7 +1600,9 @@ class MemoryEngine(MemoryEngineInterface):
# Initialize tracer if requested
from .search.tracer import SearchTracer
tracer = SearchTracer(query, thinking_budget, max_tokens) if enable_trace else None
tracer = (
SearchTracer(query, thinking_budget, max_tokens, tags=tags, tags_match=tags_match) if enable_trace else None
)
if tracer:
tracer.start()
@@ -1595,8 +1612,9 @@ class MemoryEngine(MemoryEngineInterface):
# Buffer logs for clean output in concurrent scenarios
recall_id = f"{bank_id[:8]}-{int(time.time() * 1000) % 100000}"
log_buffer = []
tags_info = f", tags={tags}, tags_match={tags_match}" if tags else ""
log_buffer.append(
f"[RECALL {recall_id}] Query: '{query[:50]}...' (budget={thinking_budget}, max_tokens={max_tokens})"
f"[RECALL {recall_id}] Query: '{query[:50]}...' (budget={thinking_budget}, max_tokens={max_tokens}{tags_info})"
)
try:
@@ -1642,6 +1660,8 @@ class MemoryEngine(MemoryEngineInterface):
thinking_budget,
question_date,
self.query_analyzer,
tags=tags,
tags_match=tags_match,
)
parallel_duration = time.time() - parallel_start
@@ -1746,6 +1766,11 @@ class MemoryEngine(MemoryEngineInterface):
f"edges={hd.get('edges_loaded', 0)}"
)
# Record temporal constraint in tracer if detected
if tracer and detected_temporal_constraint:
start_dt, end_dt = detected_temporal_constraint
tracer.record_temporal_constraint(start_dt, end_dt)
# Record retrieval results for tracer - per fact type
if tracer:
# Convert RetrievalResult to old tuple format for tracer
@@ -1788,14 +1813,22 @@ class MemoryEngine(MemoryEngineInterface):
fact_type=ft_name,
)
# Add temporal retrieval results for this fact type (even if empty, to show it ran)
if rr.temporal is not None:
# Add temporal retrieval results for this fact type
# Show temporal even with 0 results if constraint was detected
if rr.temporal is not None or rr.temporal_constraint is not None:
temporal_metadata = {"budget": thinking_budget}
if rr.temporal_constraint:
start_dt, end_dt = rr.temporal_constraint
temporal_metadata["constraint"] = {
"start": start_dt.isoformat() if start_dt else None,
"end": end_dt.isoformat() if end_dt else None,
}
tracer.add_retrieval_results(
method_name="temporal",
results=to_tuple_format(rr.temporal),
results=to_tuple_format(rr.temporal or []),
duration_seconds=rr.timings.get("temporal", 0.0),
score_field="temporal_score",
metadata={"budget": thinking_budget},
metadata=temporal_metadata,
fact_type=ft_name,
)
@@ -2055,6 +2088,7 @@ class MemoryEngine(MemoryEngineInterface):
mentioned_at=result_dict.get("mentioned_at"),
document_id=result_dict.get("document_id"),
chunk_id=result_dict.get("chunk_id"),
tags=result_dict.get("tags"),
)
)
@@ -2270,11 +2304,11 @@ class MemoryEngine(MemoryEngineInterface):
doc = await conn.fetchrow(
f"""
SELECT d.id, d.bank_id, d.original_text, d.content_hash,
d.created_at, d.updated_at, COUNT(mu.id) as unit_count
d.created_at, d.updated_at, d.tags, COUNT(mu.id) as unit_count
FROM {fq_table("documents")} d
LEFT JOIN {fq_table("memory_units")} mu ON mu.document_id = d.id
WHERE d.id = $1 AND d.bank_id = $2
GROUP BY d.id, d.bank_id, d.original_text, d.content_hash, d.created_at, d.updated_at
GROUP BY d.id, d.bank_id, d.original_text, d.content_hash, d.created_at, d.updated_at, d.tags
""",
document_id,
bank_id,
@@ -2291,6 +2325,7 @@ class MemoryEngine(MemoryEngineInterface):
"memory_unit_count": doc["unit_count"],
"created_at": doc["created_at"].isoformat() if doc["created_at"] else None,
"updated_at": doc["updated_at"].isoformat() if doc["updated_at"] else None,
"tags": list(doc["tags"]) if doc["tags"] else [],
}
async def delete_document(
@@ -2779,6 +2814,68 @@ class MemoryEngine(MemoryEngineInterface):
return {"items": items, "total": total, "limit": limit, "offset": offset}
async def get_memory_unit(
self,
bank_id: str,
memory_id: str,
request_context: "RequestContext",
):
"""
Get a single memory unit by ID.
Args:
bank_id: Bank ID
memory_id: Memory unit ID
request_context: Request context for authentication.
Returns:
Dict with memory unit data or None if not found
"""
await self._authenticate_tenant(request_context)
pool = await self._get_pool()
async with acquire_with_retry(pool) as conn:
# Get the memory unit
row = await conn.fetchrow(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, fact_type, document_id, chunk_id, tags
FROM {fq_table("memory_units")}
WHERE id = $1 AND bank_id = $2
""",
memory_id,
bank_id,
)
if not row:
return None
# Get entity information
entities_rows = await conn.fetch(
f"""
SELECT e.canonical_name
FROM {fq_table("unit_entities")} ue
JOIN {fq_table("entities")} e ON ue.entity_id = e.id
WHERE ue.unit_id = $1
""",
row["id"],
)
entities = [r["canonical_name"] for r in entities_rows]
return {
"id": str(row["id"]),
"text": row["text"],
"context": row["context"] if row["context"] else "",
"date": row["event_date"].isoformat() if row["event_date"] else "",
"type": row["fact_type"],
"mentioned_at": row["mentioned_at"].isoformat() if row["mentioned_at"] else None,
"occurred_start": row["occurred_start"].isoformat() if row["occurred_start"] else None,
"occurred_end": row["occurred_end"].isoformat() if row["occurred_end"] else None,
"entities": entities,
"document_id": row["document_id"] if row["document_id"] else None,
"chunk_id": str(row["chunk_id"]) if row["chunk_id"] else None,
"tags": row["tags"] if row["tags"] else [],
}
async def list_documents(
self,
bank_id: str,
@@ -3302,6 +3399,8 @@ Guidelines:
max_tokens: int = 4096,
response_schema: dict | None = None,
request_context: "RequestContext",
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> ReflectResult:
"""
Reflect and formulate an answer using bank identity, world facts, and opinions.
@@ -3369,6 +3468,8 @@ Guidelines:
fact_type=["experience", "world", "opinion"],
include_entities=True,
request_context=request_context,
tags=tags,
tags_match=tags_match,
)
recall_time = time.time() - recall_start
@@ -3721,6 +3822,85 @@ Guidelines:
"offset": offset,
}
async def list_tags(
self,
bank_id: str,
*,
pattern: str | None = None,
limit: int = 100,
offset: int = 0,
request_context: "RequestContext",
) -> dict[str, Any]:
"""
List all unique tags for a bank with usage counts.
Use this to discover available tags or expand wildcard patterns.
Supports '*' as wildcard for flexible matching (case-insensitive):
- 'user:*' matches user:alice, user:bob
- '*-admin' matches role-admin, super-admin
- 'env*-prod' matches env-prod, environment-prod
Args:
bank_id: Bank identifier
pattern: Wildcard pattern to filter tags (use '*' as wildcard, case-insensitive)
limit: Maximum number of tags to return
offset: Offset for pagination
request_context: Request context for authentication.
Returns:
Dict with items (list of {tag, count}), total, limit, offset
"""
await self._authenticate_tenant(request_context)
pool = await self._get_pool()
async with acquire_with_retry(pool) as conn:
# Build pattern filter if provided (convert * to % for ILIKE)
pattern_clause = ""
params: list[Any] = [bank_id]
if pattern:
# Convert wildcard pattern: * -> % for SQL ILIKE
sql_pattern = pattern.replace("*", "%")
pattern_clause = "AND tag ILIKE $2"
params.append(sql_pattern)
# Get total count of distinct tags matching pattern
total_row = await conn.fetchrow(
f"""
SELECT COUNT(DISTINCT tag) as total
FROM {fq_table("memory_units")}, unnest(tags) AS tag
WHERE bank_id = $1 AND tags IS NOT NULL AND tags != '{{}}'
{pattern_clause}
""",
*params,
)
total = total_row["total"] if total_row else 0
# Get paginated tags with counts, ordered by frequency
limit_param = len(params) + 1
offset_param = len(params) + 2
params.extend([limit, offset])
rows = await conn.fetch(
f"""
SELECT tag, COUNT(*) as count
FROM {fq_table("memory_units")}, unnest(tags) AS tag
WHERE bank_id = $1 AND tags IS NOT NULL AND tags != '{{}}'
{pattern_clause}
GROUP BY tag
ORDER BY count DESC, tag ASC
LIMIT ${limit_param} OFFSET ${offset_param}
""",
*params,
)
items = [{"tag": row["tag"], "count": row["count"]} for row in rows]
return {
"items": items,
"total": total,
"limit": limit,
"offset": offset,
}
async def get_entity_state(
self,
bank_id: str,
@@ -4365,6 +4545,7 @@ Guidelines:
contents: list[dict[str, Any]],
*,
request_context: "RequestContext",
document_tags: list[str] | None = None,
) -> dict[str, Any]:
"""Submit a batch retain operation to run asynchronously."""
await self._authenticate_tenant(request_context)
@@ -4388,14 +4569,16 @@ Guidelines:
)
# Submit task to background queue
await self._task_backend.submit_task(
{
"type": "batch_retain",
"operation_id": str(operation_id),
"bank_id": bank_id,
"contents": contents,
}
)
task_payload = {
"type": "batch_retain",
"operation_id": str(operation_id),
"bank_id": bank_id,
"contents": contents,
}
if document_tags:
task_payload["document_tags"] = document_tags
await self._task_backend.submit_task(task_payload)
logger.info(f"Retain task queued for bank_id={bank_id}, {len(contents)} items, operation_id={operation_id}")
@@ -85,6 +85,7 @@ class MemoryFact(BaseModel):
"metadata": {"source": "slack"},
"chunk_id": "bank123_session_abc123_0",
"activation": 0.95,
"tags": ["user_a", "session_123"],
}
}
)
@@ -102,6 +103,7 @@ class MemoryFact(BaseModel):
chunk_id: str | None = Field(
None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)"
)
tags: list[str] | None = Field(None, description="Visibility scope tags associated with this fact")
class ChunkInfo(BaseModel):
@@ -1268,6 +1268,7 @@ async def extract_facts_from_contents(
# mentioned_at: always the event_date (when the conversation/document occurred)
mentioned_at=content.event_date,
metadata=content.metadata,
tags=content.tags,
)
extracted_facts.append(extracted_fact)
@@ -45,6 +45,7 @@ async def insert_facts_batch(
metadata_jsons = []
chunk_ids = []
document_ids = []
tags_list = []
for fact in facts:
fact_texts.append(fact.fact_text)
@@ -65,16 +66,31 @@ async def insert_facts_batch(
chunk_ids.append(fact.chunk_id)
# Use per-fact document_id if available, otherwise fallback to batch-level document_id
document_ids.append(fact.document_id if fact.document_id else document_id)
# Convert tags to JSON string for proper batch insertion (PostgreSQL unnest doesn't handle 2D arrays well)
tags_list.append(json.dumps(fact.tags if fact.tags else []))
# Batch insert all facts
# Note: tags are passed as JSON strings and converted back to varchar[] via jsonb_array_elements_text + array_agg
results = await conn.fetch(
f"""
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id)
SELECT $1, * FROM unnest(
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[]
WITH input_data AS (
SELECT * FROM unnest(
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[], $15::jsonb[]
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags_json)
)
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags)
SELECT
$1,
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id,
COALESCE(
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
'{{}}'::varchar[]
)
FROM input_data
RETURNING id
""",
bank_id,
@@ -91,6 +107,7 @@ async def insert_facts_batch(
metadata_jsons,
chunk_ids,
document_ids,
tags_list,
)
unit_ids = [str(row["id"]) for row in results]
@@ -121,7 +138,13 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
async def handle_document_tracking(
conn, bank_id: str, document_id: str, combined_content: str, is_first_batch: bool, retain_params: dict | None = None
conn,
bank_id: str,
document_id: str,
combined_content: str,
is_first_batch: bool,
retain_params: dict | None = None,
document_tags: list[str] | None = None,
) -> None:
"""
Handle document tracking in the database.
@@ -133,6 +156,7 @@ async def handle_document_tracking(
combined_content: Combined content text from all content items
is_first_batch: Whether this is the first batch (for chunked operations)
retain_params: Optional parameters passed during retain (context, event_date, etc.)
document_tags: Optional list of tags to associate with the document
"""
import hashlib
@@ -149,13 +173,14 @@ async def handle_document_tracking(
# Insert document (or update if exists from concurrent operations)
await conn.execute(
f"""
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params)
VALUES ($1, $2, $3, $4, $5, $6)
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params, tags)
VALUES ($1, $2, $3, $4, $5, $6, $7)
ON CONFLICT (id, bank_id) DO UPDATE
SET original_text = EXCLUDED.original_text,
content_hash = EXCLUDED.content_hash,
metadata = EXCLUDED.metadata,
retain_params = EXCLUDED.retain_params,
tags = EXCLUDED.tags,
updated_at = NOW()
""",
document_id,
@@ -164,4 +189,5 @@ async def handle_document_tracking(
content_hash,
json.dumps({}), # Empty metadata dict
json.dumps(retain_params) if retain_params else None,
document_tags or [],
)
@@ -49,6 +49,7 @@ async def retain_batch(
is_first_batch: bool = True,
fact_type_override: str | None = None,
confidence_score: float | None = None,
document_tags: list[str] | None = None,
) -> tuple[list[list[str]], TokenUsage]:
"""
Process a batch of content through the retain pipeline.
@@ -67,6 +68,7 @@ async def retain_batch(
is_first_batch: Whether this is the first batch
fact_type_override: Override fact type for all facts
confidence_score: Confidence score for opinions
document_tags: Tags applied to all items in this batch
Returns:
Tuple of (unit ID lists, token usage for fact extraction)
@@ -88,12 +90,16 @@ async def retain_batch(
# Convert dicts to RetainContent objects
contents = []
for item in contents_dicts:
# Merge item-level tags with document-level tags
item_tags = item.get("tags", []) or []
merged_tags = list(set(item_tags + (document_tags or [])))
content = RetainContent(
content=item["content"],
context=item.get("context", ""),
event_date=item.get("event_date") or utcnow(),
metadata=item.get("metadata", {}),
entities=item.get("entities", []),
tags=merged_tags,
)
contents.append(content)
@@ -131,7 +137,7 @@ async def retain_batch(
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
)
else:
# Check for per-item document_ids
@@ -159,7 +165,7 @@ async def retain_batch(
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, document_tags
)
total_time = time.time() - start_time
@@ -225,7 +231,7 @@ async def retain_batch(
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
)
document_ids_added.append(document_id)
doc_id_mapping[None] = document_id # For backwards compatibility
@@ -269,7 +275,13 @@ async def retain_batch(
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, actual_doc_id, combined_content, is_first_batch, retain_params
conn,
bank_id,
actual_doc_id,
combined_content,
is_first_batch,
retain_params,
document_tags,
)
document_ids_added.append(actual_doc_id)
@@ -21,6 +21,7 @@ class RetainContentDict(TypedDict, total=False):
metadata: Custom key-value metadata (optional)
document_id: Document ID for this content item (optional)
entities: User-provided entities to merge with extracted entities (optional)
tags: Visibility scope tags for this content item (optional)
"""
content: str # Required
@@ -29,6 +30,7 @@ class RetainContentDict(TypedDict, total=False):
metadata: dict[str, str]
document_id: str
entities: list[dict[str, str]] # [{"text": "...", "type": "..."}]
tags: list[str] # Visibility scope tags
def _now_utc() -> datetime:
@@ -49,6 +51,7 @@ class RetainContent:
event_date: datetime = field(default_factory=_now_utc)
metadata: dict[str, str] = field(default_factory=dict)
entities: list[dict[str, str]] = field(default_factory=list) # User-provided entities
tags: list[str] = field(default_factory=list) # Visibility scope tags
@dataclass
@@ -113,6 +116,7 @@ class ExtractedFact:
context: str = ""
mentioned_at: datetime | None = None
metadata: dict[str, str] = field(default_factory=dict)
tags: list[str] = field(default_factory=list) # Visibility scope tags
@dataclass
@@ -158,6 +162,9 @@ class ProcessedFact:
# Track which content this fact came from (for user entity merging)
content_index: int = 0
# Visibility scope tags
tags: list[str] = field(default_factory=list)
@property
def is_duplicate(self) -> bool:
"""Check if this fact was marked as a duplicate."""
@@ -201,6 +208,7 @@ class ProcessedFact:
causal_relations=extracted_fact.causal_relations,
chunk_id=chunk_id,
content_index=extracted_fact.content_index,
tags=extracted_fact.tags,
)
@@ -232,6 +240,7 @@ class RetainBatch:
document_id: str | None = None
fact_type_override: str | None = None
confidence_score: float | None = None
document_tags: list[str] = field(default_factory=list) # Tags applied to all items
# Extracted data (populated during processing)
extracted_facts: list[ExtractedFact] = field(default_factory=list)
@@ -11,6 +11,7 @@ from abc import ABC, abstractmethod
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .tags import TagsMatch, filter_results_by_tags
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -43,6 +44,8 @@ class GraphRetriever(ABC):
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # TypedAdjacency, optional pre-loaded graph
tags: list[str] | None = None, # Visibility scope tags for filtering
tags_match: TagsMatch = "any", # How to match tags: 'any' (OR) or 'all' (AND)
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve relevant facts via graph traversal.
@@ -57,6 +60,7 @@ class GraphRetriever(ABC):
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
adjacency: Pre-loaded typed adjacency graph (optional, for MPFP)
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (List of RetrievalResult with activation scores, optional timing info)
@@ -114,6 +118,8 @@ class BFSGraphRetriever(GraphRetriever):
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # Not used by BFS
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts using BFS spreading activation.
@@ -129,7 +135,9 @@ class BFSGraphRetriever(GraphRetriever):
for interface compatibility but not used.
"""
async with acquire_with_retry(pool) as conn:
results = await self._retrieve_with_conn(conn, query_embedding_str, bank_id, fact_type, budget)
results = await self._retrieve_with_conn(
conn, query_embedding_str, bank_id, fact_type, budget, tags=tags, tags_match=tags_match
)
return results, None
async def _retrieve_with_conn(
@@ -139,33 +147,46 @@ class BFSGraphRetriever(GraphRetriever):
bank_id: str,
fact_type: str,
budget: int,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[RetrievalResult]:
"""Internal implementation with connection."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, self.entry_point_threshold, self.entry_point_limit]
if tags:
params.append(tags)
# Step 1: Find entry points
entry_points = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
query_embedding_str,
bank_id,
fact_type,
self.entry_point_threshold,
self.entry_point_limit,
*params,
)
if not entry_points:
logger.debug(
f"[BFS] No entry points found for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
)
return []
logger.debug(
f"[BFS] Found {len(entry_points)} entry points for fact_type={fact_type} "
f"(tags={tags}, tags_match={tags_match})"
)
# Step 2: BFS spreading activation
visited = set()
results = []
@@ -196,7 +217,7 @@ class BFSGraphRetriever(GraphRetriever):
f"""
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.occurred_end,
mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type,
mu.document_id, mu.chunk_id,
mu.document_id, mu.chunk_id, mu.tags,
ml.weight, ml.link_type, ml.from_unit_id
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
@@ -236,4 +257,8 @@ class BFSGraphRetriever(GraphRetriever):
neighbor_result = RetrievalResult.from_db_row(dict(n))
queue.append((neighbor_result, new_activation))
# Apply tags filtering (BFS may traverse into memories that don't match tags criteria)
if tags:
results = filter_results_by_tags(results, tags, match=tags_match)
return results
@@ -18,6 +18,7 @@ import time
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .graph_retrieval import GraphRetriever
from .tags import TagsMatch, filter_results_by_tags
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -30,26 +31,32 @@ async def _find_semantic_seeds(
fact_type: str,
limit: int = 20,
threshold: float = 0.3,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[RetrievalResult]:
"""Find semantic seeds via embedding search."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
if tags:
params.append(tags)
rows = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
query_embedding_str,
bank_id,
fact_type,
threshold,
limit,
*params,
)
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
@@ -95,6 +102,8 @@ class LinkExpansionRetriever(GraphRetriever):
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts by expanding links from seeds.
@@ -109,6 +118,7 @@ class LinkExpansionRetriever(GraphRetriever):
semantic_seeds: Pre-computed semantic entry points
temporal_seeds: Pre-computed temporal entry points
adjacency: Unused, kept for interface compatibility
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (results, timings)
@@ -125,15 +135,27 @@ class LinkExpansionRetriever(GraphRetriever):
else:
seeds_start = time.time()
all_seeds = await _find_semantic_seeds(
conn, query_embedding_str, bank_id, fact_type, limit=20, threshold=0.3
conn,
query_embedding_str,
bank_id,
fact_type,
limit=20,
threshold=0.3,
tags=tags,
tags_match=tags_match,
)
timings.seeds_time = time.time() - seeds_start
logger.debug(
f"[LinkExpansion] Found {len(all_seeds)} semantic seeds for fact_type={fact_type} "
f"(tags={tags}, tags_match={tags_match})"
)
# Add temporal seeds if provided
if temporal_seeds:
all_seeds.extend(temporal_seeds)
if not all_seeds:
logger.debug("[LinkExpansion] No seeds found, returning empty results")
return [], timings
seed_ids = list({s.id for s in all_seeds})
@@ -147,7 +169,7 @@ class LinkExpansionRetriever(GraphRetriever):
SELECT
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
COUNT(*)::float AS score
FROM {fq_table("unit_entities")} seed_ue
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
@@ -172,7 +194,7 @@ class LinkExpansionRetriever(GraphRetriever):
SELECT DISTINCT ON (mu.id)
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
ml.weight + 1.0 AS score
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
@@ -219,6 +241,10 @@ class LinkExpansionRetriever(GraphRetriever):
result.activation = row["score"]
results.append(result)
# Apply tags filtering (graph expansion may reach untagged memories)
if tags:
results = filter_results_by_tags(results, tags, match=tags_match)
timings.result_count = len(results)
timings.traverse = time.time() - start_time
@@ -23,6 +23,7 @@ from dataclasses import dataclass, field
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .graph_retrieval import GraphRetriever
from .tags import TagsMatch
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -448,7 +449,7 @@ async def fetch_memory_units_by_ids(
rows = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
AND fact_type = $2
@@ -503,6 +504,8 @@ class MPFPGraphRetriever(GraphRetriever):
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # Ignored - kept for interface compatibility
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts using MPFP algorithm with lazy edge loading.
@@ -517,6 +520,7 @@ class MPFPGraphRetriever(GraphRetriever):
semantic_seeds: Pre-computed semantic entry points
temporal_seeds: Pre-computed temporal entry points
adjacency: Ignored (kept for interface compatibility)
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (List of RetrievalResult with activation scores, MPFPTimings)
@@ -532,8 +536,13 @@ class MPFPGraphRetriever(GraphRetriever):
# If no semantic seeds provided, fall back to finding our own
if not semantic_seed_nodes:
seeds_start = time.time()
semantic_seed_nodes = await self._find_semantic_seeds(pool, query_embedding_str, bank_id, fact_type)
semantic_seed_nodes = await self._find_semantic_seeds(
pool, query_embedding_str, bank_id, fact_type, tags=tags, tags_match=tags_match
)
timings.seeds_time = time.time() - seeds_start
logger.debug(
f"[MPFP] Found {len(semantic_seed_nodes)} semantic seeds for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
)
# Collect all pattern jobs
pattern_jobs = []
@@ -549,6 +558,9 @@ class MPFPGraphRetriever(GraphRetriever):
pattern_jobs.append((temporal_seed_nodes, pattern))
if not pattern_jobs:
logger.debug(
f"[MPFP] No pattern jobs (semantic_seeds={len(semantic_seed_nodes)}, temporal_seeds={len(temporal_seed_nodes)})"
)
return [], timings
timings.pattern_count = len(pattern_jobs)
@@ -587,6 +599,7 @@ class MPFPGraphRetriever(GraphRetriever):
timings.fusion = time.time() - step_start
if not fused:
logger.debug(f"[MPFP] No fused results after RRF fusion (pattern_count={len(pattern_results)})")
return [], timings
# Get top result IDs
@@ -596,6 +609,13 @@ class MPFPGraphRetriever(GraphRetriever):
step_start = time.time()
results = await fetch_memory_units_by_ids(pool, result_ids, fact_type)
timings.fetch = time.time() - step_start
# Filter results by tags (graph traversal may have picked up unfiltered memories)
if tags:
from .tags import filter_results_by_tags
results = filter_results_by_tags(results, tags, match=tags_match)
timings.result_count = len(results)
# Add activation scores from fusion
@@ -634,8 +654,17 @@ class MPFPGraphRetriever(GraphRetriever):
fact_type: str,
limit: int = 20,
threshold: float = 0.3,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[SeedNode]:
"""Fallback: find semantic seeds via embedding search."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
if tags:
params.append(tags)
async with acquire_with_retry(pool) as conn:
rows = await conn.fetch(
f"""
@@ -645,14 +674,11 @@ class MPFPGraphRetriever(GraphRetriever):
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
query_embedding_str,
bank_id,
fact_type,
threshold,
limit,
*params,
)
return [SeedNode(node_id=str(r["id"]), score=r["similarity"]) for r in rows]
@@ -20,6 +20,7 @@ from ..memory_engine import fq_table
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
from .link_expansion_retrieval import LinkExpansionRetriever
from .mpfp_retrieval import MPFPGraphRetriever
from .tags import TagsMatch, build_tags_where_clause_simple
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -85,7 +86,12 @@ def set_default_graph_retriever(retriever: GraphRetriever) -> None:
async def retrieve_semantic(
conn, query_emb_str: str, bank_id: str, fact_type: str, limit: int
conn,
query_emb_str: str,
bank_id: str,
fact_type: str,
limit: int,
tags: list[str] | None = None,
) -> list[RetrievalResult]:
"""
Semantic retrieval via vector similarity.
@@ -96,31 +102,44 @@ async def retrieve_semantic(
agent_id: bank ID
fact_type: Fact type to filter
limit: Maximum results to return
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
List of RetrievalResult objects
"""
from .tags import TagsMatch, build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 5)
params = [query_emb_str, bank_id, fact_type, limit]
if tags:
params.append(tags)
results = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= 0.3
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $4
""",
query_emb_str,
bank_id,
fact_type,
limit,
*params,
)
return [RetrievalResult.from_db_row(dict(r)) for r in results]
async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, limit: int) -> list[RetrievalResult]:
async def retrieve_bm25(
conn,
query_text: str,
bank_id: str,
fact_type: str,
limit: int,
tags: list[str] | None = None,
) -> list[RetrievalResult]:
"""
BM25 keyword retrieval via full-text search.
@@ -130,12 +149,15 @@ async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, lim
agent_id: bank ID
fact_type: Fact type to filter
limit: Maximum results to return
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
List of RetrievalResult objects
"""
import re
from .tags import TagsMatch, build_tags_where_clause_simple
# Sanitize query text: remove special characters that have meaning in tsquery
# Keep only alphanumeric characters and spaces
sanitized_text = re.sub(r"[^\w\s]", " ", query_text.lower())
@@ -151,21 +173,24 @@ async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, lim
# This prevents empty results when some terms are missing
query_tsquery = " | ".join(tokens)
tags_clause = build_tags_where_clause_simple(tags, 5)
params = [query_tsquery, bank_id, fact_type, limit]
if tags:
params.append(tags)
results = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
ts_rank_cd(search_vector, to_tsquery('english', $1)) AS bm25_score
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND fact_type = $3
AND search_vector @@ to_tsquery('english', $1)
{tags_clause}
ORDER BY bm25_score DESC
LIMIT $4
""",
query_tsquery,
bank_id,
fact_type,
limit,
*params,
)
return [RetrievalResult.from_db_row(dict(r)) for r in results]
@@ -177,6 +202,8 @@ async def retrieve_semantic_bm25_combined(
bank_id: str,
fact_types: list[str],
limit: int,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]]:
"""
Combined semantic + BM25 retrieval for multiple fact types in a single query.
@@ -203,10 +230,14 @@ async def retrieve_semantic_bm25_combined(
# If no valid tokens for BM25, just run semantic
if not tokens:
tags_clause = build_tags_where_clause_simple(tags, 5, match=tags_match)
params = [query_emb_str, bank_id, fact_types, limit]
if tags:
params.append(tags)
results = await conn.fetch(
f"""
WITH semantic_ranked AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity,
NULL::float AS bm25_score,
'semantic' AS source,
@@ -216,16 +247,14 @@ async def retrieve_semantic_bm25_combined(
AND embedding IS NOT NULL
AND fact_type = ANY($3)
AND (1 - (embedding <=> $1::vector)) >= 0.3
{tags_clause}
)
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
similarity, bm25_score, source
FROM semantic_ranked
WHERE rn <= $4
""",
query_emb_str,
bank_id,
fact_types,
limit,
*params,
)
# Group by fact_type
result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = {
@@ -241,12 +270,18 @@ async def retrieve_semantic_bm25_combined(
query_tsquery = " | ".join(tokens)
# Build tags clause - param 6 if tags provided
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_emb_str, bank_id, fact_types, limit, query_tsquery]
if tags:
params.append(tags)
# Combined CTE query for both semantic and BM25 across all fact types
# Uses window functions to limit per fact_type per method
results = await conn.fetch(
f"""
WITH semantic_ranked AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity,
NULL::float AS bm25_score,
'semantic' AS source,
@@ -256,9 +291,10 @@ async def retrieve_semantic_bm25_combined(
AND embedding IS NOT NULL
AND fact_type = ANY($3)
AND (1 - (embedding <=> $1::vector)) >= 0.3
{tags_clause}
),
bm25_ranked AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
NULL::float AS similarity,
ts_rank_cd(search_vector, to_tsquery('english', $5)) AS bm25_score,
'bm25' AS source,
@@ -267,14 +303,15 @@ async def retrieve_semantic_bm25_combined(
WHERE bank_id = $2
AND fact_type = ANY($3)
AND search_vector @@ to_tsquery('english', $5)
{tags_clause}
),
semantic AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
similarity, bm25_score, source
FROM semantic_ranked WHERE rn <= $4
),
bm25 AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
similarity, bm25_score, source
FROM bm25_ranked WHERE rn <= $4
)
@@ -282,11 +319,7 @@ async def retrieve_semantic_bm25_combined(
UNION ALL
SELECT * FROM bm25
""",
query_emb_str,
bank_id,
fact_types,
limit,
query_tsquery,
*params,
)
# Group results by fact_type and source
@@ -313,6 +346,8 @@ async def retrieve_temporal_combined(
end_date: datetime,
budget: int,
semantic_threshold: float = 0.1,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> dict[str, list[RetrievalResult]]:
"""
Temporal retrieval for multiple fact types in a single query.
@@ -341,11 +376,17 @@ async def retrieve_temporal_combined(
if end_date.tzinfo is None:
end_date = end_date.replace(tzinfo=UTC)
# Build tags clause
tags_clause = build_tags_where_clause_simple(tags, 7, match=tags_match)
params = [query_emb_str, bank_id, fact_types, start_date, end_date, semantic_threshold]
if tags:
params.append(tags)
# Batch query: Get entry points for ALL fact types at once with window function
entry_points = await conn.fetch(
f"""
WITH ranked_entries AS (
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity,
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, embedding <=> $1::vector) AS rn
FROM {fq_table("memory_units")}
@@ -363,17 +404,13 @@ async def retrieve_temporal_combined(
(occurred_end IS NOT NULL AND occurred_end BETWEEN $4 AND $5)
)
AND (1 - (embedding <=> $1::vector)) >= $6
{tags_clause}
)
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, similarity
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags, similarity
FROM ranked_entries
WHERE rn <= 10
""",
query_emb_str,
bank_id,
fact_types,
start_date,
end_date,
semantic_threshold,
*params,
)
if not entry_points:
@@ -436,13 +473,20 @@ async def retrieve_temporal_combined(
budget_remaining = budget - len(ft_entry_points)
batch_size = 20
# Build tags clause for spreading (use param 6 since 1-5 are used)
spreading_tags_clause = build_tags_where_clause_simple(tags, 6, table_alias="mu.", match=tags_match)
while frontier and budget_remaining > 0:
batch_ids = frontier[:batch_size]
frontier = frontier[batch_size:]
spreading_params = [query_emb_str, batch_ids, ft, semantic_threshold, batch_size * 10]
if tags:
spreading_params.append(tags)
neighbors = await conn.fetch(
f"""
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id,
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
ml.weight, ml.link_type, ml.from_unit_id,
1 - (mu.embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_links")} ml
@@ -453,14 +497,11 @@ async def retrieve_temporal_combined(
AND mu.fact_type = $3
AND mu.embedding IS NOT NULL
AND (1 - (mu.embedding <=> $1::vector)) >= $4
{spreading_tags_clause}
ORDER BY ml.weight DESC
LIMIT $5
""",
query_emb_str,
batch_ids,
ft,
semantic_threshold,
batch_size * 10,
*spreading_params,
)
for n in neighbors:
@@ -529,6 +570,7 @@ async def retrieve_temporal(
end_date: datetime,
budget: int,
semantic_threshold: float = 0.1,
tags: list[str] | None = None,
) -> list[RetrievalResult]:
"""
Temporal retrieval with spreading activation.
@@ -547,6 +589,7 @@ async def retrieve_temporal(
end_date: End of time range
budget: Node budget for spreading
semantic_threshold: Minimum semantic similarity to include
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
List of RetrievalResult objects with temporal scores
@@ -558,9 +601,16 @@ async def retrieve_temporal(
if end_date.tzinfo is None:
end_date = end_date.replace(tzinfo=UTC)
from .tags import TagsMatch, build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 7)
params = [query_emb_str, bank_id, fact_type, start_date, end_date, semantic_threshold]
if tags:
params.append(tags)
entry_points = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
@@ -580,15 +630,11 @@ async def retrieve_temporal(
(occurred_end IS NOT NULL AND occurred_end BETWEEN $4 AND $5)
)
AND (1 - (embedding <=> $1::vector)) >= $6
{tags_clause}
ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, (embedding <=> $1::vector) ASC
LIMIT 10
""",
query_emb_str,
bank_id,
fact_type,
start_date,
end_date,
semantic_threshold,
*params,
)
if not entry_points:
@@ -740,6 +786,7 @@ async def retrieve_parallel(
query_analyzer: Optional["QueryAnalyzer"] = None,
graph_retriever: GraphRetriever | None = None,
temporal_constraint: tuple | None = None, # Pre-extracted temporal constraint
tags: list[str] | None = None, # Visibility scope tags for filtering
) -> ParallelRetrievalResult:
"""
Run 3-way or 4-way parallel retrieval (adds temporal if detected).
@@ -755,6 +802,7 @@ async def retrieve_parallel(
query_analyzer: Query analyzer to use (defaults to TransformerQueryAnalyzer)
graph_retriever: Graph retrieval strategy (defaults to configured retriever)
temporal_constraint: Pre-extracted temporal constraint (optional)
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
ParallelRetrievalResult with semantic, bm25, graph, temporal results and timings
@@ -775,6 +823,7 @@ async def retrieve_parallel(
retriever,
question_date,
query_analyzer,
tags=tags,
)
else:
# For BFS, extract temporal constraint upfront (legacy path)
@@ -785,7 +834,15 @@ async def retrieve_parallel(
query_text, reference_date=question_date, analyzer=query_analyzer
)
return await _retrieve_parallel_bfs(
pool, query_text, query_embedding_str, bank_id, fact_type, thinking_budget, temporal_constraint, retriever
pool,
query_text,
query_embedding_str,
bank_id,
fact_type,
thinking_budget,
temporal_constraint,
retriever,
tags=tags,
)
@@ -809,6 +866,7 @@ async def _retrieve_parallel_mpfp(
retriever: GraphRetriever,
question_date: datetime | None = None,
query_analyzer=None,
tags: list[str] | None = None,
) -> ParallelRetrievalResult:
"""
MPFP retrieval with true parallelization.
@@ -830,7 +888,9 @@ async def _retrieve_parallel_mpfp(
acquire_start = time.time()
async with acquire_with_retry(pool) as conn:
conn_wait = time.time() - acquire_start
results = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
results = await retrieve_semantic(
conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget, tags=tags
)
return _TimedResult(results, time.time() - start, conn_wait)
async def run_bm25() -> _TimedResult:
@@ -839,7 +899,7 @@ async def _retrieve_parallel_mpfp(
acquire_start = time.time()
async with acquire_with_retry(pool) as conn:
conn_wait = time.time() - acquire_start
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget)
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget, tags=tags)
return _TimedResult(results, time.time() - start, conn_wait)
async def run_graph() -> tuple[list[RetrievalResult], float, MPFPTimings | None]:
@@ -857,6 +917,7 @@ async def _retrieve_parallel_mpfp(
query_text=query_text,
semantic_seeds=None, # Let MPFP find its own seeds
temporal_seeds=None, # Don't wait for temporal extraction
tags=tags,
)
return results, time.time() - start, mpfp_timing
@@ -1028,6 +1089,7 @@ async def _retrieve_parallel_bfs(
thinking_budget: int,
temporal_constraint: tuple | None,
retriever: GraphRetriever,
tags: list[str] | None = None,
) -> ParallelRetrievalResult:
"""BFS retrieval: all methods run in parallel (original behavior)."""
import time
@@ -1035,13 +1097,15 @@ async def _retrieve_parallel_bfs(
async def run_semantic() -> _TimedResult:
start = time.time()
async with acquire_with_retry(pool) as conn:
results = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
results = await retrieve_semantic(
conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget, tags=tags
)
return _TimedResult(results, time.time() - start)
async def run_bm25() -> _TimedResult:
start = time.time()
async with acquire_with_retry(pool) as conn:
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget)
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget, tags=tags)
return _TimedResult(results, time.time() - start)
async def run_graph() -> _TimedResult:
@@ -1053,6 +1117,7 @@ async def _retrieve_parallel_bfs(
fact_type=fact_type,
budget=thinking_budget,
query_text=query_text,
tags=tags,
)
return _TimedResult(results, time.time() - start)
@@ -1068,6 +1133,7 @@ async def _retrieve_parallel_bfs(
tc_end,
budget=thinking_budget,
semantic_threshold=0.1,
tags=tags,
)
return _TimedResult(results, time.time() - start)
@@ -1122,6 +1188,8 @@ async def retrieve_all_fact_types_parallel(
question_date: datetime | None = None,
query_analyzer: Optional["QueryAnalyzer"] = None,
graph_retriever: GraphRetriever | None = None,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> MultiFactTypeRetrievalResult:
"""
Optimized retrieval for multiple fact types using batched queries.
@@ -1171,7 +1239,14 @@ async def retrieve_all_fact_types_parallel(
# Semantic + BM25 combined
semantic_bm25_results = await retrieve_semantic_bm25_combined(
conn, query_embedding_str, query_text, bank_id, fact_types, thinking_budget
conn,
query_embedding_str,
query_text,
bank_id,
fact_types,
thinking_budget,
tags=tags,
tags_match=tags_match,
)
semantic_bm25_time = time.time() - semantic_bm25_start
@@ -1188,6 +1263,8 @@ async def retrieve_all_fact_types_parallel(
tc_end,
budget=thinking_budget,
semantic_threshold=0.1,
tags=tags,
tags_match=tags_match,
)
temporal_time = time.time() - temporal_start
@@ -1206,6 +1283,8 @@ async def retrieve_all_fact_types_parallel(
query_text=query_text,
semantic_seeds=None,
temporal_seeds=None,
tags=tags,
tags_match=tags_match,
)
return ft, results, time.time() - graph_start, mpfp_timing
@@ -0,0 +1,172 @@
"""
Tags filtering utilities for retrieval.
Provides SQL building functions for filtering memories by tags.
Supports four matching modes via TagsMatch enum:
- "any": OR matching, includes untagged memories (default, backward compatible)
- "all": AND matching, includes untagged memories
- "any_strict": OR matching, excludes untagged memories
- "all_strict": AND matching, excludes untagged memories
OR matching (any/any_strict): Memory matches if ANY of its tags overlap with request tags
AND matching (all/all_strict): Memory matches if ALL request tags are present in its tags
"""
from typing import Literal
TagsMatch = Literal["any", "all", "any_strict", "all_strict"]
def _parse_tags_match(match: TagsMatch) -> tuple[str, bool]:
"""
Parse TagsMatch into operator and include_untagged flag.
Returns:
Tuple of (operator, include_untagged)
- operator: "&&" for any/any_strict, "@>" for all/all_strict
- include_untagged: True for any/all, False for any_strict/all_strict
"""
if match == "any":
return "&&", True
elif match == "all":
return "@>", True
elif match == "any_strict":
return "&&", False
elif match == "all_strict":
return "@>", False
else:
# Default to "any" behavior
return "&&", True
def build_tags_where_clause(
tags: list[str] | None,
param_offset: int = 1,
table_alias: str = "",
match: TagsMatch = "any",
) -> tuple[str, list, int]:
"""
Build a SQL WHERE clause for filtering by tags.
Supports four matching modes:
- "any" (default): OR matching, includes untagged memories
- "all": AND matching, includes untagged memories
- "any_strict": OR matching, excludes untagged memories
- "all_strict": AND matching, excludes untagged memories
Args:
tags: List of tags to filter by. If None or empty, returns empty clause (no filtering).
param_offset: Starting parameter number for SQL placeholders (default 1).
table_alias: Optional table alias prefix (e.g., "mu." for "memory_units mu").
match: Matching mode. Defaults to "any".
Returns:
Tuple of (sql_clause, params, next_param_offset):
- sql_clause: SQL WHERE clause string
- params: List of parameter values to bind
- next_param_offset: Next available parameter number
Example:
>>> clause, params, next_offset = build_tags_where_clause(['user_a'], 3, 'mu.', 'any_strict')
>>> print(clause) # "AND mu.tags IS NOT NULL AND mu.tags != '{}' AND mu.tags && $3"
"""
if not tags:
return "", [], param_offset
column = f"{table_alias}tags" if table_alias else "tags"
operator, include_untagged = _parse_tags_match(match)
if include_untagged:
# Include untagged memories (NULL or empty array) OR matching tags
clause = f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_offset})"
else:
# Strict: only memories with matching tags (exclude NULL and empty)
clause = f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_offset}"
return clause, [tags], param_offset + 1
def build_tags_where_clause_simple(
tags: list[str] | None,
param_num: int,
table_alias: str = "",
match: TagsMatch = "any",
) -> str:
"""
Build a simple SQL WHERE clause for tags filtering.
This is a convenience version that returns just the clause string,
assuming the caller will add the tags array to their params list.
Args:
tags: List of tags to filter by. If None or empty, returns empty string.
param_num: Parameter number to use in the clause.
table_alias: Optional table alias prefix.
match: Matching mode. Defaults to "any".
Returns:
SQL clause string or empty string.
"""
if not tags:
return ""
column = f"{table_alias}tags" if table_alias else "tags"
operator, include_untagged = _parse_tags_match(match)
if include_untagged:
# Include untagged memories (NULL or empty array) OR matching tags
return f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_num})"
else:
# Strict: only memories with matching tags (exclude NULL and empty)
return f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_num}"
def filter_results_by_tags(
results: list,
tags: list[str] | None,
match: TagsMatch = "any",
) -> list:
"""
Filter retrieval results by tags in Python (for post-processing).
Used when SQL filtering isn't possible (e.g., graph traversal results).
Args:
results: List of RetrievalResult objects with a 'tags' attribute.
tags: List of tags to filter by. If None or empty, returns all results.
match: Matching mode. Defaults to "any".
Returns:
Filtered list of results.
"""
if not tags:
return results
_, include_untagged = _parse_tags_match(match)
is_any_match = match in ("any", "any_strict")
tags_set = set(tags)
filtered = []
for result in results:
result_tags = getattr(result, "tags", None)
# Check if untagged
is_untagged = result_tags is None or len(result_tags) == 0
if is_untagged:
if include_untagged:
filtered.append(result)
# else: skip untagged
else:
result_tags_set = set(result_tags)
if is_any_match:
# Any overlap
if result_tags_set & tags_set:
filtered.append(result)
else:
# All tags must be present
if tags_set <= result_tags_set:
filtered.append(result)
return filtered
@@ -11,6 +11,13 @@ from typing import Any, Literal
from pydantic import BaseModel, Field
class TemporalConstraint(BaseModel):
"""Detected temporal constraint from query analysis."""
start: datetime | None = Field(default=None, description="Start of temporal range")
end: datetime | None = Field(default=None, description="End of temporal range")
class QueryInfo(BaseModel):
"""Information about the search query."""
@@ -19,6 +26,11 @@ class QueryInfo(BaseModel):
timestamp: datetime = Field(description="When the query was executed")
budget: int = Field(description="Maximum nodes to explore")
max_tokens: int = Field(description="Maximum tokens to return in results")
tags: list[str] | None = Field(default=None, description="Tags filter applied to recall")
tags_match: str | None = Field(default=None, description="Tags matching mode: any, all, any_strict, all_strict")
temporal_constraint: TemporalConstraint | None = Field(
default=None, description="Detected temporal range from query"
)
class EntryPoint(BaseModel):
@@ -22,6 +22,7 @@ from .trace import (
SearchPhaseMetrics,
SearchSummary,
SearchTrace,
TemporalConstraint,
WeightComponents,
)
@@ -45,7 +46,14 @@ class SearchTracer:
json_output = trace.to_json()
"""
def __init__(self, query: str, budget: int, max_tokens: int):
def __init__(
self,
query: str,
budget: int,
max_tokens: int,
tags: list[str] | None = None,
tags_match: str | None = None,
):
"""
Initialize tracer.
@@ -53,10 +61,14 @@ class SearchTracer:
query: Search query text
budget: Maximum nodes to explore
max_tokens: Maximum tokens to return in results
tags: Tags filter applied to recall
tags_match: Tags matching mode (any, all, any_strict, all_strict)
"""
self.query_text = query
self.budget = budget
self.max_tokens = max_tokens
self.tags = tags
self.tags_match = tags_match
# Trace data
self.query_embedding: list[float] | None = None
@@ -66,6 +78,9 @@ class SearchTracer:
self.pruned: list[PruningDecision] = []
self.phase_metrics: list[SearchPhaseMetrics] = []
# Temporal constraint detected from query
self.temporal_constraint: TemporalConstraint | None = None
# New 4-way retrieval tracking
self.retrieval_results: list[RetrievalMethodResults] = []
self.rrf_merged: list[RRFMergeResult] = []
@@ -88,6 +103,11 @@ class SearchTracer:
"""Record the query embedding."""
self.query_embedding = embedding
def record_temporal_constraint(self, start: datetime | None, end: datetime | None):
"""Record the detected temporal constraint from query analysis."""
if start is not None or end is not None:
self.temporal_constraint = TemporalConstraint(start=start, end=end)
def add_entry_point(self, node_id: str, text: str, similarity: float, rank: int):
"""
Record an entry point.
@@ -428,6 +448,9 @@ class SearchTracer:
timestamp=datetime.now(UTC),
budget=self.budget,
max_tokens=self.max_tokens,
tags=self.tags,
tags_match=self.tags_match,
temporal_constraint=self.temporal_constraint,
)
# Create summary
@@ -48,6 +48,7 @@ class RetrievalResult:
chunk_id: str | None = None
access_count: int = 0
embedding: list[float] | None = None
tags: list[str] | None = None # Visibility scope tags
# Retrieval-specific scores (only one will be set depending on retrieval method)
similarity: float | None = None # Semantic retrieval
@@ -72,6 +73,7 @@ class RetrievalResult:
chunk_id=row.get("chunk_id"),
access_count=row.get("access_count", 0),
embedding=row.get("embedding"),
tags=row.get("tags"),
similarity=row.get("similarity"),
bm25_score=row.get("bm25_score"),
activation=row.get("activation"),
@@ -156,6 +158,7 @@ class ScoredResult:
"chunk_id": self.retrieval.chunk_id,
"access_count": self.retrieval.access_count,
"embedding": self.retrieval.embedding,
"tags": self.retrieval.tags,
"semantic_similarity": self.retrieval.similarity,
"bm25_score": self.retrieval.bm25_score,
}
+883
View File
@@ -0,0 +1,883 @@
"""
Tests for tags-based visibility scoping.
This module tests the tags feature which allows filtering memories by visibility tags.
Use cases:
- Multi-user agent: Agent has a single memory bank, users should only see memories from
conversations they participated in
- Student tracking: Teacher tracks students, students should only see their own data
The tags use OR-based matching: a memory matches if ANY of its tags overlap with the request tags.
"""
from datetime import datetime
import httpx
import pytest
import pytest_asyncio
from hindsight_api.api import create_app
from hindsight_api.engine.search.tags import build_tags_where_clause_simple, filter_results_by_tags
# ============================================================================
# Unit Tests for tags SQL builder
# ============================================================================
class TestTagsWhereClauseBuilder:
"""Unit tests for the tags WHERE clause SQL builder."""
def test_no_tags_returns_empty_string(self):
"""When tags is None, should return empty string (no filtering)."""
result = build_tags_where_clause_simple(None, 5)
assert result == ""
def test_empty_tags_list_returns_empty_string(self):
"""When tags is an empty list, should return empty string (no filtering)."""
result = build_tags_where_clause_simple([], 5)
assert result == ""
def test_tags_with_different_param_num(self):
"""Should use the provided parameter number."""
result = build_tags_where_clause_simple(["user_a", "user_b"], 3)
# Default is "any" which includes untagged
assert "$3" in result
def test_tags_with_table_alias(self):
"""Should include table alias when provided."""
result = build_tags_where_clause_simple(["user_a"], 5, table_alias="mu.")
assert "mu.tags" in result
# ---- Test "any" mode (OR, includes untagged - default) ----
def test_tags_match_any_includes_untagged(self):
"""When match='any', should include untagged memories (NULL or empty)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="any")
# Should use OR with NULL/empty check
assert "IS NULL" in result
assert "= '{}'" in result
assert "&&" in result # overlap operator
def test_tags_match_any_uses_overlap(self):
"""When match='any', should use overlap operator (&&)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="any")
assert "&&" in result
# ---- Test "all" mode (AND, includes untagged) ----
def test_tags_match_all_includes_untagged(self):
"""When match='all', should include untagged memories (NULL or empty)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="all")
# Should use OR with NULL/empty check
assert "IS NULL" in result
assert "= '{}'" in result
assert "@>" in result # contains operator
def test_tags_match_all_uses_contains(self):
"""When match='all', should use contains operator (@>)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="all")
assert "@>" in result
# ---- Test "any_strict" mode (OR, excludes untagged) ----
def test_tags_match_any_strict_excludes_untagged(self):
"""When match='any_strict', should exclude untagged memories."""
result = build_tags_where_clause_simple(["user_a"], 5, match="any_strict")
# Should require tags to be NOT NULL and not empty
assert "IS NOT NULL" in result
assert "!= '{}'" in result
assert "&&" in result # overlap operator
def test_tags_match_any_strict_uses_overlap(self):
"""When match='any_strict', should use overlap operator (&&)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="any_strict")
assert "&&" in result
# Should NOT include untagged
assert "IS NULL" not in result or "IS NOT NULL" in result
# ---- Test "all_strict" mode (AND, excludes untagged) ----
def test_tags_match_all_strict_excludes_untagged(self):
"""When match='all_strict', should exclude untagged memories."""
result = build_tags_where_clause_simple(["user_a"], 5, match="all_strict")
# Should require tags to be NOT NULL and not empty
assert "IS NOT NULL" in result
assert "!= '{}'" in result
assert "@>" in result # contains operator
def test_tags_match_all_strict_uses_contains(self):
"""When match='all_strict', should use contains operator (@>)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="all_strict")
assert "@>" in result
# ---- Test table alias with all modes ----
def test_tags_match_any_with_table_alias(self):
"""Should include table alias with any mode."""
result = build_tags_where_clause_simple(["user_a"], 3, table_alias="mu.", match="any")
assert "mu.tags" in result
def test_tags_match_all_strict_with_table_alias(self):
"""Should include table alias with all_strict mode."""
result = build_tags_where_clause_simple(["user_a", "user_b"], 3, table_alias="mu.", match="all_strict")
assert "mu.tags" in result
assert "@>" in result
assert "IS NOT NULL" in result
# ============================================================================
# Unit Tests for filter_results_by_tags (Python-side filtering)
# ============================================================================
class MockResult:
"""Mock result object for testing filter_results_by_tags."""
def __init__(self, tags):
self.tags = tags
class TestFilterResultsByTags:
"""Unit tests for the Python-side tags filter function."""
def test_no_tags_returns_all(self):
"""When tags is None, should return all results."""
results = [MockResult(["a"]), MockResult(["b"]), MockResult(None)]
filtered = filter_results_by_tags(results, None)
assert len(filtered) == 3
def test_empty_tags_returns_all(self):
"""When tags is empty list, should return all results."""
results = [MockResult(["a"]), MockResult(["b"]), MockResult(None)]
filtered = filter_results_by_tags(results, [])
assert len(filtered) == 3
# ---- Test "any" mode (OR, includes untagged) ----
def test_any_mode_includes_matching_tags(self):
"""'any' mode should include results with matching tags."""
results = [MockResult(["a"]), MockResult(["b"]), MockResult(["c"])]
filtered = filter_results_by_tags(results, ["a", "b"], match="any")
# "a" and "b" match, "c" doesn't match and isn't untagged, so excluded
assert len(filtered) == 2
tags_found = [r.tags[0] for r in filtered if r.tags]
assert "a" in tags_found
assert "b" in tags_found
assert "c" not in tags_found
def test_any_mode_includes_untagged(self):
"""'any' mode should include untagged results."""
results = [MockResult(["a"]), MockResult(None), MockResult([])]
filtered = filter_results_by_tags(results, ["a"], match="any")
assert len(filtered) == 3 # a matches, None is untagged, [] is untagged
def test_any_mode_includes_partial_overlap(self):
"""'any' mode should include results with ANY overlapping tag."""
results = [MockResult(["a", "x"]), MockResult(["b", "y"])]
filtered = filter_results_by_tags(results, ["a"], match="any")
# ["a", "x"] matches, ["b", "y"] doesn't, but untagged would be included
tags_found = [r.tags for r in filtered]
assert ["a", "x"] in tags_found
# ---- Test "any_strict" mode (OR, excludes untagged) ----
def test_any_strict_excludes_untagged(self):
"""'any_strict' mode should exclude untagged results."""
results = [MockResult(["a"]), MockResult(None), MockResult([])]
filtered = filter_results_by_tags(results, ["a"], match="any_strict")
assert len(filtered) == 1 # Only ["a"] matches
assert filtered[0].tags == ["a"]
def test_any_strict_excludes_non_matching(self):
"""'any_strict' mode should exclude non-matching tagged results."""
results = [MockResult(["a"]), MockResult(["b"]), MockResult(["c"])]
filtered = filter_results_by_tags(results, ["a"], match="any_strict")
assert len(filtered) == 1
assert filtered[0].tags == ["a"]
# ---- Test "all" mode (AND, includes untagged) ----
def test_all_mode_requires_all_tags(self):
"""'all' mode should require ALL requested tags to be present."""
results = [MockResult(["a", "b"]), MockResult(["a"]), MockResult(["b"])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all")
# Only ["a", "b"] has both tags, but untagged would also be included
tags_found = [r.tags for r in filtered]
assert ["a", "b"] in tags_found
def test_all_mode_includes_untagged(self):
"""'all' mode should include untagged results."""
results = [MockResult(["a", "b"]), MockResult(None), MockResult([])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all")
assert len(filtered) == 3 # ["a", "b"] matches, None is untagged, [] is untagged
# ---- Test "all_strict" mode (AND, excludes untagged) ----
def test_all_strict_requires_all_tags(self):
"""'all_strict' mode should require ALL requested tags."""
results = [MockResult(["a", "b"]), MockResult(["a"]), MockResult(["b"])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
assert len(filtered) == 1
assert filtered[0].tags == ["a", "b"]
def test_all_strict_excludes_untagged(self):
"""'all_strict' mode should exclude untagged results."""
results = [MockResult(["a", "b"]), MockResult(None), MockResult([])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
assert len(filtered) == 1
assert filtered[0].tags == ["a", "b"]
def test_all_strict_allows_superset(self):
"""'all_strict' mode should allow results with MORE tags than requested."""
results = [MockResult(["a", "b", "c"]), MockResult(["a"])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
assert len(filtered) == 1
assert filtered[0].tags == ["a", "b", "c"] # Has a, b, AND c
# ============================================================================
# Integration Tests for tags in retain/recall/reflect
# ============================================================================
@pytest_asyncio.fixture
async def api_client(memory):
"""Create an async test client for the FastAPI app."""
app = create_app(memory, initialize_memory=False)
transport = httpx.ASGITransport(app=app)
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
yield client
@pytest.fixture
def test_bank_id():
"""Provide a unique bank ID for this test run."""
return f"tags_test_{datetime.now().timestamp()}"
@pytest.mark.asyncio
async def test_retain_with_tags(api_client, test_bank_id):
"""Test that memories can be stored with tags."""
# Store memory with tags
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{
"content": "Alice loves hiking in the mountains.",
"tags": ["user_alice"]
}
]
}
)
assert response.status_code == 200
result = response.json()
assert result["success"] is True
assert result["items_count"] == 1
@pytest.mark.asyncio
async def test_retain_with_document_tags(api_client, test_bank_id):
"""Test that document-level tags are applied to all items."""
# Store memories with document-level tags
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"document_tags": ["session_123"],
"items": [
{"content": "Bob discussed the quarterly report."},
{"content": "Charlie mentioned the new product launch."}
]
}
)
assert response.status_code == 200
result = response.json()
assert result["success"] is True
assert result["items_count"] == 2
@pytest.mark.asyncio
async def test_retain_merges_document_and_item_tags(api_client, test_bank_id):
"""Test that document tags and item tags are merged."""
# Store memory with both document and item tags
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"document_tags": ["session_abc"],
"items": [
{
"content": "Dave talked about machine learning.",
"tags": ["user_dave"]
}
]
}
)
assert response.status_code == 200
result = response.json()
assert result["success"] is True
@pytest.mark.asyncio
async def test_recall_without_tags_returns_all_memories(api_client, test_bank_id):
"""Test that recall without tags returns all memories (no filtering)."""
# Store memories for different users
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Eve works on natural language processing.", "tags": ["user_eve"]},
{"content": "Frank specializes in computer vision.", "tags": ["user_frank"]},
]
}
)
assert response.status_code == 200
# Recall without tags - should return all
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "Who works on what?", "budget": "low"}
)
assert response.status_code == 200
results = response.json()["results"]
# Should find both Eve and Frank
texts = [r["text"] for r in results]
assert any("Eve" in t for t in texts), "Should find Eve"
assert any("Frank" in t for t in texts), "Should find Frank"
@pytest.mark.asyncio
async def test_recall_with_tags_filters_memories(api_client, test_bank_id):
"""Test that recall with tags only returns matching memories."""
# Store memories for different users
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Grace is a data scientist at Google.", "tags": ["user_grace"]},
{"content": "Henry is a software engineer at Meta.", "tags": ["user_henry"]},
]
}
)
assert response.status_code == 200
# Recall with user_grace tag - should only return Grace's memory
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "Who works at which company?", "budget": "low", "tags": ["user_grace"]}
)
assert response.status_code == 200
results = response.json()["results"]
# Should find Grace but not Henry
texts = [r["text"] for r in results]
assert any("Grace" in t for t in texts), "Should find Grace with user_grace tag"
# Henry should NOT be found since he has user_henry tag
assert not any("Henry" in t for t in texts), "Should NOT find Henry (different tag)"
@pytest.mark.asyncio
async def test_recall_with_multiple_tags_uses_or_matching(api_client, test_bank_id):
"""Test that multiple tags use OR matching (any match returns the memory)."""
# Store memories for different users
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Ivan leads the security team.", "tags": ["user_ivan"]},
{"content": "Julia manages the design team.", "tags": ["user_julia"]},
{"content": "Karl oversees the marketing team.", "tags": ["user_karl"]},
]
}
)
assert response.status_code == 200
# Recall with user_ivan OR user_julia - should return both Ivan and Julia, but not Karl
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "Who leads which team?", "budget": "low", "tags": ["user_ivan", "user_julia"]}
)
assert response.status_code == 200
results = response.json()["results"]
texts = [r["text"] for r in results]
assert any("Ivan" in t for t in texts), "Should find Ivan (tag matches)"
assert any("Julia" in t for t in texts), "Should find Julia (tag matches)"
assert not any("Karl" in t for t in texts), "Should NOT find Karl (tag doesn't match)"
@pytest.mark.asyncio
async def test_recall_returns_memories_with_any_overlapping_tag(api_client, test_bank_id):
"""Test that memories with multiple tags are returned if ANY tag matches."""
# Store memory with multiple tags
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{
"content": "Lisa and Mike discussed the budget in a group chat.",
"tags": ["user_lisa", "user_mike"] # Memory visible to both
},
{"content": "Nancy reviewed the budget alone.", "tags": ["user_nancy"]},
]
}
)
assert response.status_code == 200
# Recall with user_lisa - should return the group chat memory
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "What was discussed about the budget?", "budget": "low", "tags": ["user_lisa"]}
)
assert response.status_code == 200
results = response.json()["results"]
texts = [r["text"] for r in results]
assert any("Lisa" in t and "Mike" in t for t in texts), "Should find group chat (Lisa is in tags)"
assert not any("Nancy" in t for t in texts), "Should NOT find Nancy's memory"
@pytest.mark.asyncio
async def test_reflect_with_tags_filters_memories(api_client, test_bank_id):
"""Test that reflect with tags only uses matching memories for reasoning."""
# Store different memories for different users
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Oscar's favorite color is blue.", "tags": ["user_oscar"]},
{"content": "Peter's favorite color is red.", "tags": ["user_peter"]},
]
}
)
assert response.status_code == 200
# Reflect with user_oscar tag - should only use Oscar's memories
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/reflect",
json={
"query": "What is the favorite color?",
"budget": "low",
"tags": ["user_oscar"],
"include": {"facts": {}} # Request facts to verify what was used
}
)
assert response.status_code == 200
result = response.json()
# The response should mention Oscar's color (blue), not Peter's (red)
# Note: We can check based_on facts if they're returned
if result.get("based_on"):
fact_texts = [f["text"] for f in result["based_on"]]
# Should use Oscar's memory
assert any("Oscar" in t or "blue" in t for t in fact_texts), "Should use Oscar's memory"
@pytest.mark.asyncio
async def test_recall_with_empty_tags_returns_all(api_client, test_bank_id):
"""Test that empty tags list behaves same as no tags (returns all)."""
# Store memories
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Quinn studies mathematics.", "tags": ["user_quinn"]},
{"content": "Rachel studies physics.", "tags": ["user_rachel"]},
]
}
)
assert response.status_code == 200
# Recall with empty tags list - should return all
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "Who studies what?", "budget": "low", "tags": []}
)
assert response.status_code == 200
results = response.json()["results"]
texts = [r["text"] for r in results]
assert any("Quinn" in t for t in texts), "Should find Quinn"
assert any("Rachel" in t for t in texts), "Should find Rachel"
@pytest.mark.asyncio
async def test_multi_user_agent_visibility(api_client):
"""
Test multi-user agent visibility scoping.
Scenario:
- Agent has one memory bank
- Agent chats with User A (room 1) and User B (room 2) separately
- Agent also hosts a group chat with both users (room 3)
- User A should only see memories from rooms 1 and 3
- User B should only see memories from rooms 2 and 3
- Agent (no filter) should see all memories
"""
bank_id = f"multi_user_test_{datetime.now().timestamp()}"
# Store memories from different chat rooms
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
# Room 1: Agent + User A private chat
{"content": "User A said they prefer morning meetings.", "tags": ["user_a"]},
# Room 2: Agent + User B private chat
{"content": "User B mentioned they like afternoon meetings.", "tags": ["user_b"]},
# Room 3: Group chat with both users
{"content": "In the group meeting, they agreed to meet at noon.", "tags": ["user_a", "user_b"]},
]
}
)
assert response.status_code == 200
# User A queries - should see their private chat and group chat
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "What meeting time preferences were discussed?", "budget": "low", "tags": ["user_a"]}
)
assert response.status_code == 200
user_a_results = response.json()["results"]
user_a_texts = [r["text"] for r in user_a_results]
assert any("morning" in t for t in user_a_texts), "User A should see their own preference (morning)"
assert any("noon" in t for t in user_a_texts), "User A should see group chat (noon)"
assert not any("afternoon" in t for t in user_a_texts), "User A should NOT see User B's private preference"
# User B queries - should see their private chat and group chat
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "What meeting time preferences were discussed?", "budget": "low", "tags": ["user_b"]}
)
assert response.status_code == 200
user_b_results = response.json()["results"]
user_b_texts = [r["text"] for r in user_b_results]
assert any("afternoon" in t for t in user_b_texts), "User B should see their own preference (afternoon)"
assert any("noon" in t for t in user_b_texts), "User B should see group chat (noon)"
assert not any("morning" in t for t in user_b_texts), "User B should NOT see User A's private preference"
# Agent queries (no filter) - should see everything
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "What meeting time preferences were discussed?", "budget": "low"} # No tags
)
assert response.status_code == 200
agent_results = response.json()["results"]
agent_texts = [r["text"] for r in agent_results]
assert any("morning" in t for t in agent_texts), "Agent should see User A's preference"
assert any("afternoon" in t for t in agent_texts), "Agent should see User B's preference"
assert any("noon" in t for t in agent_texts), "Agent should see group chat"
@pytest.mark.asyncio
async def test_student_tracking_visibility(api_client):
"""
Test student tracking visibility scoping.
Scenario:
- Teacher bot has one memory bank
- Teacher records observations for Student A, Student B
- Student A should only see their own data
- Teacher (no filter) should see all student data
"""
bank_id = f"student_test_{datetime.now().timestamp()}"
# Store memories for different students
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Student A showed improvement in algebra today.", "tags": ["student_a"]},
{"content": "Student B struggled with geometry concepts.", "tags": ["student_b"]},
{"content": "Student A participated actively in class discussion.", "tags": ["student_a"]},
]
}
)
assert response.status_code == 200
# Student A queries - should only see their own data
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "How am I doing in class?", "budget": "low", "tags": ["student_a"]}
)
assert response.status_code == 200
student_a_results = response.json()["results"]
student_a_texts = [r["text"] for r in student_a_results]
assert any("algebra" in t for t in student_a_texts), "Student A should see their algebra progress"
assert any("participated" in t for t in student_a_texts), "Student A should see their participation"
assert not any("Student B" in t or "geometry" in t for t in student_a_texts), "Student A should NOT see Student B's data"
# Teacher queries (no filter) - should see all students
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "Which students need help?", "budget": "low"} # No tags
)
assert response.status_code == 200
teacher_results = response.json()["results"]
teacher_texts = [r["text"] for r in teacher_results]
assert any("Student A" in t for t in teacher_texts), "Teacher should see Student A's data"
assert any("Student B" in t for t in teacher_texts), "Teacher should see Student B's data"
# ============================================================================
# Tests for list_tags API endpoint
# ============================================================================
@pytest.mark.asyncio
async def test_list_tags_returns_all_tags(api_client):
"""Test that list_tags returns all unique tags with counts."""
bank_id = f"list_tags_test_{datetime.now().timestamp()}"
# Store memories with various tags
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Memory 1 for user alice.", "tags": ["user:alice"]},
{"content": "Memory 2 for user alice.", "tags": ["user:alice"]},
{"content": "Memory 3 for user bob.", "tags": ["user:bob"]},
{"content": "Memory 4 in session 123.", "tags": ["session:123"]},
{"content": "Memory 5 for alice in session 456.", "tags": ["user:alice", "session:456"]},
]
}
)
assert response.status_code == 200
# List all tags
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
assert response.status_code == 200
result = response.json()
# Verify structure
assert "items" in result
assert "total" in result
assert "limit" in result
assert "offset" in result
# Verify tags and counts
tags_map = {item["tag"]: item["count"] for item in result["items"]}
assert "user:alice" in tags_map
assert tags_map["user:alice"] == 3 # 3 memories have this tag
assert "user:bob" in tags_map
assert tags_map["user:bob"] == 1
assert "session:123" in tags_map
assert tags_map["session:123"] == 1
assert "session:456" in tags_map
assert tags_map["session:456"] == 1
assert result["total"] == 4 # 4 unique tags
@pytest.mark.asyncio
async def test_list_tags_with_wildcard_prefix(api_client):
"""Test that list_tags filters with prefix wildcard pattern (user:*)."""
bank_id = f"list_tags_wildcard_test_{datetime.now().timestamp()}"
# Store memories with various tags
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Memory for alice who works at tech.", "tags": ["user:alice"]},
{"content": "Memory for bob who is an engineer.", "tags": ["user:bob"]},
{"content": "Memory for charlie the designer.", "tags": ["user:charlie"]},
{"content": "Session memory about the meeting.", "tags": ["session:abc"]},
{"content": "Room memory for conference room.", "tags": ["room:123"]},
]
}
)
assert response.status_code == 200
# List tags with 'user:*' wildcard pattern
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "user:*"})
assert response.status_code == 200
result = response.json()
# Should only return user:* tags
tags = [item["tag"] for item in result["items"]]
assert "user:alice" in tags
assert "user:bob" in tags
assert "user:charlie" in tags
assert "session:abc" not in tags
assert "room:123" not in tags
assert result["total"] == 3
@pytest.mark.asyncio
async def test_list_tags_with_wildcard_suffix(api_client):
"""Test that list_tags filters with suffix wildcard pattern (*-admin)."""
bank_id = f"list_tags_suffix_test_{datetime.now().timestamp()}"
# Store memories with various tags
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Admin role memory for super admin.", "tags": ["role-admin"]},
{"content": "Super admin memory about permissions.", "tags": ["super-admin"]},
{"content": "User memory for standard users.", "tags": ["role-user"]},
{"content": "Guest memory for visitors.", "tags": ["role-guest"]},
]
}
)
assert response.status_code == 200
# List tags with '*-admin' wildcard pattern (suffix match)
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "*-admin"})
assert response.status_code == 200
result = response.json()
# Should only return *-admin tags
tags = [item["tag"] for item in result["items"]]
assert "role-admin" in tags
assert "super-admin" in tags
assert "role-user" not in tags
assert "role-guest" not in tags
assert result["total"] == 2
@pytest.mark.asyncio
async def test_list_tags_with_wildcard_middle(api_client):
"""Test that list_tags filters with middle wildcard pattern (env*-prod)."""
bank_id = f"list_tags_middle_test_{datetime.now().timestamp()}"
# Store memories with various tags - use meaningful content for fact extraction
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "The production environment is configured with high availability and uses AWS infrastructure.", "tags": ["env-prod"]},
{"content": "The enterprise environment for production runs on dedicated servers with 24/7 monitoring.", "tags": ["environment-prod"]},
{"content": "The staging environment mirrors production but uses smaller instance sizes.", "tags": ["env-staging"]},
{"content": "The development environment allows developers to test their code locally.", "tags": ["env-dev"]},
]
}
)
assert response.status_code == 200
# List tags with 'env*-prod' wildcard pattern (middle match)
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "env*-prod"})
assert response.status_code == 200
result = response.json()
# Should only return env*-prod tags
tags = [item["tag"] for item in result["items"]]
assert "env-prod" in tags
assert "environment-prod" in tags
assert "env-staging" not in tags
assert "env-dev" not in tags
assert result["total"] == 2
@pytest.mark.asyncio
async def test_list_tags_case_insensitive(api_client):
"""Test that list_tags wildcard matching is case-insensitive."""
bank_id = f"list_tags_case_test_{datetime.now().timestamp()}"
# Store memories with mixed case tags - use meaningful content
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Alice is a software engineer who specializes in machine learning algorithms.", "tags": ["User:Alice"]},
{"content": "Bob works as a data scientist at a large technology company.", "tags": ["user:bob"]},
{"content": "Charlie is the lead designer responsible for the user interface.", "tags": ["USER:CHARLIE"]},
]
}
)
assert response.status_code == 200
# List tags with lowercase pattern - should match all cases
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "user:*"})
assert response.status_code == 200
result = response.json()
# Should match all user tags regardless of case
tags = [item["tag"] for item in result["items"]]
assert len(tags) == 3
assert result["total"] == 3
@pytest.mark.asyncio
async def test_list_tags_pagination(api_client):
"""Test that list_tags supports pagination."""
bank_id = f"list_tags_pagination_test_{datetime.now().timestamp()}"
# Store memories with many tags - use meaningful content for fact extraction
names = ["Alice", "Bob", "Charlie", "Diana", "Eve", "Frank", "Grace", "Henry", "Ivan", "Julia"]
items = [
{"content": f"{name} works as a software engineer at company {i}.", "tags": [f"tag:{i:03d}"]}
for i, name in enumerate(names)
]
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={"items": items}
)
assert response.status_code == 200
# Get first page (limit 3)
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"limit": 3, "offset": 0})
assert response.status_code == 200
result = response.json()
assert len(result["items"]) == 3
assert result["total"] == 10
assert result["limit"] == 3
assert result["offset"] == 0
# Get second page
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"limit": 3, "offset": 3})
assert response.status_code == 200
result = response.json()
assert len(result["items"]) == 3
assert result["offset"] == 3
@pytest.mark.asyncio
async def test_list_tags_empty_bank(api_client):
"""Test that list_tags returns empty for bank with no tags."""
bank_id = f"list_tags_empty_test_{datetime.now().timestamp()}"
# List tags without storing anything
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
assert response.status_code == 200
result = response.json()
assert result["items"] == []
assert result["total"] == 0
@pytest.mark.asyncio
async def test_list_tags_ordered_by_count(api_client):
"""Test that list_tags returns tags ordered by frequency (most used first)."""
bank_id = f"list_tags_order_test_{datetime.now().timestamp()}"
# Store memories with tags having different frequencies - use meaningful content
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Alice works at a startup company as a developer.", "tags": ["rare"]},
{"content": "Bob is a senior engineer at Google.", "tags": ["common"]},
{"content": "Charlie manages the marketing team at Microsoft.", "tags": ["common"]},
{"content": "Diana leads the design department at Apple.", "tags": ["common"]},
{"content": "Eve is a data scientist at Amazon.", "tags": ["medium"]},
{"content": "Frank handles customer support at Meta.", "tags": ["medium"]},
]
}
)
assert response.status_code == 200
# List tags - should be ordered by count descending
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
assert response.status_code == 200
result = response.json()
tags = [item["tag"] for item in result["items"]]
# common (3) should come before medium (2) which should come before rare (1)
assert tags.index("common") < tags.index("medium")
assert tags.index("medium") < tags.index("rare")
+5 -1
View File
@@ -5,7 +5,7 @@ use crossterm::{
execute,
terminal::{disable_raw_mode, enable_raw_mode, EnterAlternateScreen, LeaveAlternateScreen},
};
use hindsight_client::types::{BankListItem, RecallResult, EntityListItem, Budget};
use hindsight_client::types::{BankListItem, RecallResult, EntityListItem, Budget, TagsMatch};
use serde_json::{Map, Value};
use ratatui::{
backend::{Backend, CrosstermBackend},
@@ -341,6 +341,8 @@ impl App {
trace: false,
query_timestamp: None,
include: None,
tags: None,
tags_match: TagsMatch::Any,
};
let result = client.recall(&bank_id, &request, false)
@@ -357,6 +359,8 @@ impl App {
max_tokens: 4096,
include: None,
response_schema: None,
tags: None,
tags_match: TagsMatch::Any,
};
let result = client.reflect(&bank_id, &request, false)
+9 -1
View File
@@ -9,7 +9,7 @@ use crate::output::{self, OutputFormat};
use crate::ui;
// Import types from generated client
use hindsight_client::types::{Budget, ChunkIncludeOptions, IncludeOptions};
use hindsight_client::types::{Budget, ChunkIncludeOptions, IncludeOptions, TagsMatch};
use serde_json;
// Helper function to parse budget string to Budget enum
@@ -60,6 +60,8 @@ pub fn recall(
trace,
query_timestamp: None,
include,
tags: None,
tags_match: TagsMatch::Any,
};
let response = client.recall(agent_id, &request, verbose);
@@ -116,6 +118,8 @@ pub fn reflect(
max_tokens: max_tokens.unwrap_or(4096),
include: None,
response_schema,
tags: None,
tags_match: TagsMatch::Any,
};
let response = client.reflect(agent_id, &request, verbose);
@@ -162,11 +166,13 @@ pub fn retain(
timestamp: None,
document_id: Some(doc_id.clone()),
entities: None,
tags: None,
};
let request = RetainRequest {
items: vec![item],
async_: r#async,
document_tags: None,
};
let response = client.retain(agent_id, &request, r#async, verbose);
@@ -272,6 +278,7 @@ pub fn retain_files(
timestamp: None,
document_id: Some(doc_id),
entities: None,
tags: None,
});
pb.inc(1);
@@ -288,6 +295,7 @@ pub fn retain_files(
let request = RetainRequest {
items,
async_: r#async,
document_tags: None,
};
let response = client.retain(agent_id, &request, r#async, verbose);
@@ -39,6 +39,7 @@ hindsight_client_api/models/http_validation_error.py
hindsight_client_api/models/include_options.py
hindsight_client_api/models/list_documents_response.py
hindsight_client_api/models/list_memory_units_response.py
hindsight_client_api/models/list_tags_response.py
hindsight_client_api/models/memory_item.py
hindsight_client_api/models/operation_response.py
hindsight_client_api/models/operations_list_response.py
@@ -51,6 +52,7 @@ hindsight_client_api/models/reflect_request.py
hindsight_client_api/models/reflect_response.py
hindsight_client_api/models/retain_request.py
hindsight_client_api/models/retain_response.py
hindsight_client_api/models/tag_item.py
hindsight_client_api/models/token_usage.py
hindsight_client_api/models/update_disposition_request.py
hindsight_client_api/models/validation_error.py
@@ -64,6 +64,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
from hindsight_client_api.models.include_options import IncludeOptions
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
from hindsight_client_api.models.list_tags_response import ListTagsResponse
from hindsight_client_api.models.memory_item import MemoryItem
from hindsight_client_api.models.operation_response import OperationResponse
from hindsight_client_api.models.operations_list_response import OperationsListResponse
@@ -76,6 +77,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
from hindsight_client_api.models.reflect_response import ReflectResponse
from hindsight_client_api.models.retain_request import RetainRequest
from hindsight_client_api.models.retain_response import RetainResponse
from hindsight_client_api.models.tag_item import TagItem
from hindsight_client_api.models.token_usage import TokenUsage
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
from hindsight_client_api.models.validation_error import ValidationError
@@ -17,11 +17,12 @@ from typing import Any, Dict, List, Optional, Tuple, Union
from typing_extensions import Annotated
from pydantic import Field, StrictInt, StrictStr
from typing import Optional
from typing import Any, Optional
from typing_extensions import Annotated
from hindsight_client_api.models.delete_response import DeleteResponse
from hindsight_client_api.models.graph_data_response import GraphDataResponse
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
from hindsight_client_api.models.list_tags_response import ListTagsResponse
from hindsight_client_api.models.recall_request import RecallRequest
from hindsight_client_api.models.recall_response import RecallResponse
from hindsight_client_api.models.reflect_request import ReflectRequest
@@ -654,6 +655,299 @@ class MemoryApi:
@validate_call
async def get_memory(
self,
bank_id: StrictStr,
memory_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> object:
"""Get memory unit
Get a single memory unit by ID with all its metadata including entities and tags.
:param bank_id: (required)
:type bank_id: str
:param memory_id: (required)
:type memory_id: str
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._get_memory_serialize(
bank_id=bank_id,
memory_id=memory_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "object",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
await response_data.read()
return self.api_client.response_deserialize(
response_data=response_data,
response_types_map=_response_types_map,
).data
@validate_call
async def get_memory_with_http_info(
self,
bank_id: StrictStr,
memory_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> ApiResponse[object]:
"""Get memory unit
Get a single memory unit by ID with all its metadata including entities and tags.
:param bank_id: (required)
:type bank_id: str
:param memory_id: (required)
:type memory_id: str
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._get_memory_serialize(
bank_id=bank_id,
memory_id=memory_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "object",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
await response_data.read()
return self.api_client.response_deserialize(
response_data=response_data,
response_types_map=_response_types_map,
)
@validate_call
async def get_memory_without_preload_content(
self,
bank_id: StrictStr,
memory_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> RESTResponseType:
"""Get memory unit
Get a single memory unit by ID with all its metadata including entities and tags.
:param bank_id: (required)
:type bank_id: str
:param memory_id: (required)
:type memory_id: str
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._get_memory_serialize(
bank_id=bank_id,
memory_id=memory_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "object",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
return response_data.response
def _get_memory_serialize(
self,
bank_id,
memory_id,
authorization,
_request_auth,
_content_type,
_headers,
_host_index,
) -> RequestSerialized:
_host = None
_collection_formats: Dict[str, str] = {
}
_path_params: Dict[str, str] = {}
_query_params: List[Tuple[str, str]] = []
_header_params: Dict[str, Optional[str]] = _headers or {}
_form_params: List[Tuple[str, str]] = []
_files: Dict[
str, Union[str, bytes, List[str], List[bytes], List[Tuple[str, bytes]]]
] = {}
_body_params: Optional[bytes] = None
# process the path parameters
if bank_id is not None:
_path_params['bank_id'] = bank_id
if memory_id is not None:
_path_params['memory_id'] = memory_id
# process the query parameters
# process the header parameters
if authorization is not None:
_header_params['authorization'] = authorization
# process the form parameters
# process the body parameter
# set the HTTP header `Accept`
if 'Accept' not in _header_params:
_header_params['Accept'] = self.api_client.select_header_accept(
[
'application/json'
]
)
# authentication setting
_auth_settings: List[str] = [
]
return self.api_client.param_serialize(
method='GET',
resource_path='/v1/default/banks/{bank_id}/memories/{memory_id}',
path_params=_path_params,
query_params=_query_params,
header_params=_header_params,
body=_body_params,
post_params=_form_params,
files=_files,
auth_settings=_auth_settings,
collection_formats=_collection_formats,
_host=_host,
_request_auth=_request_auth
)
@validate_call
async def list_memories(
self,
@@ -1000,6 +1294,335 @@ class MemoryApi:
@validate_call
async def list_tags(
self,
bank_id: StrictStr,
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> ListTagsResponse:
"""List tags
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
:param bank_id: (required)
:type bank_id: str
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
:type q: str
:param limit: Maximum number of tags to return
:type limit: int
:param offset: Offset for pagination
:type offset: int
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._list_tags_serialize(
bank_id=bank_id,
q=q,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "ListTagsResponse",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
await response_data.read()
return self.api_client.response_deserialize(
response_data=response_data,
response_types_map=_response_types_map,
).data
@validate_call
async def list_tags_with_http_info(
self,
bank_id: StrictStr,
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> ApiResponse[ListTagsResponse]:
"""List tags
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
:param bank_id: (required)
:type bank_id: str
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
:type q: str
:param limit: Maximum number of tags to return
:type limit: int
:param offset: Offset for pagination
:type offset: int
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._list_tags_serialize(
bank_id=bank_id,
q=q,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "ListTagsResponse",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
await response_data.read()
return self.api_client.response_deserialize(
response_data=response_data,
response_types_map=_response_types_map,
)
@validate_call
async def list_tags_without_preload_content(
self,
bank_id: StrictStr,
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> RESTResponseType:
"""List tags
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
:param bank_id: (required)
:type bank_id: str
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
:type q: str
:param limit: Maximum number of tags to return
:type limit: int
:param offset: Offset for pagination
:type offset: int
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._list_tags_serialize(
bank_id=bank_id,
q=q,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "ListTagsResponse",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
return response_data.response
def _list_tags_serialize(
self,
bank_id,
q,
limit,
offset,
authorization,
_request_auth,
_content_type,
_headers,
_host_index,
) -> RequestSerialized:
_host = None
_collection_formats: Dict[str, str] = {
}
_path_params: Dict[str, str] = {}
_query_params: List[Tuple[str, str]] = []
_header_params: Dict[str, Optional[str]] = _headers or {}
_form_params: List[Tuple[str, str]] = []
_files: Dict[
str, Union[str, bytes, List[str], List[bytes], List[Tuple[str, bytes]]]
] = {}
_body_params: Optional[bytes] = None
# process the path parameters
if bank_id is not None:
_path_params['bank_id'] = bank_id
# process the query parameters
if q is not None:
_query_params.append(('q', q))
if limit is not None:
_query_params.append(('limit', limit))
if offset is not None:
_query_params.append(('offset', offset))
# process the header parameters
if authorization is not None:
_header_params['authorization'] = authorization
# process the form parameters
# process the body parameter
# set the HTTP header `Accept`
if 'Accept' not in _header_params:
_header_params['Accept'] = self.api_client.select_header_accept(
[
'application/json'
]
)
# authentication setting
_auth_settings: List[str] = [
]
return self.api_client.param_serialize(
method='GET',
resource_path='/v1/default/banks/{bank_id}/tags',
path_params=_path_params,
query_params=_query_params,
header_params=_header_params,
body=_body_params,
post_params=_form_params,
files=_files,
auth_settings=_auth_settings,
collection_formats=_collection_formats,
_host=_host,
_request_auth=_request_auth
)
@validate_call
async def recall_memories(
self,
@@ -42,6 +42,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
from hindsight_client_api.models.include_options import IncludeOptions
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
from hindsight_client_api.models.list_tags_response import ListTagsResponse
from hindsight_client_api.models.memory_item import MemoryItem
from hindsight_client_api.models.operation_response import OperationResponse
from hindsight_client_api.models.operations_list_response import OperationsListResponse
@@ -54,6 +55,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
from hindsight_client_api.models.reflect_response import ReflectResponse
from hindsight_client_api.models.retain_request import RetainRequest
from hindsight_client_api.models.retain_response import RetainResponse
from hindsight_client_api.models.tag_item import TagItem
from hindsight_client_api.models.token_usage import TokenUsage
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
from hindsight_client_api.models.validation_error import ValidationError
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, StrictInt, StrictStr
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
from typing import Any, ClassVar, Dict, List, Optional
from typing import Optional, Set
from typing_extensions import Self
@@ -33,7 +33,8 @@ class DocumentResponse(BaseModel):
created_at: StrictStr
updated_at: StrictStr
memory_unit_count: StrictInt
__properties: ClassVar[List[str]] = ["id", "bank_id", "original_text", "content_hash", "created_at", "updated_at", "memory_unit_count"]
tags: Optional[List[StrictStr]] = Field(default=None, description="Tags associated with this document")
__properties: ClassVar[List[str]] = ["id", "bank_id", "original_text", "content_hash", "created_at", "updated_at", "memory_unit_count", "tags"]
model_config = ConfigDict(
populate_by_name=True,
@@ -97,7 +98,8 @@ class DocumentResponse(BaseModel):
"content_hash": obj.get("content_hash"),
"created_at": obj.get("created_at"),
"updated_at": obj.get("updated_at"),
"memory_unit_count": obj.get("memory_unit_count")
"memory_unit_count": obj.get("memory_unit_count"),
"tags": obj.get("tags")
})
return _obj
@@ -0,0 +1,101 @@
# coding: utf-8
"""
Hindsight HTTP API
HTTP API for Hindsight
The version of the OpenAPI document: 0.1.0
Generated by OpenAPI Generator (https://openapi-generator.tech)
Do not edit the class manually.
""" # noqa: E501
from __future__ import annotations
import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, StrictInt
from typing import Any, ClassVar, Dict, List
from hindsight_client_api.models.tag_item import TagItem
from typing import Optional, Set
from typing_extensions import Self
class ListTagsResponse(BaseModel):
"""
Response model for list tags endpoint.
""" # noqa: E501
items: List[TagItem]
total: StrictInt
limit: StrictInt
offset: StrictInt
__properties: ClassVar[List[str]] = ["items", "total", "limit", "offset"]
model_config = ConfigDict(
populate_by_name=True,
validate_assignment=True,
protected_namespaces=(),
)
def to_str(self) -> str:
"""Returns the string representation of the model using alias"""
return pprint.pformat(self.model_dump(by_alias=True))
def to_json(self) -> str:
"""Returns the JSON representation of the model using alias"""
# TODO: pydantic v2: use .model_dump_json(by_alias=True, exclude_unset=True) instead
return json.dumps(self.to_dict())
@classmethod
def from_json(cls, json_str: str) -> Optional[Self]:
"""Create an instance of ListTagsResponse from a JSON string"""
return cls.from_dict(json.loads(json_str))
def to_dict(self) -> Dict[str, Any]:
"""Return the dictionary representation of the model using alias.
This has the following differences from calling pydantic's
`self.model_dump(by_alias=True)`:
* `None` is only added to the output dict for nullable fields that
were set at model initialization. Other fields with value `None`
are ignored.
"""
excluded_fields: Set[str] = set([
])
_dict = self.model_dump(
by_alias=True,
exclude=excluded_fields,
exclude_none=True,
)
# override the default output from pydantic by calling `to_dict()` of each item in items (list)
_items = []
if self.items:
for _item_items in self.items:
if _item_items:
_items.append(_item_items.to_dict())
_dict['items'] = _items
return _dict
@classmethod
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional[Self]:
"""Create an instance of ListTagsResponse from a dict"""
if obj is None:
return None
if not isinstance(obj, dict):
return cls.model_validate(obj)
_obj = cls.model_validate({
"items": [TagItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None,
"total": obj.get("total"),
"limit": obj.get("limit"),
"offset": obj.get("offset")
})
return _obj
@@ -34,7 +34,8 @@ class MemoryItem(BaseModel):
metadata: Optional[Dict[str, StrictStr]] = None
document_id: Optional[StrictStr] = None
entities: Optional[List[EntityInput]] = None
__properties: ClassVar[List[str]] = ["content", "timestamp", "context", "metadata", "document_id", "entities"]
tags: Optional[List[StrictStr]] = None
__properties: ClassVar[List[str]] = ["content", "timestamp", "context", "metadata", "document_id", "entities", "tags"]
model_config = ConfigDict(
populate_by_name=True,
@@ -107,6 +108,11 @@ class MemoryItem(BaseModel):
if self.entities is None and "entities" in self.model_fields_set:
_dict['entities'] = None
# set to None if tags (nullable) is None
# and model_fields_set contains the field
if self.tags is None and "tags" in self.model_fields_set:
_dict['tags'] = None
return _dict
@classmethod
@@ -124,7 +130,8 @@ class MemoryItem(BaseModel):
"context": obj.get("context"),
"metadata": obj.get("metadata"),
"document_id": obj.get("document_id"),
"entities": [EntityInput.from_dict(_item) for _item in obj["entities"]] if obj.get("entities") is not None else None
"entities": [EntityInput.from_dict(_item) for _item in obj["entities"]] if obj.get("entities") is not None else None,
"tags": obj.get("tags")
})
return _obj
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictInt, StrictStr
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictInt, StrictStr, field_validator
from typing import Any, ClassVar, Dict, List, Optional
from hindsight_client_api.models.budget import Budget
from hindsight_client_api.models.include_options import IncludeOptions
@@ -35,7 +35,19 @@ class RecallRequest(BaseModel):
trace: Optional[StrictBool] = False
query_timestamp: Optional[StrictStr] = None
include: Optional[IncludeOptions] = Field(default=None, description="Options for including additional data (entities are included by default)")
__properties: ClassVar[List[str]] = ["query", "types", "budget", "max_tokens", "trace", "query_timestamp", "include"]
tags: Optional[List[StrictStr]] = None
tags_match: Optional[StrictStr] = Field(default='any', description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).")
__properties: ClassVar[List[str]] = ["query", "types", "budget", "max_tokens", "trace", "query_timestamp", "include", "tags", "tags_match"]
@field_validator('tags_match')
def tags_match_validate_enum(cls, value):
"""Validates the enum"""
if value is None:
return value
if value not in set(['any', 'all', 'any_strict', 'all_strict']):
raise ValueError("must be one of enum values ('any', 'all', 'any_strict', 'all_strict')")
return value
model_config = ConfigDict(
populate_by_name=True,
@@ -89,6 +101,11 @@ class RecallRequest(BaseModel):
if self.query_timestamp is None and "query_timestamp" in self.model_fields_set:
_dict['query_timestamp'] = None
# set to None if tags (nullable) is None
# and model_fields_set contains the field
if self.tags is None and "tags" in self.model_fields_set:
_dict['tags'] = None
return _dict
@classmethod
@@ -107,7 +124,9 @@ class RecallRequest(BaseModel):
"max_tokens": obj.get("max_tokens") if obj.get("max_tokens") is not None else 4096,
"trace": obj.get("trace") if obj.get("trace") is not None else False,
"query_timestamp": obj.get("query_timestamp"),
"include": IncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None
"include": IncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None,
"tags": obj.get("tags"),
"tags_match": obj.get("tags_match") if obj.get("tags_match") is not None else 'any'
})
return _obj
@@ -37,7 +37,8 @@ class RecallResult(BaseModel):
document_id: Optional[StrictStr] = None
metadata: Optional[Dict[str, StrictStr]] = None
chunk_id: Optional[StrictStr] = None
__properties: ClassVar[List[str]] = ["id", "text", "type", "entities", "context", "occurred_start", "occurred_end", "mentioned_at", "document_id", "metadata", "chunk_id"]
tags: Optional[List[StrictStr]] = None
__properties: ClassVar[List[str]] = ["id", "text", "type", "entities", "context", "occurred_start", "occurred_end", "mentioned_at", "document_id", "metadata", "chunk_id", "tags"]
model_config = ConfigDict(
populate_by_name=True,
@@ -123,6 +124,11 @@ class RecallResult(BaseModel):
if self.chunk_id is None and "chunk_id" in self.model_fields_set:
_dict['chunk_id'] = None
# set to None if tags (nullable) is None
# and model_fields_set contains the field
if self.tags is None and "tags" in self.model_fields_set:
_dict['tags'] = None
return _dict
@classmethod
@@ -145,7 +151,8 @@ class RecallResult(BaseModel):
"mentioned_at": obj.get("mentioned_at"),
"document_id": obj.get("document_id"),
"metadata": obj.get("metadata"),
"chunk_id": obj.get("chunk_id")
"chunk_id": obj.get("chunk_id"),
"tags": obj.get("tags")
})
return _obj
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr, field_validator
from typing import Any, ClassVar, Dict, List, Optional
from hindsight_client_api.models.budget import Budget
from hindsight_client_api.models.reflect_include_options import ReflectIncludeOptions
@@ -34,7 +34,19 @@ class ReflectRequest(BaseModel):
max_tokens: Optional[StrictInt] = Field(default=4096, description="Maximum tokens for the response")
include: Optional[ReflectIncludeOptions] = Field(default=None, description="Options for including additional data (disabled by default)")
response_schema: Optional[Dict[str, Any]] = None
__properties: ClassVar[List[str]] = ["query", "budget", "context", "max_tokens", "include", "response_schema"]
tags: Optional[List[StrictStr]] = None
tags_match: Optional[StrictStr] = Field(default='any', description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).")
__properties: ClassVar[List[str]] = ["query", "budget", "context", "max_tokens", "include", "response_schema", "tags", "tags_match"]
@field_validator('tags_match')
def tags_match_validate_enum(cls, value):
"""Validates the enum"""
if value is None:
return value
if value not in set(['any', 'all', 'any_strict', 'all_strict']):
raise ValueError("must be one of enum values ('any', 'all', 'any_strict', 'all_strict')")
return value
model_config = ConfigDict(
populate_by_name=True,
@@ -88,6 +100,11 @@ class ReflectRequest(BaseModel):
if self.response_schema is None and "response_schema" in self.model_fields_set:
_dict['response_schema'] = None
# set to None if tags (nullable) is None
# and model_fields_set contains the field
if self.tags is None and "tags" in self.model_fields_set:
_dict['tags'] = None
return _dict
@classmethod
@@ -105,7 +122,9 @@ class ReflectRequest(BaseModel):
"context": obj.get("context"),
"max_tokens": obj.get("max_tokens") if obj.get("max_tokens") is not None else 4096,
"include": ReflectIncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None,
"response_schema": obj.get("response_schema")
"response_schema": obj.get("response_schema"),
"tags": obj.get("tags"),
"tags_match": obj.get("tags_match") if obj.get("tags_match") is not None else 'any'
})
return _obj
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, Field, StrictBool
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictStr
from typing import Any, ClassVar, Dict, List, Optional
from hindsight_client_api.models.memory_item import MemoryItem
from typing import Optional, Set
@@ -29,7 +29,8 @@ class RetainRequest(BaseModel):
""" # noqa: E501
items: List[MemoryItem]
var_async: Optional[StrictBool] = Field(default=False, description="If true, process asynchronously in background. If false, wait for completion (default: false)", alias="async")
__properties: ClassVar[List[str]] = ["items", "async"]
document_tags: Optional[List[StrictStr]] = None
__properties: ClassVar[List[str]] = ["items", "async", "document_tags"]
model_config = ConfigDict(
populate_by_name=True,
@@ -77,6 +78,11 @@ class RetainRequest(BaseModel):
if _item_items:
_items.append(_item_items.to_dict())
_dict['items'] = _items
# set to None if document_tags (nullable) is None
# and model_fields_set contains the field
if self.document_tags is None and "document_tags" in self.model_fields_set:
_dict['document_tags'] = None
return _dict
@classmethod
@@ -90,7 +96,8 @@ class RetainRequest(BaseModel):
_obj = cls.model_validate({
"items": [MemoryItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None,
"async": obj.get("async") if obj.get("async") is not None else False
"async": obj.get("async") if obj.get("async") is not None else False,
"document_tags": obj.get("document_tags")
})
return _obj
@@ -0,0 +1,89 @@
# coding: utf-8
"""
Hindsight HTTP API
HTTP API for Hindsight
The version of the OpenAPI document: 0.1.0
Generated by OpenAPI Generator (https://openapi-generator.tech)
Do not edit the class manually.
""" # noqa: E501
from __future__ import annotations
import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
from typing import Any, ClassVar, Dict, List
from typing import Optional, Set
from typing_extensions import Self
class TagItem(BaseModel):
"""
Single tag with usage count.
""" # noqa: E501
tag: StrictStr = Field(description="The tag value")
count: StrictInt = Field(description="Number of memories with this tag")
__properties: ClassVar[List[str]] = ["tag", "count"]
model_config = ConfigDict(
populate_by_name=True,
validate_assignment=True,
protected_namespaces=(),
)
def to_str(self) -> str:
"""Returns the string representation of the model using alias"""
return pprint.pformat(self.model_dump(by_alias=True))
def to_json(self) -> str:
"""Returns the JSON representation of the model using alias"""
# TODO: pydantic v2: use .model_dump_json(by_alias=True, exclude_unset=True) instead
return json.dumps(self.to_dict())
@classmethod
def from_json(cls, json_str: str) -> Optional[Self]:
"""Create an instance of TagItem from a JSON string"""
return cls.from_dict(json.loads(json_str))
def to_dict(self) -> Dict[str, Any]:
"""Return the dictionary representation of the model using alias.
This has the following differences from calling pydantic's
`self.model_dump(by_alias=True)`:
* `None` is only added to the output dict for nullable fields that
were set at model initialization. Other fields with value `None`
are ignored.
"""
excluded_fields: Set[str] = set([
])
_dict = self.model_dump(
by_alias=True,
exclude=excluded_fields,
exclude_none=True,
)
return _dict
@classmethod
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional[Self]:
"""Create an instance of TagItem from a dict"""
if obj is None:
return None
if not isinstance(obj, dict):
return cls.model_validate(obj)
_obj = cls.model_validate({
"tag": obj.get("tag"),
"count": obj.get("count")
})
return _obj
+7
View File
@@ -64,6 +64,7 @@ mod tests {
metadata: None,
timestamp: None,
entities: None,
tags: None,
},
types::MemoryItem {
content: "Bob works with Alice on the search team".to_string(),
@@ -72,8 +73,10 @@ mod tests {
metadata: None,
timestamp: None,
entities: None,
tags: None,
},
],
document_tags: None,
};
let retain_response = client
.retain_memories(&bank_id, None, &retain_request)
@@ -90,6 +93,8 @@ mod tests {
include: None,
query_timestamp: None,
types: None,
tags: None,
tags_match: types::TagsMatch::Any,
};
let recall_response = client
.recall_memories(&bank_id, None, &recall_request)
@@ -106,6 +111,8 @@ mod tests {
max_tokens: 4096,
include: None,
response_schema: None,
tags: None,
tags_match: types::TagsMatch::Any,
};
let reflect_response = client
.reflect(&bank_id, None, &reflect_request)
@@ -39,6 +39,9 @@ import type {
GetGraphData,
GetGraphErrors,
GetGraphResponses,
GetMemoryData,
GetMemoryErrors,
GetMemoryResponses,
HealthEndpointHealthGetData,
HealthEndpointHealthGetResponses,
ListBanksData,
@@ -56,6 +59,9 @@ import type {
ListOperationsData,
ListOperationsErrors,
ListOperationsResponses,
ListTagsData,
ListTagsErrors,
ListTagsResponses,
MetricsEndpointMetricsGetData,
MetricsEndpointMetricsGetResponses,
RecallMemoriesData,
@@ -148,6 +154,20 @@ export const listMemories = <ThrowOnError extends boolean = false>(
ThrowOnError
>({ url: "/v1/default/banks/{bank_id}/memories/list", ...options });
/**
* Get memory unit
*
* Get a single memory unit by ID with all its metadata including entities and tags.
*/
export const getMemory = <ThrowOnError extends boolean = false>(
options: Options<GetMemoryData, ThrowOnError>,
) =>
(options.client ?? client).get<
GetMemoryResponses,
GetMemoryErrors,
ThrowOnError
>({ url: "/v1/default/banks/{bank_id}/memories/{memory_id}", ...options });
/**
* Recall memory
*
@@ -329,6 +349,20 @@ export const getDocument = <ThrowOnError extends boolean = false>(
ThrowOnError
>({ url: "/v1/default/banks/{bank_id}/documents/{document_id}", ...options });
/**
* List tags
*
* List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
*/
export const listTags = <ThrowOnError extends boolean = false>(
options: Options<ListTagsData, ThrowOnError>,
) =>
(options.client ?? client).get<
ListTagsResponses,
ListTagsErrors,
ThrowOnError
>({ url: "/v1/default/banks/{bank_id}/tags", ...options });
/**
* Get chunk details
*
@@ -377,6 +377,12 @@ export type DocumentResponse = {
* Memory Unit Count
*/
memory_unit_count: number;
/**
* Tags
*
* Tags associated with this document
*/
tags?: Array<string>;
};
/**
@@ -666,6 +672,30 @@ export type ListMemoryUnitsResponse = {
offset: number;
};
/**
* ListTagsResponse
*
* Response model for list tags endpoint.
*/
export type ListTagsResponse = {
/**
* Items
*/
items: Array<TagItem>;
/**
* Total
*/
total: number;
/**
* Limit
*/
limit: number;
/**
* Offset
*/
offset: number;
};
/**
* MemoryItem
*
@@ -702,6 +732,12 @@ export type MemoryItem = {
* Optional entities to combine with auto-extracted entities.
*/
entities?: Array<EntityInput> | null;
/**
* Tags
*
* Optional tags for visibility scoping. Memories with tags can be filtered during recall.
*/
tags?: Array<string> | null;
};
/**
@@ -791,6 +827,18 @@ export type RecallRequest = {
* Options for including additional data (entities are included by default)
*/
include?: IncludeOptions;
/**
* Tags
*
* Filter memories by tags. If not specified, all memories are returned.
*/
tags?: Array<string> | null;
/**
* Tags Match
*
* How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).
*/
tags_match?: "any" | "all" | "any_strict" | "all_strict";
};
/**
@@ -879,6 +927,10 @@ export type RecallResult = {
* Chunk Id
*/
chunk_id?: string | null;
/**
* Tags
*/
tags?: Array<string> | null;
};
/**
@@ -958,6 +1010,18 @@ export type ReflectRequest = {
response_schema?: {
[key: string]: unknown;
} | null;
/**
* Tags
*
* Filter memories by tags during reflection. If not specified, all memories are considered.
*/
tags?: Array<string> | null;
/**
* Tags Match
*
* How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).
*/
tags_match?: "any" | "all" | "any_strict" | "all_strict";
};
/**
@@ -1004,6 +1068,12 @@ export type RetainRequest = {
* If true, process asynchronously in background. If false, wait for completion (default: false)
*/
async?: boolean;
/**
* Document Tags
*
* Tags applied to all items in this request. These are merged with any item-level tags.
*/
document_tags?: Array<string> | null;
};
/**
@@ -1042,6 +1112,26 @@ export type RetainResponse = {
usage?: TokenUsage | null;
};
/**
* TagItem
*
* Single tag with usage count.
*/
export type TagItem = {
/**
* Tag
*
* The tag value
*/
tag: string;
/**
* Count
*
* Number of memories with this tag
*/
count: number;
};
/**
* TokenUsage
*
@@ -1225,6 +1315,44 @@ export type ListMemoriesResponses = {
export type ListMemoriesResponse =
ListMemoriesResponses[keyof ListMemoriesResponses];
export type GetMemoryData = {
body?: never;
headers?: {
/**
* Authorization
*/
authorization?: string | null;
};
path: {
/**
* Bank Id
*/
bank_id: string;
/**
* Memory Id
*/
memory_id: string;
};
query?: never;
url: "/v1/default/banks/{bank_id}/memories/{memory_id}";
};
export type GetMemoryErrors = {
/**
* Validation Error
*/
422: HttpValidationError;
};
export type GetMemoryError = GetMemoryErrors[keyof GetMemoryErrors];
export type GetMemoryResponses = {
/**
* Successful Response
*/
200: unknown;
};
export type RecallMemoriesData = {
body: RecallRequest;
headers?: {
@@ -1632,6 +1760,61 @@ export type GetDocumentResponses = {
export type GetDocumentResponse =
GetDocumentResponses[keyof GetDocumentResponses];
export type ListTagsData = {
body?: never;
headers?: {
/**
* Authorization
*/
authorization?: string | null;
};
path: {
/**
* Bank Id
*/
bank_id: string;
};
query?: {
/**
* Q
*
* Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
*/
q?: string | null;
/**
* Limit
*
* Maximum number of tags to return
*/
limit?: number;
/**
* Offset
*
* Offset for pagination
*/
offset?: number;
};
url: "/v1/default/banks/{bank_id}/tags";
};
export type ListTagsErrors = {
/**
* Validation Error
*/
422: HttpValidationError;
};
export type ListTagsError = ListTagsErrors[keyof ListTagsErrors];
export type ListTagsResponses = {
/**
* Successful Response
*/
200: ListTagsResponse;
};
export type ListTagsResponse2 = ListTagsResponses[keyof ListTagsResponses];
export type GetChunkData = {
body?: never;
headers?: {
+4 -1
View File
@@ -62,6 +62,7 @@ export interface MemoryItemInput {
metadata?: Record<string, string>;
document_id?: string;
entities?: EntityInput[];
tags?: string[];
}
export class HindsightClient {
@@ -142,13 +143,14 @@ export class HindsightClient {
/**
* Retain multiple memories in batch.
*/
async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; async?: boolean }): Promise<RetainResponse> {
async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; documentTags?: string[]; async?: boolean }): Promise<RetainResponse> {
const processedItems = items.map((item) => ({
content: item.content,
context: item.context,
metadata: item.metadata,
document_id: item.document_id,
entities: item.entities,
tags: item.tags,
timestamp:
item.timestamp instanceof Date
? item.timestamp.toISOString()
@@ -166,6 +168,7 @@ export class HindsightClient {
path: { bank_id: bankId },
body: {
items: itemsWithDocId,
document_tags: options?.documentTags,
async: options?.async,
},
});
@@ -0,0 +1,41 @@
import { NextRequest, NextResponse } from "next/server";
const DATAPLANE_URL = process.env.HINDSIGHT_CP_DATAPLANE_API_URL || "http://localhost:8888";
export async function GET(
request: NextRequest,
{ params }: { params: Promise<{ memoryId: string }> }
) {
try {
const { memoryId } = await params;
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get("bank_id");
if (!bankId) {
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const response = await fetch(
`${DATAPLANE_URL}/v1/default/banks/${bankId}/memories/${memoryId}`,
{
method: "GET",
headers: {
"Content-Type": "application/json",
},
}
);
if (!response.ok) {
if (response.status === 404) {
return NextResponse.json({ error: "Memory not found" }, { status: 404 });
}
throw new Error(`API returned ${response.status}`);
}
const data = await response.json();
return NextResponse.json(data, { status: 200 });
} catch (error) {
console.error("Error fetching memory:", error);
return NextResponse.json({ error: "Failed to fetch memory" }, { status: 500 });
}
}
@@ -10,9 +10,12 @@ export async function POST(request: NextRequest) {
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const { items, document_id } = body;
const { items, document_id, document_tags } = body;
const response = await hindsightClient.retainBatch(bankId, items, { documentId: document_id });
const response = await hindsightClient.retainBatch(bankId, items, {
documentId: document_id,
documentTags: document_tags,
});
return NextResponse.json(response, { status: 200 });
} catch (error) {
@@ -5,7 +5,18 @@ export async function POST(request: NextRequest) {
try {
const body = await request.json();
const bankId = body.bank_id || body.agent_id || "default";
const { query, types, fact_type, max_tokens, trace, budget, include, query_timestamp } = body;
const {
query,
types,
fact_type,
max_tokens,
trace,
budget,
include,
query_timestamp,
tags,
tags_match,
} = body;
const response = await sdk.recallMemories({
client: lowLevelClient,
@@ -18,6 +29,8 @@ export async function POST(request: NextRequest) {
budget: budget || "mid",
include,
query_timestamp,
tags,
tags_match,
},
});
@@ -5,12 +5,14 @@ export async function POST(request: NextRequest) {
try {
const body = await request.json();
const bankId = body.bank_id || body.agent_id || "default";
const { query, context, budget, thinking_budget, include_facts } = body;
const { query, context, budget, thinking_budget, include_facts, tags, tags_match } = body;
const requestBody: any = {
query,
budget: budget || (thinking_budget ? "mid" : "low"),
context: context || undefined,
tags,
tags_match,
};
// Add include options if specified
@@ -7,6 +7,7 @@ import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
import { Checkbox } from "@/components/ui/checkbox";
import { Tag } from "lucide-react";
export function AddMemoryView() {
const { currentBank } = useBank();
@@ -14,6 +15,7 @@ export function AddMemoryView() {
const [context, setContext] = useState("");
const [eventDate, setEventDate] = useState("");
const [documentId, setDocumentId] = useState("");
const [tags, setTags] = useState("");
const [async, setAsync] = useState(false);
const [loading, setLoading] = useState(false);
const [result, setResult] = useState<string | null>(null);
@@ -23,6 +25,7 @@ export function AddMemoryView() {
setContext("");
setEventDate("");
setDocumentId("");
setTags("");
setAsync(false);
setResult(null);
};
@@ -37,16 +40,24 @@ export function AddMemoryView() {
setResult(null);
try {
// Parse tags from comma-separated string
const parsedTags = tags
.split(",")
.map((t) => t.trim())
.filter((t) => t.length > 0);
const item: any = { content };
if (context) item.context = context;
// datetime-local gives "2024-01-15T10:30", add seconds for proper ISO format
if (eventDate) item.timestamp = eventDate + ":00";
if (parsedTags.length > 0) item.tags = parsedTags;
const data: any = await client.retain({
bank_id: currentBank,
items: [item],
document_id: documentId,
async,
...(parsedTags.length > 0 && { document_tags: parsedTags }),
});
setResult(data.message as string);
@@ -112,6 +123,22 @@ export function AddMemoryView() {
</small>
</div>
<div className="mb-4">
<label className="font-bold block mb-1 text-card-foreground flex items-center gap-2">
<Tag className="h-4 w-4" />
Tags
</label>
<Input
type="text"
value={tags}
onChange={(e) => setTags(e.target.value)}
placeholder="user_alice, session_123, project_x"
/>
<small className="text-muted-foreground text-xs mt-1 block">
Comma-separated tags for filtering during recall/reflect. Tags cannot contain commas.
</small>
</div>
<div className="mb-4">
<div className="flex items-center gap-2">
<Checkbox
@@ -23,7 +23,7 @@ import {
DialogFooter,
} from "@/components/ui/dialog";
import { Input } from "@/components/ui/input";
import { Check, ChevronsUpDown, Plus, FileText, Moon, Sun, Github } from "lucide-react";
import { Check, ChevronsUpDown, Plus, FileText, Moon, Sun, Github, Tag } from "lucide-react";
import { useTheme } from "@/lib/theme-context";
import Image from "next/image";
import { Textarea } from "@/components/ui/textarea";
@@ -47,6 +47,7 @@ function BankSelectorInner() {
const [docContext, setDocContext] = React.useState("");
const [docEventDate, setDocEventDate] = React.useState("");
const [docDocumentId, setDocDocumentId] = React.useState("");
const [docTags, setDocTags] = React.useState("");
const [docAsync, setDocAsync] = React.useState(false);
const [isCreatingDoc, setIsCreatingDoc] = React.useState(false);
const [docError, setDocError] = React.useState<string | null>(null);
@@ -83,10 +84,17 @@ function BankSelectorInner() {
setDocError(null);
try {
// Parse tags from comma-separated string
const parsedTags = docTags
.split(",")
.map((t) => t.trim())
.filter((t) => t.length > 0);
const item: any = { content: docContent };
if (docContext) item.context = docContext;
// datetime-local gives "2024-01-15T10:30", add seconds for proper ISO format
if (docEventDate) item.timestamp = docEventDate + ":00";
if (parsedTags.length > 0) item.tags = parsedTags;
const params: any = {
bank_id: currentBank,
@@ -94,6 +102,7 @@ function BankSelectorInner() {
};
if (docDocumentId) params.document_id = docDocumentId;
if (parsedTags.length > 0) params.document_tags = parsedTags;
if (docAsync) {
await client.retain({ ...params, async: true });
@@ -107,6 +116,7 @@ function BankSelectorInner() {
setDocContext("");
setDocEventDate("");
setDocDocumentId("");
setDocTags("");
setDocAsync(false);
// Navigate to documents view to see the new document
@@ -335,6 +345,22 @@ function BankSelectorInner() {
</div>
</div>
<div>
<label className="font-bold block mb-1 text-sm text-foreground flex items-center gap-2">
<Tag className="h-4 w-4" />
Tags
</label>
<Input
type="text"
value={docTags}
onChange={(e) => setDocTags(e.target.value)}
placeholder="user_alice, session_123, project_x"
/>
<p className="text-xs text-muted-foreground mt-1">
Comma-separated tags for filtering during recall/reflect
</p>
</div>
<div className="flex items-center gap-2">
<Checkbox
id="async-doc"
@@ -357,6 +383,7 @@ function BankSelectorInner() {
setDocContext("");
setDocEventDate("");
setDocDocumentId("");
setDocTags("");
setDocAsync(false);
setDocError(null);
}}
@@ -362,6 +362,7 @@ export function DataView({ factType }: DataViewProps) {
memory={selectedGraphNode}
onClose={() => setSelectedGraphNode(null)}
inPanel
bankId={currentBank || undefined}
/>
) : (
/* Legend & Controls View */
@@ -738,13 +739,20 @@ export function DataView({ factType }: DataViewProps) {
memory={selectedTableMemory}
onClose={() => setSelectedTableMemory(null)}
inPanel
bankId={currentBank || undefined}
/>
</div>
)}
</div>
)}
{viewMode === "timeline" && <TimelineView data={data} filteredRows={filteredTableRows} />}
{viewMode === "timeline" && (
<TimelineView
data={data}
filteredRows={filteredTableRows}
bankId={currentBank || undefined}
/>
)}
</>
) : (
<div className="flex items-center justify-center py-20">
@@ -761,7 +769,15 @@ export function DataView({ factType }: DataViewProps) {
// Timeline View Component - Custom compact timeline with zoom and navigation
type Granularity = "year" | "month" | "week" | "day";
function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }) {
function TimelineView({
data,
filteredRows,
bankId,
}: {
data: any;
filteredRows: any[];
bankId?: string;
}) {
const [selectedItem, setSelectedItem] = useState<any>(null);
const [granularity, setGranularity] = useState<Granularity>("month");
const [currentIndex, setCurrentIndex] = useState(0);
@@ -1114,7 +1130,12 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
{/* Detail Panel - Fixed on Right */}
{selectedItem && (
<div className="fixed right-0 top-0 h-screen w-[420px] bg-card border-l-2 border-primary shadow-2xl z-50 overflow-y-auto animate-in slide-in-from-right duration-300 ease-out">
<MemoryDetailPanel memory={selectedItem} onClose={() => setSelectedItem(null)} inPanel />
<MemoryDetailPanel
memory={selectedItem}
onClose={() => setSelectedItem(null)}
inPanel
bankId={bankId}
/>
</div>
)}
</div>
@@ -318,6 +318,25 @@ export function DocumentsView() {
</div>
)}
{/* Tags */}
{selectedDocument.tags && selectedDocument.tags.length > 0 && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Tags
</div>
<div className="flex flex-wrap gap-2">
{selectedDocument.tags.map((tag: string, i: number) => (
<span
key={i}
className="text-sm px-3 py-1.5 rounded-full bg-amber-500/10 text-amber-600 dark:text-amber-400 font-medium"
>
{tag}
</span>
))}
</div>
</div>
)}
{/* Delete Button */}
<div className="pt-2 border-t border-border">
<Button
@@ -1,15 +1,17 @@
"use client";
import { useState } from "react";
import { useState, useEffect } from "react";
import { Button } from "@/components/ui/button";
import { Copy, Check, X } from "lucide-react";
import { Copy, Check, X, Loader2 } from "lucide-react";
import { DocumentChunkModal } from "./document-chunk-modal";
import { client } from "@/lib/api";
interface MemoryDetailPanelProps {
memory: any;
onClose: () => void;
compact?: boolean;
inPanel?: boolean;
bankId?: string;
}
export function MemoryDetailPanel({
@@ -17,10 +19,40 @@ export function MemoryDetailPanel({
onClose,
compact = false,
inPanel = false,
bankId,
}: MemoryDetailPanelProps) {
const [copiedId, setCopiedId] = useState<string | null>(null);
const [modalType, setModalType] = useState<"document" | "chunk" | null>(null);
const [modalId, setModalId] = useState<string | null>(null);
const [fullMemory, setFullMemory] = useState<any>(null);
const [loading, setLoading] = useState(false);
// Fetch full memory data when panel opens
useEffect(() => {
const memoryId = memory?.id || memory?.node_id;
if (!memoryId || !bankId) {
setFullMemory(null);
return;
}
setLoading(true);
client
.getMemory(memoryId, bankId)
.then((data) => {
setFullMemory(data);
})
.catch((err) => {
console.error("Failed to fetch memory details:", err);
// Fall back to showing the partial data we have
setFullMemory(null);
})
.finally(() => {
setLoading(false);
});
}, [memory?.id, memory?.node_id, bankId]);
// Use full memory data if available, otherwise fall back to the partial data passed in
const displayMemory = fullMemory || memory;
const copyToClipboard = async (text: string) => {
try {
@@ -50,7 +82,7 @@ export function MemoryDetailPanel({
if (!memory) return null;
// Handle both 'id' and 'node_id' (trace results use node_id)
const memoryId = memory.id || memory.node_id;
const memoryId = displayMemory.id || displayMemory.node_id;
const labelSize = compact ? "text-[10px]" : "text-xs";
const textSize = compact ? "text-xs" : "text-sm";
@@ -71,123 +103,156 @@ export function MemoryDetailPanel({
</Button>
</div>
<div className="space-y-5">
{/* Full Text */}
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Full Text
</div>
<div className="text-sm whitespace-pre-wrap leading-relaxed text-foreground">
{memory.text}
</div>
{loading ? (
<div className="flex items-center justify-center py-12">
<Loader2 className="h-6 w-6 animate-spin text-muted-foreground" />
<span className="ml-2 text-muted-foreground">Loading memory details...</span>
</div>
{/* Context */}
{memory.context && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Context
</div>
<div className="text-sm text-foreground">{memory.context}</div>
</div>
)}
{/* Dates */}
<div className="grid grid-cols-2 gap-4">
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Occurred
</div>
<div className="text-sm font-medium text-foreground">
{memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"}
</div>
</div>
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Mentioned
</div>
<div className="text-sm font-medium text-foreground">
{memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"}
</div>
</div>
</div>
{/* Entities */}
{memory.entities && (
) : (
<div className="space-y-5">
{/* Full Text */}
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">
Entities
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Full Text
</div>
<div className="flex flex-wrap gap-2">
{(Array.isArray(memory.entities)
? memory.entities
: String(memory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string" ? entity : entity?.name || JSON.stringify(entity);
return (
<div className="text-sm whitespace-pre-wrap leading-relaxed text-foreground">
{displayMemory.text}
</div>
</div>
{/* Context */}
{displayMemory.context && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Context
</div>
<div className="text-sm text-foreground">{displayMemory.context}</div>
</div>
)}
{/* Dates */}
<div className="grid grid-cols-2 gap-4">
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Occurred
</div>
<div className="text-sm font-medium text-foreground">
{displayMemory.occurred_start
? new Date(displayMemory.occurred_start).toLocaleString()
: "N/A"}
</div>
</div>
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Mentioned
</div>
<div className="text-sm font-medium text-foreground">
{displayMemory.mentioned_at
? new Date(displayMemory.mentioned_at).toLocaleString()
: "N/A"}
</div>
</div>
</div>
{/* Entities */}
{displayMemory.entities &&
(Array.isArray(displayMemory.entities)
? displayMemory.entities.length > 0
: displayMemory.entities) && (
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">
Entities
</div>
<div className="flex flex-wrap gap-2">
{(Array.isArray(displayMemory.entities)
? displayMemory.entities
: String(displayMemory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string"
? entity
: entity?.name || JSON.stringify(entity);
return (
<span
key={i}
className="text-sm px-3 py-1.5 rounded-full bg-primary/10 text-primary font-medium"
>
{entityText}
</span>
);
})}
</div>
</div>
)}
{/* Tags */}
{displayMemory.tags && displayMemory.tags.length > 0 && (
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">Tags</div>
<div className="flex flex-wrap gap-2">
{displayMemory.tags.map((tag: string, i: number) => (
<span
key={i}
className="text-sm px-3 py-1.5 rounded-full bg-primary/10 text-primary font-medium"
className="text-sm px-3 py-1.5 rounded-full bg-amber-500/10 text-amber-600 dark:text-amber-400 font-medium"
>
{entityText}
{tag}
</span>
);
})}
))}
</div>
</div>
</div>
)}
)}
{/* ID */}
{memoryId && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Memory ID
{/* ID */}
{memoryId && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Memory ID
</div>
<div className="flex items-center gap-2">
<code className="text-xs font-mono break-all flex-1 text-muted-foreground">
{memoryId}
</code>
<Button
variant="ghost"
size="sm"
className="h-8 w-8 p-0 flex-shrink-0"
onClick={() => copyToClipboard(memoryId)}
>
{copiedId === memoryId ? (
<Check className="h-4 w-4 text-green-600" />
) : (
<Copy className="h-4 w-4" />
)}
</Button>
</div>
</div>
<div className="flex items-center gap-2">
<code className="text-xs font-mono break-all flex-1 text-muted-foreground">
{memoryId}
</code>
<Button
variant="ghost"
size="sm"
className="h-8 w-8 p-0 flex-shrink-0"
onClick={() => copyToClipboard(memoryId)}
>
{copiedId === memoryId ? (
<Check className="h-4 w-4 text-green-600" />
) : (
<Copy className="h-4 w-4" />
)}
</Button>
</div>
</div>
)}
)}
{/* Document/Chunk buttons */}
{(memory.document_id || memory.chunk_id) && (
<div className="flex gap-3 pt-2">
{memory.document_id && (
<Button
onClick={() => openDocumentModal(memory.document_id)}
variant="secondary"
className="flex-1"
>
View Document
</Button>
)}
{memory.chunk_id && (
<Button
onClick={() => openChunkModal(memory.chunk_id)}
variant="secondary"
className="flex-1"
>
View Chunk
</Button>
)}
</div>
)}
</div>
{/* Document/Chunk buttons */}
{(displayMemory.document_id || displayMemory.chunk_id) && (
<div className="flex gap-3 pt-2">
{displayMemory.document_id && (
<Button
onClick={() => openDocumentModal(displayMemory.document_id)}
variant="secondary"
className="flex-1"
>
View Document
</Button>
)}
{displayMemory.chunk_id && (
<Button
onClick={() => openChunkModal(displayMemory.chunk_id)}
variant="secondary"
className="flex-1"
>
View Chunk
</Button>
)}
</div>
)}
</div>
)}
</div>
{/* Document/Chunk Modal */}
@@ -225,123 +290,158 @@ export function MemoryDetailPanel({
</Button>
</div>
<div className={gap}>
{/* Full Text */}
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Full Text
</div>
<div className={`${textSize} whitespace-pre-wrap`}>{memory.text}</div>
{loading ? (
<div className="flex items-center justify-center py-8">
<Loader2 className="h-5 w-5 animate-spin text-muted-foreground" />
<span className="ml-2 text-sm text-muted-foreground">Loading...</span>
</div>
{/* Context */}
{memory.context && (
) : (
<div className={gap}>
{/* Full Text */}
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Context
Full Text
</div>
<div className={textSize}>{memory.context}</div>
<div className={`${textSize} whitespace-pre-wrap`}>{displayMemory.text}</div>
</div>
)}
{/* Dates */}
<div className="grid grid-cols-2 gap-2">
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Occurred
{/* Context */}
{displayMemory.context && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Context
</div>
<div className={textSize}>{displayMemory.context}</div>
</div>
<div className={textSize}>
{memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"}
</div>
</div>
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Mentioned
</div>
<div className={textSize}>
{memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"}
</div>
</div>
</div>
)}
{/* Entities */}
{memory.entities && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
Entities
{/* Dates */}
<div className="grid grid-cols-2 gap-2">
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Occurred
</div>
<div className={textSize}>
{displayMemory.occurred_start
? new Date(displayMemory.occurred_start).toLocaleString()
: "N/A"}
</div>
</div>
<div className="flex flex-wrap gap-1">
{(Array.isArray(memory.entities)
? memory.entities
: String(memory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string" ? entity : entity?.name || JSON.stringify(entity);
return (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Mentioned
</div>
<div className={textSize}>
{displayMemory.mentioned_at
? new Date(displayMemory.mentioned_at).toLocaleString()
: "N/A"}
</div>
</div>
</div>
{/* Entities */}
{displayMemory.entities &&
(Array.isArray(displayMemory.entities)
? displayMemory.entities.length > 0
: displayMemory.entities) && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
Entities
</div>
<div className="flex flex-wrap gap-1">
{(Array.isArray(displayMemory.entities)
? displayMemory.entities
: String(displayMemory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string"
? entity
: entity?.name || JSON.stringify(entity);
return (
<span
key={i}
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-secondary text-secondary-foreground`}
>
{entityText}
</span>
);
})}
</div>
</div>
)}
{/* Tags */}
{displayMemory.tags && displayMemory.tags.length > 0 && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
Tags
</div>
<div className="flex flex-wrap gap-1">
{displayMemory.tags.map((tag: string, i: number) => (
<span
key={i}
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-secondary text-secondary-foreground`}
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-amber-500/10 text-amber-600 dark:text-amber-400`}
>
{entityText}
{tag}
</span>
);
})}
))}
</div>
</div>
</div>
)}
)}
{/* ID */}
{memoryId && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Memory ID
{/* ID */}
{memoryId && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Memory ID
</div>
<div className="flex items-center gap-2">
<span className={`${compact ? "text-[10px]" : "text-sm"} font-mono break-all`}>
{memoryId}
</span>
<Button
variant="ghost"
size="sm"
className="h-6 w-6 p-0 flex-shrink-0"
onClick={() => copyToClipboard(memoryId)}
>
{copiedId === memoryId ? (
<Check className="h-3 w-3 text-green-600" />
) : (
<Copy className="h-3 w-3" />
)}
</Button>
</div>
</div>
<div className="flex items-center gap-2">
<span className={`${compact ? "text-[10px]" : "text-sm"} font-mono break-all`}>
{memoryId}
</span>
<Button
variant="ghost"
size="sm"
className="h-6 w-6 p-0 flex-shrink-0"
onClick={() => copyToClipboard(memoryId)}
>
{copiedId === memoryId ? (
<Check className="h-3 w-3 text-green-600" />
) : (
<Copy className="h-3 w-3" />
)}
</Button>
</div>
</div>
)}
)}
{/* Document/Chunk buttons */}
{(memory.document_id || memory.chunk_id) && (
<div className={`flex gap-2 ${compact ? "pt-1" : ""}`}>
{memory.document_id && (
<Button
onClick={() => openDocumentModal(memory.document_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Document
</Button>
)}
{memory.chunk_id && (
<Button
onClick={() => openChunkModal(memory.chunk_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Chunk
</Button>
)}
</div>
)}
</div>
{/* Document/Chunk buttons */}
{(displayMemory.document_id || displayMemory.chunk_id) && (
<div className={`flex gap-2 ${compact ? "pt-1" : ""}`}>
{displayMemory.document_id && (
<Button
onClick={() => openDocumentModal(displayMemory.document_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Document
</Button>
)}
{displayMemory.chunk_id && (
<Button
onClick={() => openChunkModal(displayMemory.chunk_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Chunk
</Button>
)}
</div>
)}
</div>
)}
</div>
{/* Document/Chunk Modal */}
@@ -25,6 +25,8 @@ import {
FileText,
Users,
ArrowDown,
Tag,
Calendar,
} from "lucide-react";
import JsonView from "react18-json-view";
import "react18-json-view/src/style.css";
@@ -32,6 +34,7 @@ import { MemoryDetailPanel } from "./memory-detail-panel";
type FactType = "world" | "experience" | "opinion";
type Budget = "low" | "mid" | "high";
type TagsMatch = "any" | "all" | "any_strict" | "all_strict";
type ViewMode = "results" | "trace" | "json";
export function SearchDebugView() {
@@ -45,6 +48,8 @@ export function SearchDebugView() {
const [queryDate, setQueryDate] = useState("");
const [includeChunks, setIncludeChunks] = useState(false);
const [includeEntities, setIncludeEntities] = useState(false);
const [tags, setTags] = useState("");
const [tagsMatch, setTagsMatch] = useState<TagsMatch>("any");
// Results state
const [results, setResults] = useState<any[] | null>(null);
@@ -83,6 +88,14 @@ export function SearchDebugView() {
const INITIAL_RESULTS_COUNT = 5;
// Helper to find full memory data from results when clicking trace items
const selectMemoryFromTrace = (traceResult: any) => {
const nodeId = traceResult.id || traceResult.node_id;
// Try to find the full result with all metadata
const fullResult = results?.find((r: any) => r.id === nodeId || r.node_id === nodeId);
setSelectedMemory(fullResult || traceResult);
};
const runSearch = async () => {
if (!currentBank) {
alert("Please select a memory bank first");
@@ -99,6 +112,12 @@ export function SearchDebugView() {
setLoading(true);
try {
// Parse tags from comma-separated string
const parsedTags = tags
.split(",")
.map((t) => t.trim())
.filter((t) => t.length > 0);
const requestBody: any = {
bank_id: currentBank,
query: query,
@@ -111,6 +130,7 @@ export function SearchDebugView() {
chunks: includeChunks ? { max_tokens: 8192 } : null,
},
...(queryDate && { query_timestamp: queryDate }),
...(parsedTags.length > 0 && { tags: parsedTags, tags_match: tagsMatch }),
};
const data: any = await client.recall(requestBody);
@@ -246,6 +266,31 @@ export function SearchDebugView() {
</label>
</div>
</div>
{/* Tags Filter */}
<div className="flex items-center gap-4 mt-4 pt-4 border-t">
<Tag className="h-4 w-4 text-muted-foreground" />
<div className="flex-1 max-w-md">
<Input
type="text"
value={tags}
onChange={(e) => setTags(e.target.value)}
placeholder="Filter by tags (comma-separated)"
className="h-8"
/>
</div>
<Select value={tagsMatch} onValueChange={(v) => setTagsMatch(v as TagsMatch)}>
<SelectTrigger className="w-40 h-8">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="any">Any (incl. untagged)</SelectItem>
<SelectItem value="all">All (incl. untagged)</SelectItem>
<SelectItem value="any_strict">Any (strict)</SelectItem>
<SelectItem value="all_strict">All (strict)</SelectItem>
</SelectContent>
</Select>
</div>
</CardContent>
</Card>
@@ -507,9 +552,29 @@ export function SearchDebugView() {
}}
>
<div className="flex items-center justify-between mb-1">
<span className="font-medium text-sm text-foreground capitalize">
{method.method_name}
</span>
<div className="flex items-center gap-2">
<span className="font-medium text-sm text-foreground capitalize">
{method.method_name}
</span>
{/* Show temporal range inline */}
{method.method_name === "temporal" &&
method.metadata?.constraint && (
<span className="flex items-center gap-1 text-[10px] text-muted-foreground">
<Calendar className="h-3 w-3" />
{method.metadata.constraint.start
? new Date(
method.metadata.constraint.start
).toLocaleDateString()
: "any"}
{" → "}
{method.metadata.constraint.end
? new Date(
method.metadata.constraint.end
).toLocaleDateString()
: "any"}
</span>
)}
</div>
{isMethodExpanded ? (
<ChevronDown className="h-3 w-3 text-muted-foreground" />
) : (
@@ -546,7 +611,7 @@ export function SearchDebugView() {
className="p-2 bg-background rounded cursor-pointer hover:bg-muted/50 transition-colors border border-border"
onClick={(e) => {
e.stopPropagation();
setSelectedMemory(r);
selectMemoryFromTrace(r);
}}
>
<div className="flex items-start gap-2">
@@ -684,7 +749,7 @@ export function SearchDebugView() {
className="p-3 bg-muted/30 rounded-lg cursor-pointer hover:bg-muted/50 transition-colors"
onClick={(e) => {
e.stopPropagation();
setSelectedMemory(r);
selectMemoryFromTrace(r);
}}
>
<div className="flex items-start gap-3">
@@ -792,7 +857,7 @@ export function SearchDebugView() {
className="p-3 bg-muted/30 rounded-lg cursor-pointer hover:bg-muted/50 transition-colors"
onClick={(e) => {
e.stopPropagation();
setSelectedMemory(r);
selectMemoryFromTrace(r);
}}
>
<div className="flex items-start gap-3">
@@ -931,6 +996,7 @@ export function SearchDebugView() {
memory={selectedMemory}
onClose={() => setSelectedMemory(null)}
inPanel
bankId={currentBank || undefined}
/>
</div>
)}
@@ -15,10 +15,12 @@ import {
} from "@/components/ui/select";
import { Checkbox } from "@/components/ui/checkbox";
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card";
import { Sparkles, Info } from "lucide-react";
import { Sparkles, Info, Tag } from "lucide-react";
import JsonView from "react18-json-view";
import "react18-json-view/src/style.css";
type TagsMatch = "any" | "all" | "any_strict" | "all_strict";
export function ThinkView() {
const { currentBank } = useBank();
const [query, setQuery] = useState("");
@@ -28,6 +30,8 @@ export function ThinkView() {
const [showRawJson, setShowRawJson] = useState(false);
const [result, setResult] = useState<any>(null);
const [loading, setLoading] = useState(false);
const [tags, setTags] = useState("");
const [tagsMatch, setTagsMatch] = useState<TagsMatch>("any");
const runReflect = async () => {
if (!currentBank || !query) return;
@@ -35,12 +39,19 @@ export function ThinkView() {
setLoading(true);
setShowRawJson(false);
try {
// Parse tags from comma-separated string
const parsedTags = tags
.split(",")
.map((t) => t.trim())
.filter((t) => t.length > 0);
const data: any = await client.reflect({
bank_id: currentBank,
query,
budget,
context: context || undefined,
include_facts: includeFacts,
...(parsedTags.length > 0 && { tags: parsedTags, tags_match: tagsMatch }),
});
setResult(data);
} catch (error) {
@@ -103,6 +114,29 @@ export function ThinkView() {
rows={3}
/>
</div>
<div className="flex items-center gap-4 mt-4 pt-4 border-t">
<Tag className="h-4 w-4 text-muted-foreground" />
<div className="flex-1 max-w-md">
<Input
type="text"
value={tags}
onChange={(e) => setTags(e.target.value)}
placeholder="Filter by tags (comma-separated)"
className="h-8"
/>
</div>
<Select value={tagsMatch} onValueChange={(v) => setTagsMatch(v as TagsMatch)}>
<SelectTrigger className="w-40 h-8">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="any">Any (incl. untagged)</SelectItem>
<SelectItem value="all">All (incl. untagged)</SelectItem>
<SelectItem value="any_strict">Any (strict)</SelectItem>
<SelectItem value="all_strict">All (strict)</SelectItem>
</SelectContent>
</Select>
</div>
</CardContent>
</Card>
+24
View File
@@ -53,6 +53,8 @@ export class ControlPlaneClient {
chunks?: { max_tokens: number } | null;
};
query_timestamp?: string;
tags?: string[];
tags_match?: "any" | "all" | "any_strict" | "all_strict";
}) {
return this.fetchApi("/api/recall", {
method: "POST",
@@ -69,6 +71,8 @@ export class ControlPlaneClient {
budget?: string;
context?: string;
include_facts?: boolean;
tags?: string[];
tags_match?: "any" | "all" | "any_strict" | "all_strict";
}) {
return this.fetchApi("/api/reflect", {
method: "POST",
@@ -209,6 +213,26 @@ export class ControlPlaneClient {
return this.fetchApi(`/api/chunks/${chunkId}`);
}
/**
* Get a single memory by ID
*/
async getMemory(memoryId: string, bankId: string) {
return this.fetchApi<{
id: string;
text: string;
context: string;
date: string;
type: string;
mentioned_at: string | null;
occurred_start: string | null;
occurred_end: string | null;
entities: string[];
document_id: string | null;
chunk_id: string | null;
tags: string[];
}>(`/api/memories/${memoryId}?bank_id=${bankId}`);
}
/**
* Get bank profile
*/
+34
View File
@@ -167,6 +167,39 @@ As facts accumulate about an entity, Hindsight synthesizes **observations** —
---
## Tagging Memories
You can tag memories for filtering during recall—useful when one memory bank serves multiple users but each user should only see relevant memories.
```python
# Tag memories for specific users
client.retain(
bank_id="my-agent",
items=[
{
"content": "Alice prefers morning meetings",
"tags": ["user_alice"]
}
]
)
# Apply tags to all items in a batch
client.retain(
bank_id="my-agent",
document_tags=["session_123", "user_alice"], # Applied to all items
items=[
{"content": "Alice discussed the project timeline"},
{"content": "Alice mentioned she needs help with Python"}
]
)
```
During recall, use `tags_match` to control matching:
- `"any"` (default): OR matching - returns memories where **any** tag overlaps
- `"all"`: AND matching - returns memories containing **all** specified tags
---
## What You Get
After `retain()` completes:
@@ -176,6 +209,7 @@ After `retain()` completes:
- **Knowledge graph** with entity, temporal, semantic, and causal links
- **Temporal grounding** for both historical and recency-based queries
- **Background processing** that generates entity summaries
- **Optional tags** for filtering during recall
All stored in your isolated **memory bank**, ready for `recall()` and `reflect()`.
@@ -134,6 +134,8 @@ Hindsight is built for AI agents, not humans. Traditional search systems return
- `max_tokens`: How much memory content to return (default: 4096 tokens)
- `budget`: Search depth level (low, mid, high)
- `fact_type`: Filter by world, experience, opinion, or all
- `tags`: Filter memories by tags
- `tags_match`: How to match tags - `"any"` for OR (default), `"all"` for AND
### Expanding Context: Chunks and Entity Observations
+370 -1
View File
@@ -249,6 +249,72 @@
}
}
},
"/v1/default/banks/{bank_id}/memories/{memory_id}": {
"get": {
"tags": [
"Memory"
],
"summary": "Get memory unit",
"description": "Get a single memory unit by ID with all its metadata including entities and tags.",
"operationId": "get_memory",
"parameters": [
{
"name": "bank_id",
"in": "path",
"required": true,
"schema": {
"type": "string",
"title": "Bank Id"
}
},
{
"name": "memory_id",
"in": "path",
"required": true,
"schema": {
"type": "string",
"title": "Memory Id"
}
},
{
"name": "authorization",
"in": "header",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Authorization"
}
}
],
"responses": {
"200": {
"description": "Successful Response",
"content": {
"application/json": {
"schema": {}
}
}
},
"422": {
"description": "Validation Error",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
}
}
}
}
},
"/v1/default/banks/{bank_id}/memories/recall": {
"post": {
"tags": [
@@ -944,6 +1010,107 @@
}
}
},
"/v1/default/banks/{bank_id}/tags": {
"get": {
"tags": [
"Memory"
],
"summary": "List tags",
"description": "List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.",
"operationId": "list_tags",
"parameters": [
{
"name": "bank_id",
"in": "path",
"required": true,
"schema": {
"type": "string",
"title": "Bank Id"
}
},
{
"name": "q",
"in": "query",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.",
"title": "Q"
},
"description": "Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive."
},
{
"name": "limit",
"in": "query",
"required": false,
"schema": {
"type": "integer",
"description": "Maximum number of tags to return",
"default": 100,
"title": "Limit"
},
"description": "Maximum number of tags to return"
},
{
"name": "offset",
"in": "query",
"required": false,
"schema": {
"type": "integer",
"description": "Offset for pagination",
"default": 0,
"title": "Offset"
},
"description": "Offset for pagination"
},
{
"name": "authorization",
"in": "header",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Authorization"
}
}
],
"responses": {
"200": {
"description": "Successful Response",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ListTagsResponse"
}
}
}
},
"422": {
"description": "Validation Error",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
}
}
}
}
},
"/v1/default/chunks/{chunk_id}": {
"get": {
"tags": [
@@ -2219,6 +2386,14 @@
"memory_unit_count": {
"type": "integer",
"title": "Memory Unit Count"
},
"tags": {
"items": {
"type": "string"
},
"type": "array",
"title": "Tags",
"description": "Tags associated with this document"
}
},
"type": "object",
@@ -2240,6 +2415,10 @@
"id": "session_1",
"memory_unit_count": 15,
"original_text": "Full document text here...",
"tags": [
"user_a",
"session_123"
],
"updated_at": "2024-01-15T10:30:00Z"
}
},
@@ -2752,6 +2931,57 @@
"total": 150
}
},
"ListTagsResponse": {
"properties": {
"items": {
"items": {
"$ref": "#/components/schemas/TagItem"
},
"type": "array",
"title": "Items"
},
"total": {
"type": "integer",
"title": "Total"
},
"limit": {
"type": "integer",
"title": "Limit"
},
"offset": {
"type": "integer",
"title": "Offset"
}
},
"type": "object",
"required": [
"items",
"total",
"limit",
"offset"
],
"title": "ListTagsResponse",
"description": "Response model for list tags endpoint.",
"example": {
"items": [
{
"count": 42,
"tag": "user:alice"
},
{
"count": 15,
"tag": "user:bob"
},
{
"count": 8,
"tag": "session:abc123"
}
],
"limit": 100,
"offset": 0,
"total": 25
}
},
"MemoryItem": {
"properties": {
"content": {
@@ -2821,6 +3051,21 @@
],
"title": "Entities",
"description": "Optional entities to combine with auto-extracted entities."
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags",
"description": "Optional tags for visibility scoping. Memories with tags can be filtered during recall."
}
},
"type": "object",
@@ -2846,6 +3091,10 @@
"channel": "engineering",
"source": "slack"
},
"tags": [
"user_a",
"user_b"
],
"timestamp": "2024-01-15T10:30:00Z"
}
},
@@ -2999,6 +3248,33 @@
"include": {
"$ref": "#/components/schemas/IncludeOptions",
"description": "Options for including additional data (entities are included by default)"
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags",
"description": "Filter memories by tags. If not specified, all memories are returned."
},
"tags_match": {
"type": "string",
"enum": [
"any",
"all",
"any_strict",
"all_strict"
],
"title": "Tags Match",
"description": "How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
"default": "any"
}
},
"type": "object",
@@ -3017,6 +3293,10 @@
"max_tokens": 4096,
"query": "What did Alice say about machine learning?",
"query_timestamp": "2023-05-30T23:40:00",
"tags": [
"user_a"
],
"tags_match": "any",
"trace": true,
"types": [
"world",
@@ -3238,6 +3518,20 @@
}
],
"title": "Chunk Id"
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags"
}
},
"type": "object",
@@ -3262,6 +3556,10 @@
},
"occurred_end": "2024-01-15T10:30:00Z",
"occurred_start": "2024-01-15T10:30:00Z",
"tags": [
"user_a",
"user_b"
],
"text": "Alice works at Google on the AI team",
"type": "world"
}
@@ -3404,6 +3702,33 @@
],
"title": "Response Schema",
"description": "Optional JSON Schema for structured output. When provided, the response will include a 'structured_output' field with the LLM response parsed according to this schema."
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags",
"description": "Filter memories by tags during reflection. If not specified, all memories are considered."
},
"tags_match": {
"type": "string",
"enum": [
"any",
"all",
"any_strict",
"all_strict"
],
"title": "Tags Match",
"description": "How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
"default": "any"
}
},
"type": "object",
@@ -3437,7 +3762,11 @@
"key_points"
],
"type": "object"
}
},
"tags": [
"user_a"
],
"tags_match": "any"
}
},
"ReflectResponse": {
@@ -3527,6 +3856,21 @@
"title": "Async",
"description": "If true, process asynchronously in background. If false, wait for completion (default: false)",
"default": false
},
"document_tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Document Tags",
"description": "Tags applied to all items in this request. These are merged with any item-level tags."
}
},
"type": "object",
@@ -3537,6 +3881,10 @@
"description": "Request model for retain endpoint.",
"example": {
"async": false,
"document_tags": [
"user_a",
"user_b"
],
"items": [
{
"content": "Alice works at Google",
@@ -3615,6 +3963,27 @@
}
}
},
"TagItem": {
"properties": {
"tag": {
"type": "string",
"title": "Tag",
"description": "The tag value"
},
"count": {
"type": "integer",
"title": "Count",
"description": "Number of memories with this tag"
}
},
"type": "object",
"required": [
"tag",
"count"
],
"title": "TagItem",
"description": "Single tag with usage count."
},
"TokenUsage": {
"properties": {
"input_tokens": {