diff --git a/CLAUDE.md b/CLAUDE.md index eb0ebab5..fd56e53f 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -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) diff --git a/hindsight-api/hindsight_api/alembic/versions/g2a3b4c5d6e7_add_tags_column.py b/hindsight-api/hindsight_api/alembic/versions/g2a3b4c5d6e7_add_tags_column.py new file mode 100644 index 00000000..8c53f683 --- /dev/null +++ b/hindsight-api/hindsight_api/alembic/versions/g2a3b4c5d6e7_add_tags_column.py @@ -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") diff --git a/hindsight-api/hindsight_api/api/http.py b/hindsight-api/hindsight_api/api/http.py index 35058543..8675bd45 100644 --- a/hindsight-api/hindsight_api/api/http.py +++ b/hindsight-api/hindsight_api/api/http.py @@ -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( diff --git a/hindsight-api/hindsight_api/engine/memory_engine.py b/hindsight-api/hindsight_api/engine/memory_engine.py index 6577b6fa..91e5d20b 100644 --- a/hindsight-api/hindsight_api/engine/memory_engine.py +++ b/hindsight-api/hindsight_api/engine/memory_engine.py @@ -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}") diff --git a/hindsight-api/hindsight_api/engine/response_models.py b/hindsight-api/hindsight_api/engine/response_models.py index c9f7568d..4ab4cb85 100644 --- a/hindsight-api/hindsight_api/engine/response_models.py +++ b/hindsight-api/hindsight_api/engine/response_models.py @@ -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): diff --git a/hindsight-api/hindsight_api/engine/retain/fact_extraction.py b/hindsight-api/hindsight_api/engine/retain/fact_extraction.py index dd9b9871..3403db59 100644 --- a/hindsight-api/hindsight_api/engine/retain/fact_extraction.py +++ b/hindsight-api/hindsight_api/engine/retain/fact_extraction.py @@ -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) diff --git a/hindsight-api/hindsight_api/engine/retain/fact_storage.py b/hindsight-api/hindsight_api/engine/retain/fact_storage.py index 9122204b..7db9469c 100644 --- a/hindsight-api/hindsight_api/engine/retain/fact_storage.py +++ b/hindsight-api/hindsight_api/engine/retain/fact_storage.py @@ -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 [], ) diff --git a/hindsight-api/hindsight_api/engine/retain/orchestrator.py b/hindsight-api/hindsight_api/engine/retain/orchestrator.py index ac4fc898..729daa73 100644 --- a/hindsight-api/hindsight_api/engine/retain/orchestrator.py +++ b/hindsight-api/hindsight_api/engine/retain/orchestrator.py @@ -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) diff --git a/hindsight-api/hindsight_api/engine/retain/types.py b/hindsight-api/hindsight_api/engine/retain/types.py index b0bfef17..528a2d02 100644 --- a/hindsight-api/hindsight_api/engine/retain/types.py +++ b/hindsight-api/hindsight_api/engine/retain/types.py @@ -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) diff --git a/hindsight-api/hindsight_api/engine/search/graph_retrieval.py b/hindsight-api/hindsight_api/engine/search/graph_retrieval.py index d0f312cb..60f3169a 100644 --- a/hindsight-api/hindsight_api/engine/search/graph_retrieval.py +++ b/hindsight-api/hindsight_api/engine/search/graph_retrieval.py @@ -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 diff --git a/hindsight-api/hindsight_api/engine/search/link_expansion_retrieval.py b/hindsight-api/hindsight_api/engine/search/link_expansion_retrieval.py index ccd88ef2..76d157b9 100644 --- a/hindsight-api/hindsight_api/engine/search/link_expansion_retrieval.py +++ b/hindsight-api/hindsight_api/engine/search/link_expansion_retrieval.py @@ -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 diff --git a/hindsight-api/hindsight_api/engine/search/mpfp_retrieval.py b/hindsight-api/hindsight_api/engine/search/mpfp_retrieval.py index aa896930..269a0ac4 100644 --- a/hindsight-api/hindsight_api/engine/search/mpfp_retrieval.py +++ b/hindsight-api/hindsight_api/engine/search/mpfp_retrieval.py @@ -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] diff --git a/hindsight-api/hindsight_api/engine/search/retrieval.py b/hindsight-api/hindsight_api/engine/search/retrieval.py index 32f69e43..4adc80cd 100644 --- a/hindsight-api/hindsight_api/engine/search/retrieval.py +++ b/hindsight-api/hindsight_api/engine/search/retrieval.py @@ -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 diff --git a/hindsight-api/hindsight_api/engine/search/tags.py b/hindsight-api/hindsight_api/engine/search/tags.py new file mode 100644 index 00000000..5417a5a9 --- /dev/null +++ b/hindsight-api/hindsight_api/engine/search/tags.py @@ -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 diff --git a/hindsight-api/hindsight_api/engine/search/trace.py b/hindsight-api/hindsight_api/engine/search/trace.py index 19c80638..e0ed874b 100644 --- a/hindsight-api/hindsight_api/engine/search/trace.py +++ b/hindsight-api/hindsight_api/engine/search/trace.py @@ -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): diff --git a/hindsight-api/hindsight_api/engine/search/tracer.py b/hindsight-api/hindsight_api/engine/search/tracer.py index ef695ad0..d5c5016e 100644 --- a/hindsight-api/hindsight_api/engine/search/tracer.py +++ b/hindsight-api/hindsight_api/engine/search/tracer.py @@ -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 diff --git a/hindsight-api/hindsight_api/engine/search/types.py b/hindsight-api/hindsight_api/engine/search/types.py index 34af67c2..b80bb2e3 100644 --- a/hindsight-api/hindsight_api/engine/search/types.py +++ b/hindsight-api/hindsight_api/engine/search/types.py @@ -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, } diff --git a/hindsight-api/tests/test_tags_visibility.py b/hindsight-api/tests/test_tags_visibility.py new file mode 100644 index 00000000..a25bda07 --- /dev/null +++ b/hindsight-api/tests/test_tags_visibility.py @@ -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") diff --git a/hindsight-cli/src/commands/explore.rs b/hindsight-cli/src/commands/explore.rs index 892686ca..157aa1d1 100644 --- a/hindsight-cli/src/commands/explore.rs +++ b/hindsight-cli/src/commands/explore.rs @@ -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) diff --git a/hindsight-cli/src/commands/memory.rs b/hindsight-cli/src/commands/memory.rs index 0b92a5f5..00696b6c 100644 --- a/hindsight-cli/src/commands/memory.rs +++ b/hindsight-cli/src/commands/memory.rs @@ -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); diff --git a/hindsight-clients/python/.openapi-generator/FILES b/hindsight-clients/python/.openapi-generator/FILES index 6ddbc545..74ebf9db 100644 --- a/hindsight-clients/python/.openapi-generator/FILES +++ b/hindsight-clients/python/.openapi-generator/FILES @@ -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 diff --git a/hindsight-clients/python/hindsight_client_api/__init__.py b/hindsight-clients/python/hindsight_client_api/__init__.py index 76992a52..b1d01d46 100644 --- a/hindsight-clients/python/hindsight_client_api/__init__.py +++ b/hindsight-clients/python/hindsight_client_api/__init__.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 diff --git a/hindsight-clients/python/hindsight_client_api/api/memory_api.py b/hindsight-clients/python/hindsight_client_api/api/memory_api.py index f067094e..30b5ce5d 100644 --- a/hindsight-clients/python/hindsight_client_api/api/memory_api.py +++ b/hindsight-clients/python/hindsight_client_api/api/memory_api.py @@ -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, diff --git a/hindsight-clients/python/hindsight_client_api/models/__init__.py b/hindsight-clients/python/hindsight_client_api/models/__init__.py index 3c7ecc33..31fc55d6 100644 --- a/hindsight-clients/python/hindsight_client_api/models/__init__.py +++ b/hindsight-clients/python/hindsight_client_api/models/__init__.py @@ -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 diff --git a/hindsight-clients/python/hindsight_client_api/models/document_response.py b/hindsight-clients/python/hindsight_client_api/models/document_response.py index e5efa476..a9a84471 100644 --- a/hindsight-clients/python/hindsight_client_api/models/document_response.py +++ b/hindsight-clients/python/hindsight_client_api/models/document_response.py @@ -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 diff --git a/hindsight-clients/python/hindsight_client_api/models/list_tags_response.py b/hindsight-clients/python/hindsight_client_api/models/list_tags_response.py new file mode 100644 index 00000000..996769ae --- /dev/null +++ b/hindsight-clients/python/hindsight_client_api/models/list_tags_response.py @@ -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 + + diff --git a/hindsight-clients/python/hindsight_client_api/models/memory_item.py b/hindsight-clients/python/hindsight_client_api/models/memory_item.py index 89b51d05..4dec4e01 100644 --- a/hindsight-clients/python/hindsight_client_api/models/memory_item.py +++ b/hindsight-clients/python/hindsight_client_api/models/memory_item.py @@ -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 diff --git a/hindsight-clients/python/hindsight_client_api/models/recall_request.py b/hindsight-clients/python/hindsight_client_api/models/recall_request.py index 2ed62d02..0eeb1377 100644 --- a/hindsight-clients/python/hindsight_client_api/models/recall_request.py +++ b/hindsight-clients/python/hindsight_client_api/models/recall_request.py @@ -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 diff --git a/hindsight-clients/python/hindsight_client_api/models/recall_result.py b/hindsight-clients/python/hindsight_client_api/models/recall_result.py index 069efe87..1002b89a 100644 --- a/hindsight-clients/python/hindsight_client_api/models/recall_result.py +++ b/hindsight-clients/python/hindsight_client_api/models/recall_result.py @@ -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 diff --git a/hindsight-clients/python/hindsight_client_api/models/reflect_request.py b/hindsight-clients/python/hindsight_client_api/models/reflect_request.py index dca346cb..da38a88c 100644 --- a/hindsight-clients/python/hindsight_client_api/models/reflect_request.py +++ b/hindsight-clients/python/hindsight_client_api/models/reflect_request.py @@ -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 diff --git a/hindsight-clients/python/hindsight_client_api/models/retain_request.py b/hindsight-clients/python/hindsight_client_api/models/retain_request.py index 0e5a9023..e6c9c8e3 100644 --- a/hindsight-clients/python/hindsight_client_api/models/retain_request.py +++ b/hindsight-clients/python/hindsight_client_api/models/retain_request.py @@ -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 diff --git a/hindsight-clients/python/hindsight_client_api/models/tag_item.py b/hindsight-clients/python/hindsight_client_api/models/tag_item.py new file mode 100644 index 00000000..b18676ef --- /dev/null +++ b/hindsight-clients/python/hindsight_client_api/models/tag_item.py @@ -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 + + diff --git a/hindsight-clients/rust/src/lib.rs b/hindsight-clients/rust/src/lib.rs index bd78fcf8..37bd23db 100644 --- a/hindsight-clients/rust/src/lib.rs +++ b/hindsight-clients/rust/src/lib.rs @@ -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) diff --git a/hindsight-clients/typescript/generated/sdk.gen.ts b/hindsight-clients/typescript/generated/sdk.gen.ts index 68cfb145..bd1c2378 100644 --- a/hindsight-clients/typescript/generated/sdk.gen.ts +++ b/hindsight-clients/typescript/generated/sdk.gen.ts @@ -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 >({ 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 = ( + options: Options, +) => + (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 >({ 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 = ( + options: Options, +) => + (options.client ?? client).get< + ListTagsResponses, + ListTagsErrors, + ThrowOnError + >({ url: "/v1/default/banks/{bank_id}/tags", ...options }); + /** * Get chunk details * diff --git a/hindsight-clients/typescript/generated/types.gen.ts b/hindsight-clients/typescript/generated/types.gen.ts index cc453acb..4457d82e 100644 --- a/hindsight-clients/typescript/generated/types.gen.ts +++ b/hindsight-clients/typescript/generated/types.gen.ts @@ -377,6 +377,12 @@ export type DocumentResponse = { * Memory Unit Count */ memory_unit_count: number; + /** + * Tags + * + * Tags associated with this document + */ + tags?: Array; }; /** @@ -666,6 +672,30 @@ export type ListMemoryUnitsResponse = { offset: number; }; +/** + * ListTagsResponse + * + * Response model for list tags endpoint. + */ +export type ListTagsResponse = { + /** + * Items + */ + items: Array; + /** + * 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 | null; + /** + * Tags + * + * Optional tags for visibility scoping. Memories with tags can be filtered during recall. + */ + tags?: Array | 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 | 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 | 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 | 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 | 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?: { diff --git a/hindsight-clients/typescript/src/index.ts b/hindsight-clients/typescript/src/index.ts index 2d91d233..1f0d445c 100644 --- a/hindsight-clients/typescript/src/index.ts +++ b/hindsight-clients/typescript/src/index.ts @@ -62,6 +62,7 @@ export interface MemoryItemInput { metadata?: Record; 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 { + async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; documentTags?: string[]; async?: boolean }): Promise { 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, }, }); diff --git a/hindsight-control-plane/src/app/api/memories/[memoryId]/route.ts b/hindsight-control-plane/src/app/api/memories/[memoryId]/route.ts new file mode 100644 index 00000000..67cbec28 --- /dev/null +++ b/hindsight-control-plane/src/app/api/memories/[memoryId]/route.ts @@ -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 }); + } +} diff --git a/hindsight-control-plane/src/app/api/memories/retain/route.ts b/hindsight-control-plane/src/app/api/memories/retain/route.ts index fcc351dc..68f1febf 100644 --- a/hindsight-control-plane/src/app/api/memories/retain/route.ts +++ b/hindsight-control-plane/src/app/api/memories/retain/route.ts @@ -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) { diff --git a/hindsight-control-plane/src/app/api/recall/route.ts b/hindsight-control-plane/src/app/api/recall/route.ts index c957f007..ce6ec1bc 100644 --- a/hindsight-control-plane/src/app/api/recall/route.ts +++ b/hindsight-control-plane/src/app/api/recall/route.ts @@ -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, }, }); diff --git a/hindsight-control-plane/src/app/api/reflect/route.ts b/hindsight-control-plane/src/app/api/reflect/route.ts index 1f6e6053..93bba07a 100644 --- a/hindsight-control-plane/src/app/api/reflect/route.ts +++ b/hindsight-control-plane/src/app/api/reflect/route.ts @@ -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 diff --git a/hindsight-control-plane/src/components/add-memory-view.tsx b/hindsight-control-plane/src/components/add-memory-view.tsx index 3f26bced..f10d6630 100644 --- a/hindsight-control-plane/src/components/add-memory-view.tsx +++ b/hindsight-control-plane/src/components/add-memory-view.tsx @@ -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(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() { +
+ + setTags(e.target.value)} + placeholder="user_alice, session_123, project_x" + /> + + Comma-separated tags for filtering during recall/reflect. Tags cannot contain commas. + +
+
(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() {
+
+ + setDocTags(e.target.value)} + placeholder="user_alice, session_123, project_x" + /> +

+ Comma-separated tags for filtering during recall/reflect +

+
+
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} />
)} )} - {viewMode === "timeline" && } + {viewMode === "timeline" && ( + + )} ) : (
@@ -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(null); const [granularity, setGranularity] = useState("month"); const [currentIndex, setCurrentIndex] = useState(0); @@ -1114,7 +1130,12 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] } {/* Detail Panel - Fixed on Right */} {selectedItem && (
- setSelectedItem(null)} inPanel /> + setSelectedItem(null)} + inPanel + bankId={bankId} + />
)}
diff --git a/hindsight-control-plane/src/components/documents-view.tsx b/hindsight-control-plane/src/components/documents-view.tsx index 57c019d9..1312672d 100644 --- a/hindsight-control-plane/src/components/documents-view.tsx +++ b/hindsight-control-plane/src/components/documents-view.tsx @@ -318,6 +318,25 @@ export function DocumentsView() { )} + {/* Tags */} + {selectedDocument.tags && selectedDocument.tags.length > 0 && ( +
+
+ Tags +
+
+ {selectedDocument.tags.map((tag: string, i: number) => ( + + {tag} + + ))} +
+
+ )} + {/* Delete Button */}
-
- {/* Full Text */} -
-
- Full Text -
-
- {memory.text} -
+ {loading ? ( +
+ + Loading memory details...
- - {/* Context */} - {memory.context && ( -
-
- Context -
-
{memory.context}
-
- )} - - {/* Dates */} -
-
-
- Occurred -
-
- {memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"} -
-
-
-
- Mentioned -
-
- {memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"} -
-
-
- - {/* Entities */} - {memory.entities && ( + ) : ( +
+ {/* Full Text */}
-
- Entities +
+ Full Text
-
- {(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 ( +
+ {displayMemory.text} +
+
+ + {/* Context */} + {displayMemory.context && ( +
+
+ Context +
+
{displayMemory.context}
+
+ )} + + {/* Dates */} +
+
+
+ Occurred +
+
+ {displayMemory.occurred_start + ? new Date(displayMemory.occurred_start).toLocaleString() + : "N/A"} +
+
+
+
+ Mentioned +
+
+ {displayMemory.mentioned_at + ? new Date(displayMemory.mentioned_at).toLocaleString() + : "N/A"} +
+
+
+ + {/* Entities */} + {displayMemory.entities && + (Array.isArray(displayMemory.entities) + ? displayMemory.entities.length > 0 + : displayMemory.entities) && ( +
+
+ Entities +
+
+ {(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 ( + + {entityText} + + ); + })} +
+
+ )} + + {/* Tags */} + {displayMemory.tags && displayMemory.tags.length > 0 && ( +
+
Tags
+
+ {displayMemory.tags.map((tag: string, i: number) => ( - {entityText} + {tag} - ); - })} + ))} +
-
- )} + )} - {/* ID */} - {memoryId && ( -
-
- Memory ID + {/* ID */} + {memoryId && ( +
+
+ Memory ID +
+
+ + {memoryId} + + +
-
- - {memoryId} - - -
-
- )} + )} - {/* Document/Chunk buttons */} - {(memory.document_id || memory.chunk_id) && ( -
- {memory.document_id && ( - - )} - {memory.chunk_id && ( - - )} -
- )} -
+ {/* Document/Chunk buttons */} + {(displayMemory.document_id || displayMemory.chunk_id) && ( +
+ {displayMemory.document_id && ( + + )} + {displayMemory.chunk_id && ( + + )} +
+ )} +
+ )}
{/* Document/Chunk Modal */} @@ -225,123 +290,158 @@ export function MemoryDetailPanel({
-
- {/* Full Text */} -
-
- Full Text -
-
{memory.text}
+ {loading ? ( +
+ + Loading...
- - {/* Context */} - {memory.context && ( + ) : ( +
+ {/* Full Text */}
- Context + Full Text
-
{memory.context}
+
{displayMemory.text}
- )} - {/* Dates */} -
-
-
- Occurred + {/* Context */} + {displayMemory.context && ( +
+
+ Context +
+
{displayMemory.context}
-
- {memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"} -
-
-
-
- Mentioned -
-
- {memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"} -
-
-
+ )} - {/* Entities */} - {memory.entities && ( -
-
- Entities + {/* Dates */} +
+
+
+ Occurred +
+
+ {displayMemory.occurred_start + ? new Date(displayMemory.occurred_start).toLocaleString() + : "N/A"} +
-
- {(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 ( +
+
+ Mentioned +
+
+ {displayMemory.mentioned_at + ? new Date(displayMemory.mentioned_at).toLocaleString() + : "N/A"} +
+
+
+ + {/* Entities */} + {displayMemory.entities && + (Array.isArray(displayMemory.entities) + ? displayMemory.entities.length > 0 + : displayMemory.entities) && ( +
+
+ Entities +
+
+ {(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 ( + + {entityText} + + ); + })} +
+
+ )} + + {/* Tags */} + {displayMemory.tags && displayMemory.tags.length > 0 && ( +
+
+ Tags +
+
+ {displayMemory.tags.map((tag: string, i: number) => ( - {entityText} + {tag} - ); - })} + ))} +
-
- )} + )} - {/* ID */} - {memoryId && ( -
-
- Memory ID + {/* ID */} + {memoryId && ( +
+
+ Memory ID +
+
+ + {memoryId} + + +
-
- - {memoryId} - - -
-
- )} + )} - {/* Document/Chunk buttons */} - {(memory.document_id || memory.chunk_id) && ( -
- {memory.document_id && ( - - )} - {memory.chunk_id && ( - - )} -
- )} -
+ {/* Document/Chunk buttons */} + {(displayMemory.document_id || displayMemory.chunk_id) && ( +
+ {displayMemory.document_id && ( + + )} + {displayMemory.chunk_id && ( + + )} +
+ )} +
+ )}
{/* Document/Chunk Modal */} diff --git a/hindsight-control-plane/src/components/search-debug-view.tsx b/hindsight-control-plane/src/components/search-debug-view.tsx index d0c12825..d63711b5 100644 --- a/hindsight-control-plane/src/components/search-debug-view.tsx +++ b/hindsight-control-plane/src/components/search-debug-view.tsx @@ -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("any"); // Results state const [results, setResults] = useState(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() {
+ + {/* Tags Filter */} +
+ +
+ setTags(e.target.value)} + placeholder="Filter by tags (comma-separated)" + className="h-8" + /> +
+ +
@@ -507,9 +552,29 @@ export function SearchDebugView() { }} >
- - {method.method_name} - +
+ + {method.method_name} + + {/* Show temporal range inline */} + {method.method_name === "temporal" && + method.metadata?.constraint && ( + + + {method.metadata.constraint.start + ? new Date( + method.metadata.constraint.start + ).toLocaleDateString() + : "any"} + {" → "} + {method.metadata.constraint.end + ? new Date( + method.metadata.constraint.end + ).toLocaleDateString() + : "any"} + + )} +
{isMethodExpanded ? ( ) : ( @@ -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); }} >
@@ -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); }} >
@@ -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); }} >
@@ -931,6 +996,7 @@ export function SearchDebugView() { memory={selectedMemory} onClose={() => setSelectedMemory(null)} inPanel + bankId={currentBank || undefined} />
)} diff --git a/hindsight-control-plane/src/components/think-view.tsx b/hindsight-control-plane/src/components/think-view.tsx index 3ccfd912..e7def47b 100644 --- a/hindsight-control-plane/src/components/think-view.tsx +++ b/hindsight-control-plane/src/components/think-view.tsx @@ -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(null); const [loading, setLoading] = useState(false); + const [tags, setTags] = useState(""); + const [tagsMatch, setTagsMatch] = useState("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} />
+
+ +
+ setTags(e.target.value)} + placeholder="Filter by tags (comma-separated)" + className="h-8" + /> +
+ +
diff --git a/hindsight-control-plane/src/lib/api.ts b/hindsight-control-plane/src/lib/api.ts index f90dae1b..c83dbc68 100644 --- a/hindsight-control-plane/src/lib/api.ts +++ b/hindsight-control-plane/src/lib/api.ts @@ -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 */ diff --git a/hindsight-docs/docs/developer/retain.md b/hindsight-docs/docs/developer/retain.md index 559d9359..366aa3c7 100644 --- a/hindsight-docs/docs/developer/retain.md +++ b/hindsight-docs/docs/developer/retain.md @@ -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()`. diff --git a/hindsight-docs/docs/developer/retrieval.md b/hindsight-docs/docs/developer/retrieval.md index 4ccaed69..891465ef 100644 --- a/hindsight-docs/docs/developer/retrieval.md +++ b/hindsight-docs/docs/developer/retrieval.md @@ -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 diff --git a/hindsight-docs/static/openapi.json b/hindsight-docs/static/openapi.json index 546523a2..a42318f7 100644 --- a/hindsight-docs/static/openapi.json +++ b/hindsight-docs/static/openapi.json @@ -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": {