Compare commits

...
13 Commits
Author SHA1 Message Date
Nicolò Boschi 2795e98699 fixes 2026-01-16 11:04:17 +01:00
Nicolò Boschi 4221814de7 fix 2026-01-16 10:12:52 +01:00
Nicolò Boschi 6859b5a60e fix 2026-01-16 09:59:36 +01:00
Nicolò Boschi c9c949b34f doc: refinement for 0.3.0 new features 2026-01-14 09:03:42 +01:00
Chris Bartholomew 70ce979fbe doc: update expired Slack invite link (#157) 2026-01-13 16:57:23 -05:00
Nicolò Boschi de132501c6 doc: changelog for 0.3.0 (#156) 2026-01-13 19:09:13 +01:00
Nicolò Boschi a75dcfebf5 Release v0.3.0
- Update version to 0.3.0 in all components
- Python packages: hindsight-api, hindsight-dev, hindsight-all, hindsight-litellm, hindsight-embed
- Python client: hindsight-clients/python
- TypeScript client: hindsight-clients/typescript
- Rust CLI: hindsight-cli
- Control Plane: hindsight-control-plane
- Helm chart
2026-01-13 18:43:33 +01:00
Nicolò Boschi 20c8f8b06a feat: add memory tags (#152)
* feat: add memory tags

* feat: add memory tags

* support tags

* support tags
2026-01-13 18:28:40 +01:00
Chris Bartholomew f5f3fca4ad Fix: Load extensions in server.py for multi-worker deployments (#155)
* Fix: Load extensions in server.py for multi-worker deployments

When running with multiple workers (--workers 2), uvicorn uses
`hindsight_api.server:app` import string instead of passing an app
object. The server.py module was not loading tenant/operation validator
extensions, causing authentication bypass in production.

This fix:
- Adds extension loading to server.py matching main.py behavior
- Sets extension context on tenant extension for schema provisioning
- Adds comprehensive unit tests for server.py extension loading

The tests specifically verify:
- TENANT extension is loaded when HINDSIGHT_API_TENANT_EXTENSION is set
- OPERATION_VALIDATOR is loaded when configured
- Extensions are passed to MemoryEngine constructor
- Extension context is set on tenant extension
- Server works correctly without extensions configured

* Add unit tests for main.py extension loading (single-worker path)
2026-01-13 17:55:33 +01:00
Nicolò Boschi d47c8a28cc feat: support litellm gateway (#154) 2026-01-13 16:55:28 +01:00
Nicolò Boschi 1ffc2a418c feat: add tenant to metrics labels (#151) 2026-01-13 15:31:58 +01:00
Nicolò Boschi fa53917c63 feat: support custom url for openai embeddings & cohere (#150)
* feat: support custom url for openai embeddings & cohere

* feat: support custom url for openai embeddings & cohere
2026-01-13 14:01:44 +01:00
Nicolò Boschi 59913086be fix: batch queries on recall (#149)
* fix: batch queries on recall

* fix: batch queries on recall
2026-01-13 13:20:22 +01:00
87 changed files with 6399 additions and 443 deletions
+19
View File
@@ -174,6 +174,25 @@ This runs the same checks as the pre-commit hook (Ruff for Python, ESLint/Pretti
- Multi-bank queries are client responsibility to orchestrate
- Disposition traits only affect reflect, not recall
### Control Plane API Routes
When adding or modifying parameters in the dataplane API (hindsight-api), you must also update the control plane routes that proxy to it:
1. **API Routes** (`hindsight-control-plane/src/app/api/`):
- `recall/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/recall`
- `reflect/route.ts` - proxies to `/v1/default/banks/{bank_id}/reflect`
- `memories/retain/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/retain`
- Other routes follow the same pattern
2. **Client types** (`hindsight-control-plane/src/lib/api.ts`):
- Update the TypeScript type definitions for `recall()`, `reflect()`, `retain()` etc.
3. **Checklist when adding new API parameters**:
- Add parameter extraction in the route handler (destructure from `body`)
- Pass the parameter to the SDK call
- Update the client type definition in `lib/api.ts`
- Update any UI components that need to use the new parameter
### Python Style
- Python 3.11+, type hints required
- Async throughout (asyncpg, async FastAPI)
+2 -2
View File
@@ -5,7 +5,7 @@
[Documentation](https://hindsight.vectorize.io) • [Paper](https://arxiv.org/abs/2512.12818) • [Cookbook](https://hindsight.vectorize.io/cookbook) • [Hindsight Cloud](https://vectorize.io/hindsight/cloud)
[![CI](https://github.com/vectorize-io/hindsight/actions/workflows/release.yml/badge.svg)](https://github.com/vectorize-io/hindsight/actions/workflows/release.yml)
[![Slack Community](https://img.shields.io/badge/Slack-Join%20Community-4A154B?logo=slack)](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
[![Slack Community](https://img.shields.io/badge/Slack-Join%20Community-4A154B?logo=slack)](https://join.slack.com/t/hindsight-space/shared_invite/zt-3nhbm4w29-LeSJ5Ixi6j8PdiYOCPlOgg)
[![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
![PyPI - Downloads](https://img.shields.io/pypi/dm/hindsight-api?label=PyPI)
![NPM Downloads](https://img.shields.io/npm/dm/%40vectorize-io%2Fhindsight-client?logoColor=orange&label=NPM&color=blue&link=https%3A%2F%2Fwww.npmjs.com%2Fpackage%2F%40vectorize-io%2Fhindsight-client)
@@ -242,7 +242,7 @@ client.reflect(bank_id="my-bank", query="What should I know about Alice?")
- [CLI](https://hindsight.vectorize.io/sdks/cli)
**Community:**
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3nhbm4w29-LeSJ5Ixi6j8PdiYOCPlOgg)
- [GitHub Issues](https://github.com/vectorize-io/hindsight/issues)
---
+2 -2
View File
@@ -2,8 +2,8 @@ apiVersion: v2
name: hindsight
description: Hindsight helm chart
type: application
version: 0.2.1
appVersion: "0.2.1"
version: 0.3.0
appVersion: "0.3.0"
keywords:
- ai
- memory
@@ -0,0 +1,48 @@
"""add_tags_column
Revision ID: g2a3b4c5d6e7
Revises: f1a2b3c4d5e6
Create Date: 2025-01-13
Add tags column to memory_units and documents tables for visibility scoping.
Tags enable filtering memories by scope (e.g., user IDs, session IDs) during recall/reflect.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "g2a3b4c5d6e7"
down_revision: str | Sequence[str] | None = "f1a2b3c4d5e6"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _get_schema_prefix() -> str:
"""Get schema prefix for table names (e.g., 'tenant_x.' or '' for public)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
"""Add tags column to memory_units and documents tables."""
schema = _get_schema_prefix()
# Add tags column to memory_units table
op.execute(f"ALTER TABLE {schema}memory_units ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
# Create GIN index for efficient array containment queries (tags && ARRAY['x'])
op.execute(f"CREATE INDEX IF NOT EXISTS idx_memory_units_tags ON {schema}memory_units USING GIN (tags)")
# Add tags column to documents table for document-level tags
op.execute(f"ALTER TABLE {schema}documents ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
def downgrade() -> None:
"""Remove tags columns and index."""
schema = _get_schema_prefix()
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_tags")
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS tags")
op.execute(f"ALTER TABLE {schema}documents DROP COLUMN IF EXISTS tags")
+167 -2
View File
@@ -37,6 +37,7 @@ from hindsight_api import MemoryEngine
from hindsight_api.engine.db_utils import acquire_with_retry
from hindsight_api.engine.memory_engine import Budget, fq_table
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES, TokenUsage
from hindsight_api.engine.search.tags import TagsMatch
from hindsight_api.extensions import HttpExtension, OperationValidationError, load_extension
from hindsight_api.metrics import create_metrics_collector, get_metrics_collector, initialize_metrics
from hindsight_api.models import RequestContext
@@ -81,6 +82,8 @@ class RecallRequest(BaseModel):
"trace": True,
"query_timestamp": "2023-05-30T23:40:00",
"include": {"entities": {"max_tokens": 500}},
"tags": ["user_a"],
"tags_match": "any",
}
}
)
@@ -99,6 +102,15 @@ class RecallRequest(BaseModel):
default_factory=IncludeOptions,
description="Options for including additional data (entities are included by default)",
)
tags: list[str] | None = Field(
default=None,
description="Filter memories by tags. If not specified, all memories are returned.",
)
tags_match: TagsMatch = Field(
default="any",
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
)
class RecallResult(BaseModel):
@@ -119,6 +131,7 @@ class RecallResult(BaseModel):
"document_id": "session_abc123",
"metadata": {"source": "slack"},
"chunk_id": "456e7890-e12b-34d5-a678-901234567890",
"tags": ["user_a", "user_b"],
}
},
}
@@ -134,6 +147,7 @@ class RecallResult(BaseModel):
document_id: str | None = None # Document this memory belongs to
metadata: dict[str, str] | None = None # User-defined metadata
chunk_id: str | None = None # Chunk this fact was extracted from
tags: list[str] | None = None # Visibility scope tags
class EntityObservationResponse(BaseModel):
@@ -306,6 +320,7 @@ class MemoryItem(BaseModel):
"metadata": {"source": "slack", "channel": "engineering"},
"document_id": "meeting_notes_2024_01_15",
"entities": [{"text": "Alice"}, {"text": "ML model", "type": "CONCEPT"}],
"tags": ["user_a", "user_b"],
}
},
)
@@ -319,6 +334,10 @@ class MemoryItem(BaseModel):
default=None,
description="Optional entities to combine with auto-extracted entities.",
)
tags: list[str] | None = Field(
default=None,
description="Optional tags for visibility scoping. Memories with tags can be filtered during recall.",
)
@field_validator("timestamp", mode="before")
@classmethod
@@ -353,6 +372,7 @@ class RetainRequest(BaseModel):
},
],
"async": False,
"document_tags": ["user_a", "user_b"],
}
}
)
@@ -363,6 +383,10 @@ class RetainRequest(BaseModel):
alias="async",
description="If true, process asynchronously in background. If false, wait for completion (default: false)",
)
document_tags: list[str] | None = Field(
default=None,
description="Tags applied to all items in this request. These are merged with any item-level tags.",
)
class RetainResponse(BaseModel):
@@ -431,6 +455,8 @@ class ReflectRequest(BaseModel):
},
"required": ["summary", "key_points"],
},
"tags": ["user_a"],
"tags_match": "any",
}
}
)
@@ -446,6 +472,15 @@ class ReflectRequest(BaseModel):
default=None,
description="Optional JSON Schema for structured output. When provided, the response will include a 'structured_output' field with the LLM response parsed according to this schema.",
)
tags: list[str] | None = Field(
default=None,
description="Filter memories by tags during reflection. If not specified, all memories are considered.",
)
tags_match: TagsMatch = Field(
default="any",
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
)
class OpinionItem(BaseModel):
@@ -728,6 +763,37 @@ class ListDocumentsResponse(BaseModel):
offset: int
class TagItem(BaseModel):
"""Single tag with usage count."""
tag: str = Field(description="The tag value")
count: int = Field(description="Number of memories with this tag")
class ListTagsResponse(BaseModel):
"""Response model for list tags endpoint."""
model_config = ConfigDict(
json_schema_extra={
"example": {
"items": [
{"tag": "user:alice", "count": 42},
{"tag": "user:bob", "count": 15},
{"tag": "session:abc123", "count": 8},
],
"total": 25,
"limit": 100,
"offset": 0,
}
}
)
items: list[TagItem]
total: int
limit: int
offset: int
class DocumentResponse(BaseModel):
"""Response model for get document endpoint."""
@@ -741,6 +807,7 @@ class DocumentResponse(BaseModel):
"created_at": "2024-01-15T10:30:00Z",
"updated_at": "2024-01-15T10:30:00Z",
"memory_unit_count": 15,
"tags": ["user_a", "session_123"],
}
}
)
@@ -752,6 +819,7 @@ class DocumentResponse(BaseModel):
created_at: str
updated_at: str
memory_unit_count: int
tags: list[str] = Field(default_factory=list, description="Tags associated with this document")
class DeleteDocumentResponse(BaseModel):
@@ -1179,6 +1247,37 @@ def _register_routes(app: FastAPI):
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/list: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.get(
"/v1/default/banks/{bank_id}/memories/{memory_id}",
summary="Get memory unit",
description="Get a single memory unit by ID with all its metadata including entities and tags.",
operation_id="get_memory",
tags=["Memory"],
)
async def api_get_memory(
bank_id: str,
memory_id: str,
request_context: RequestContext = Depends(get_request_context),
):
"""Get a single memory unit by ID."""
try:
data = await app.state.memory.get_memory_unit(
bank_id=bank_id,
memory_id=memory_id,
request_context=request_context,
)
if data is None:
raise HTTPException(status_code=404, detail=f"Memory unit '{memory_id}' not found")
return data
except (AuthenticationError, HTTPException):
raise
except Exception as e:
import traceback
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/{memory_id}: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.post(
"/v1/default/banks/{bank_id}/memories/recall",
response_model=RecallResponse,
@@ -1243,6 +1342,8 @@ def _register_routes(app: FastAPI):
include_chunks=include_chunks,
max_chunk_tokens=max_chunk_tokens,
request_context=request_context,
tags=request.tags,
tags_match=request.tags_match,
)
# Convert core MemoryFact objects to API RecallResult objects (excluding internal metrics)
@@ -1258,6 +1359,7 @@ def _register_routes(app: FastAPI):
mentioned_at=fact.mentioned_at,
document_id=fact.document_id,
chunk_id=fact.chunk_id,
tags=fact.tags,
)
for fact in core_result.results
]
@@ -1350,6 +1452,8 @@ def _register_routes(app: FastAPI):
max_tokens=request.max_tokens,
response_schema=request.response_schema,
request_context=request_context,
tags=request.tags,
tags_match=request.tags_match,
)
# Convert core MemoryFact objects to API ReflectFact objects if facts are requested
@@ -1734,6 +1838,59 @@ def _register_routes(app: FastAPI):
logger.error(f"Error in /v1/default/banks/{bank_id}/documents/{document_id}: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.get(
"/v1/default/banks/{bank_id}/tags",
response_model=ListTagsResponse,
summary="List tags",
description="List all unique tags in a memory bank with usage counts. "
"Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.",
operation_id="list_tags",
tags=["Memory"],
)
async def api_list_tags(
bank_id: str,
q: str | None = Query(
default=None,
description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). "
"Use '*' as wildcard. Case-insensitive.",
),
limit: int = Query(default=100, description="Maximum number of tags to return"),
offset: int = Query(default=0, description="Offset for pagination"),
request_context: RequestContext = Depends(get_request_context),
):
"""
List all unique tags in a memory bank.
Use this endpoint to discover available tags or expand wildcard patterns.
Supports '*' wildcards for flexible matching (case-insensitive):
- 'user:*' matches user:alice, user:bob
- '*-admin' matches role-admin, super-admin
- 'env*-prod' matches env-prod, environment-prod
Args:
bank_id: Memory Bank ID (from path)
q: Wildcard pattern to filter tags (use '*' as wildcard)
limit: Maximum number of tags to return (default: 100)
offset: Offset for pagination (default: 0)
"""
try:
data = await app.state.memory.list_tags(
bank_id=bank_id,
pattern=q,
limit=limit,
offset=offset,
request_context=request_context,
)
return data
except (AuthenticationError, HTTPException):
raise
except Exception as e:
import traceback
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
logger.error(f"Error in /v1/default/banks/{bank_id}/tags: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.get(
"/v1/default/chunks/{chunk_id:path}",
response_model=ChunkResponse,
@@ -2096,11 +2253,15 @@ def _register_routes(app: FastAPI):
content_dict["document_id"] = item.document_id
if item.entities:
content_dict["entities"] = [{"text": e.text, "type": e.type or "CONCEPT"} for e in item.entities]
if item.tags:
content_dict["tags"] = item.tags
contents.append(content_dict)
if request.async_:
# Async processing: queue task and return immediately
result = await app.state.memory.submit_async_retain(bank_id, contents, request_context=request_context)
result = await app.state.memory.submit_async_retain(
bank_id, contents, document_tags=request.document_tags, request_context=request_context
)
return RetainResponse.model_validate(
{
"success": True,
@@ -2114,7 +2275,11 @@ def _register_routes(app: FastAPI):
# Synchronous processing: wait for completion (record metrics)
with metrics.record_operation("retain", bank_id=bank_id, source="api"):
result, usage = await app.state.memory.retain_batch_async(
bank_id=bank_id, contents=contents, request_context=request_context, return_usage=True
bank_id=bank_id,
contents=contents,
document_tags=request.document_tags,
request_context=request_context,
return_usage=True,
)
return RetainResponse.model_validate(
+26
View File
@@ -41,10 +41,19 @@ ENV_EMBEDDINGS_LOCAL_MODEL = "HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL"
ENV_EMBEDDINGS_TEI_URL = "HINDSIGHT_API_EMBEDDINGS_TEI_URL"
ENV_EMBEDDINGS_OPENAI_API_KEY = "HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY"
ENV_EMBEDDINGS_OPENAI_MODEL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL"
ENV_EMBEDDINGS_OPENAI_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL"
ENV_COHERE_API_KEY = "HINDSIGHT_API_COHERE_API_KEY"
ENV_EMBEDDINGS_COHERE_MODEL = "HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL"
ENV_EMBEDDINGS_COHERE_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL"
ENV_RERANKER_COHERE_MODEL = "HINDSIGHT_API_RERANKER_COHERE_MODEL"
ENV_RERANKER_COHERE_BASE_URL = "HINDSIGHT_API_RERANKER_COHERE_BASE_URL"
# LiteLLM gateway configuration (for embeddings and reranker via LiteLLM proxy)
ENV_LITELLM_API_BASE = "HINDSIGHT_API_LITELLM_API_BASE"
ENV_LITELLM_API_KEY = "HINDSIGHT_API_LITELLM_API_KEY"
ENV_EMBEDDINGS_LITELLM_MODEL = "HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL"
ENV_RERANKER_LITELLM_MODEL = "HINDSIGHT_API_RERANKER_LITELLM_MODEL"
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
@@ -64,6 +73,7 @@ ENV_MCP_ENABLED = "HINDSIGHT_API_MCP_ENABLED"
ENV_GRAPH_RETRIEVER = "HINDSIGHT_API_GRAPH_RETRIEVER"
ENV_MPFP_TOP_K_NEIGHBORS = "HINDSIGHT_API_MPFP_TOP_K_NEIGHBORS"
ENV_RECALL_MAX_CONCURRENT = "HINDSIGHT_API_RECALL_MAX_CONCURRENT"
ENV_RECALL_CONNECTION_BUDGET = "HINDSIGHT_API_RECALL_CONNECTION_BUDGET"
ENV_MCP_LOCAL_BANK_ID = "HINDSIGHT_API_MCP_LOCAL_BANK_ID"
ENV_MCP_INSTRUCTIONS = "HINDSIGHT_API_MCP_INSTRUCTIONS"
@@ -120,6 +130,11 @@ DEFAULT_RERANKER_FLASHRANK_CACHE_DIR = None # Use default cache directory
DEFAULT_EMBEDDINGS_COHERE_MODEL = "embed-english-v3.0"
DEFAULT_RERANKER_COHERE_MODEL = "rerank-english-v3.0"
# LiteLLM defaults
DEFAULT_LITELLM_API_BASE = "http://localhost:4000"
DEFAULT_EMBEDDINGS_LITELLM_MODEL = "text-embedding-3-small"
DEFAULT_RERANKER_LITELLM_MODEL = "cohere/rerank-english-v3.0"
DEFAULT_HOST = "0.0.0.0"
DEFAULT_PORT = 8888
DEFAULT_LOG_LEVEL = "info"
@@ -128,6 +143,7 @@ DEFAULT_MCP_ENABLED = True
DEFAULT_GRAPH_RETRIEVER = "link_expansion" # Options: "link_expansion", "mpfp", "bfs"
DEFAULT_MPFP_TOP_K_NEIGHBORS = 20 # Fan-out limit per node in MPFP graph traversal
DEFAULT_RECALL_MAX_CONCURRENT = 32 # Max concurrent recall operations per worker
DEFAULT_RECALL_CONNECTION_BUDGET = 4 # Max concurrent DB connections per recall operation
DEFAULT_MCP_LOCAL_BANK_ID = "mcp"
# Observation thresholds
@@ -222,6 +238,8 @@ class HindsightConfig:
embeddings_provider: str
embeddings_local_model: str
embeddings_tei_url: str | None
embeddings_openai_base_url: str | None
embeddings_cohere_base_url: str | None
# Reranker
reranker_provider: str
@@ -230,6 +248,7 @@ class HindsightConfig:
reranker_tei_batch_size: int
reranker_tei_max_concurrent: int
reranker_max_candidates: int
reranker_cohere_base_url: str | None
# Server
host: str
@@ -241,6 +260,7 @@ class HindsightConfig:
graph_retriever: str
mpfp_top_k_neighbors: int
recall_max_concurrent: int
recall_connection_budget: int
# Observation thresholds
observation_min_facts: int
@@ -297,6 +317,8 @@ class HindsightConfig:
embeddings_provider=os.getenv(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER),
embeddings_local_model=os.getenv(ENV_EMBEDDINGS_LOCAL_MODEL, DEFAULT_EMBEDDINGS_LOCAL_MODEL),
embeddings_tei_url=os.getenv(ENV_EMBEDDINGS_TEI_URL),
embeddings_openai_base_url=os.getenv(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None,
embeddings_cohere_base_url=os.getenv(ENV_EMBEDDINGS_COHERE_BASE_URL) or None,
# Reranker
reranker_provider=os.getenv(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER),
reranker_local_model=os.getenv(ENV_RERANKER_LOCAL_MODEL, DEFAULT_RERANKER_LOCAL_MODEL),
@@ -306,6 +328,7 @@ class HindsightConfig:
os.getenv(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT))
),
reranker_max_candidates=int(os.getenv(ENV_RERANKER_MAX_CANDIDATES, str(DEFAULT_RERANKER_MAX_CANDIDATES))),
reranker_cohere_base_url=os.getenv(ENV_RERANKER_COHERE_BASE_URL) or None,
# Server
host=os.getenv(ENV_HOST, DEFAULT_HOST),
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
@@ -315,6 +338,9 @@ class HindsightConfig:
graph_retriever=os.getenv(ENV_GRAPH_RETRIEVER, DEFAULT_GRAPH_RETRIEVER),
mpfp_top_k_neighbors=int(os.getenv(ENV_MPFP_TOP_K_NEIGHBORS, str(DEFAULT_MPFP_TOP_K_NEIGHBORS))),
recall_max_concurrent=int(os.getenv(ENV_RECALL_MAX_CONCURRENT, str(DEFAULT_RECALL_MAX_CONCURRENT))),
recall_connection_budget=int(
os.getenv(ENV_RECALL_CONNECTION_BUDGET, str(DEFAULT_RECALL_CONNECTION_BUDGET))
),
# Optimization flags
skip_llm_verification=os.getenv(ENV_SKIP_LLM_VERIFICATION, "false").lower() == "true",
lazy_reranker=os.getenv(ENV_LAZY_RERANKER, "false").lower() == "true",
@@ -15,18 +15,24 @@ from concurrent.futures import ThreadPoolExecutor
import httpx
from ..config import (
DEFAULT_LITELLM_API_BASE,
DEFAULT_RERANKER_COHERE_MODEL,
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR,
DEFAULT_RERANKER_FLASHRANK_MODEL,
DEFAULT_RERANKER_LITELLM_MODEL,
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
DEFAULT_RERANKER_LOCAL_MODEL,
DEFAULT_RERANKER_PROVIDER,
DEFAULT_RERANKER_TEI_BATCH_SIZE,
DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
ENV_COHERE_API_KEY,
ENV_LITELLM_API_BASE,
ENV_LITELLM_API_KEY,
ENV_RERANKER_COHERE_BASE_URL,
ENV_RERANKER_COHERE_MODEL,
ENV_RERANKER_FLASHRANK_CACHE_DIR,
ENV_RERANKER_FLASHRANK_MODEL,
ENV_RERANKER_LITELLM_MODEL,
ENV_RERANKER_LOCAL_MAX_CONCURRENT,
ENV_RERANKER_LOCAL_MODEL,
ENV_RERANKER_PROVIDER,
@@ -392,6 +398,7 @@ class CohereCrossEncoder(CrossEncoderModel):
self,
api_key: str,
model: str = DEFAULT_RERANKER_COHERE_MODEL,
base_url: str | None = None,
timeout: float = 60.0,
):
"""
@@ -400,10 +407,12 @@ class CohereCrossEncoder(CrossEncoderModel):
Args:
api_key: Cohere API key
model: Cohere rerank model name (default: rerank-english-v3.0)
base_url: Custom base URL for Cohere-compatible API (e.g., Azure-hosted endpoint)
timeout: Request timeout in seconds (default: 60.0)
"""
self.api_key = api_key
self.model = model
self.base_url = base_url
self.timeout = timeout
self._client = None
@@ -421,8 +430,14 @@ class CohereCrossEncoder(CrossEncoderModel):
except ImportError:
raise ImportError("cohere is required for CohereCrossEncoder. Install it with: pip install cohere")
logger.info(f"Reranker: initializing Cohere provider with model {self.model}")
self._client = cohere.Client(api_key=self.api_key, timeout=self.timeout)
base_url_msg = f" at {self.base_url}" if self.base_url else ""
logger.info(f"Reranker: initializing Cohere provider with model {self.model}{base_url_msg}")
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
client_kwargs = {"api_key": self.api_key, "timeout": self.timeout}
if self.base_url:
client_kwargs["base_url"] = self.base_url
self._client = cohere.Client(**client_kwargs)
logger.info("Reranker: Cohere provider initialized")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
@@ -641,6 +656,116 @@ class FlashRankCrossEncoder(CrossEncoderModel):
return await loop.run_in_executor(FlashRankCrossEncoder._executor, self._predict_sync, pairs)
class LiteLLMCrossEncoder(CrossEncoderModel):
"""
LiteLLM cross-encoder implementation using LiteLLM proxy's /rerank endpoint.
LiteLLM provides a unified interface for multiple reranking providers via
the Cohere-compatible /rerank endpoint.
See: https://docs.litellm.ai/docs/rerank
Supported providers via LiteLLM:
- Cohere (rerank-english-v3.0, etc.) - prefix with cohere/
- Together AI - prefix with together_ai/
- Azure AI - prefix with azure_ai/
- Jina AI - prefix with jina_ai/
- AWS Bedrock - prefix with bedrock/
- Voyage AI - prefix with voyage/
"""
def __init__(
self,
api_base: str = DEFAULT_LITELLM_API_BASE,
api_key: str | None = None,
model: str = DEFAULT_RERANKER_LITELLM_MODEL,
timeout: float = 60.0,
):
"""
Initialize LiteLLM cross-encoder client.
Args:
api_base: Base URL of the LiteLLM proxy (default: http://localhost:4000)
api_key: API key for the LiteLLM proxy (optional, depends on proxy config)
model: Reranking model name (default: cohere/rerank-english-v3.0)
Use provider prefix (e.g., cohere/, together_ai/, voyage/)
timeout: Request timeout in seconds (default: 60.0)
"""
self.api_base = api_base.rstrip("/")
self.api_key = api_key
self.model = model
self.timeout = timeout
self._async_client: httpx.AsyncClient | None = None
@property
def provider_name(self) -> str:
return "litellm"
async def initialize(self) -> None:
"""Initialize the async HTTP client."""
if self._async_client is not None:
return
logger.info(f"Reranker: initializing LiteLLM provider at {self.api_base} with model {self.model}")
headers = {"Content-Type": "application/json"}
if self.api_key:
headers["Authorization"] = f"Bearer {self.api_key}"
self._async_client = httpx.AsyncClient(timeout=self.timeout, headers=headers)
logger.info("Reranker: LiteLLM provider initialized")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using the LiteLLM proxy's /rerank endpoint.
Args:
pairs: List of (query, document) tuples to score
Returns:
List of relevance scores
"""
if self._async_client is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
if not pairs:
return []
# Group pairs by query (LiteLLM rerank expects one query with multiple documents)
query_groups: dict[str, list[tuple[int, str]]] = {}
for idx, (query, text) in enumerate(pairs):
if query not in query_groups:
query_groups[query] = []
query_groups[query].append((idx, text))
all_scores = [0.0] * len(pairs)
for query, indexed_texts in query_groups.items():
texts = [text for _, text in indexed_texts]
indices = [idx for idx, _ in indexed_texts]
# LiteLLM /rerank follows Cohere API format
response = await self._async_client.post(
f"{self.api_base}/rerank",
json={
"model": self.model,
"query": query,
"documents": texts,
"top_n": len(texts), # Return all scores
},
)
response.raise_for_status()
result = response.json()
# Map scores back to original positions
# Response format: {"results": [{"index": 0, "relevance_score": 0.9}, ...]}
for item in result.get("results", []):
original_idx = item["index"]
score = item.get("relevance_score", item.get("score", 0.0))
all_scores[indices[original_idx]] = score
return all_scores
def create_cross_encoder_from_env() -> CrossEncoderModel:
"""
Create a CrossEncoderModel instance based on environment variables.
@@ -671,14 +796,20 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
if not api_key:
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_RERANKER_PROVIDER} is 'cohere'")
model = os.environ.get(ENV_RERANKER_COHERE_MODEL, DEFAULT_RERANKER_COHERE_MODEL)
return CohereCrossEncoder(api_key=api_key, model=model)
base_url = os.environ.get(ENV_RERANKER_COHERE_BASE_URL) or None
return CohereCrossEncoder(api_key=api_key, model=model, base_url=base_url)
elif provider == "flashrank":
model = os.environ.get(ENV_RERANKER_FLASHRANK_MODEL, DEFAULT_RERANKER_FLASHRANK_MODEL)
cache_dir = os.environ.get(ENV_RERANKER_FLASHRANK_CACHE_DIR, DEFAULT_RERANKER_FLASHRANK_CACHE_DIR)
return FlashRankCrossEncoder(model_name=model, cache_dir=cache_dir)
elif provider == "litellm":
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
api_key = os.environ.get(ENV_LITELLM_API_KEY)
model = os.environ.get(ENV_RERANKER_LITELLM_MODEL, DEFAULT_RERANKER_LITELLM_MODEL)
return LiteLLMCrossEncoder(api_base=api_base, api_key=api_key, model=model)
elif provider == "rrf":
return RRFPassthroughCrossEncoder()
else:
raise ValueError(
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'flashrank', 'rrf'"
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'flashrank', 'litellm', 'rrf'"
)
@@ -0,0 +1,284 @@
"""
Database connection budget management.
Limits concurrent database connections per operation to prevent
a single operation (e.g., recall with parallel queries) from
exhausting the connection pool.
"""
import asyncio
import logging
import uuid
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, AsyncIterator
if TYPE_CHECKING:
import asyncpg
logger = logging.getLogger(__name__)
@dataclass
class OperationBudget:
"""
Tracks connection budget for a single operation.
Each operation gets a semaphore limiting its concurrent connections.
"""
operation_id: str
max_connections: int
semaphore: asyncio.Semaphore = field(init=False)
active_count: int = field(default=0, init=False)
def __post_init__(self):
self.semaphore = asyncio.Semaphore(self.max_connections)
class ConnectionBudgetManager:
"""
Manages per-operation connection budgets.
Usage:
manager = ConnectionBudgetManager(default_budget=4)
# Start an operation
async with manager.operation(max_connections=2) as op:
# Acquire connections within the budget
async with op.acquire(pool) as conn:
await conn.fetch(...)
# Multiple connections respect the budget
async with op.acquire(pool) as conn1, op.acquire(pool) as conn2:
# At most 2 concurrent connections for this operation
...
"""
def __init__(self, default_budget: int = 4):
"""
Initialize the budget manager.
Args:
default_budget: Default max connections per operation
"""
self.default_budget = default_budget
self._operations: dict[str, OperationBudget] = {}
self._lock = asyncio.Lock()
@asynccontextmanager
async def operation(
self,
max_connections: int | None = None,
operation_id: str | None = None,
) -> AsyncIterator["BudgetedOperation"]:
"""
Create a budgeted operation context.
Args:
max_connections: Max concurrent connections for this operation.
Defaults to manager's default_budget.
operation_id: Optional custom operation ID. Auto-generated if not provided.
Yields:
BudgetedOperation context for acquiring connections
"""
op_id = operation_id or f"op-{uuid.uuid4().hex[:12]}"
budget = max_connections or self.default_budget
async with self._lock:
if op_id in self._operations:
raise ValueError(f"Operation {op_id} already exists")
self._operations[op_id] = OperationBudget(op_id, budget)
try:
yield BudgetedOperation(self, op_id)
finally:
async with self._lock:
self._operations.pop(op_id, None)
def _get_budget(self, operation_id: str) -> OperationBudget:
"""Get budget for an operation (internal use)."""
budget = self._operations.get(operation_id)
if not budget:
raise ValueError(f"Operation {operation_id} not found")
return budget
class BudgetedOperation:
"""
A single operation with connection budget.
Provides methods to acquire connections within the budget.
"""
def __init__(self, manager: ConnectionBudgetManager, operation_id: str):
self._manager = manager
self.operation_id = operation_id
@property
def budget(self) -> OperationBudget:
"""Get the budget for this operation."""
return self._manager._get_budget(self.operation_id)
@asynccontextmanager
async def acquire(self, pool: "asyncpg.Pool") -> AsyncIterator["asyncpg.Connection"]:
"""
Acquire a connection within the operation's budget.
Blocks if the operation has reached its connection limit.
Args:
pool: asyncpg connection pool
Yields:
Database connection
"""
budget = self.budget
async with budget.semaphore:
budget.active_count += 1
conn = await pool.acquire()
try:
yield conn
finally:
budget.active_count -= 1
await pool.release(conn)
def wrap_pool(self, pool: "asyncpg.Pool") -> "BudgetedPool":
"""
Wrap a pool with this operation's budget.
The returned BudgetedPool can be passed to functions expecting a pool,
and all acquire() calls will be limited by this operation's budget.
Args:
pool: asyncpg connection pool to wrap
Returns:
BudgetedPool that limits connections to this operation's budget
"""
return BudgetedPool(pool, self)
async def acquire_many(
self,
pool: "asyncpg.Pool",
count: int,
) -> AsyncIterator[list["asyncpg.Connection"]]:
"""
Acquire multiple connections within the budget.
Note: This acquires connections sequentially to respect the budget.
For parallel acquisition, use multiple acquire() calls with asyncio.gather().
Args:
pool: asyncpg connection pool
count: Number of connections to acquire
Yields:
List of database connections
"""
connections = []
try:
for _ in range(count):
conn = await pool.acquire()
connections.append(conn)
yield connections
finally:
for conn in connections:
await pool.release(conn)
# Global default manager instance
_default_manager: ConnectionBudgetManager | None = None
def get_budget_manager(default_budget: int = 4) -> ConnectionBudgetManager:
"""
Get or create the global budget manager.
Args:
default_budget: Default max connections per operation
Returns:
Global ConnectionBudgetManager instance
"""
global _default_manager
if _default_manager is None:
_default_manager = ConnectionBudgetManager(default_budget=default_budget)
return _default_manager
@asynccontextmanager
async def budgeted_operation(
max_connections: int | None = None,
operation_id: str | None = None,
default_budget: int = 4,
) -> AsyncIterator[BudgetedOperation]:
"""
Convenience function to create a budgeted operation.
Args:
max_connections: Max concurrent connections for this operation
operation_id: Optional custom operation ID
default_budget: Default budget if manager not yet created
Yields:
BudgetedOperation context
Example:
async with budgeted_operation(max_connections=2) as op:
async with op.acquire(pool) as conn:
await conn.fetch(...)
"""
manager = get_budget_manager(default_budget)
async with manager.operation(max_connections, operation_id) as op:
yield op
class BudgetedPool:
"""
A pool wrapper that limits concurrent connection acquisitions.
This can be passed to functions expecting a pool, and acquire()
calls will be limited by the budget semaphore.
Usage:
async with budgeted_operation(max_connections=4) as op:
budgeted_pool = op.wrap_pool(pool)
# Pass budgeted_pool to functions that expect a pool
await some_function(budgeted_pool, ...)
"""
def __init__(self, pool: "asyncpg.Pool", operation: BudgetedOperation):
self._pool = pool
self._operation = operation
async def acquire(self) -> "asyncpg.Connection":
"""
Acquire a connection within the budget.
Note: Caller must release the connection when done.
Prefer using as context manager via acquire_with_retry or op.acquire().
"""
budget = self._operation.budget
await budget.semaphore.acquire()
budget.active_count += 1
try:
return await self._pool.acquire()
except Exception:
budget.active_count -= 1
budget.semaphore.release()
raise
async def release(self, conn: "asyncpg.Connection") -> None:
"""Release a connection back to the pool."""
budget = self._operation.budget
try:
await self._pool.release(conn)
finally:
budget.active_count -= 1
budget.semaphore.release()
def __getattr__(self, name):
"""Proxy other attributes to the underlying pool."""
return getattr(self._pool, name)
@@ -17,16 +17,23 @@ import httpx
from ..config import (
DEFAULT_EMBEDDINGS_COHERE_MODEL,
DEFAULT_EMBEDDINGS_LITELLM_MODEL,
DEFAULT_EMBEDDINGS_LOCAL_MODEL,
DEFAULT_EMBEDDINGS_OPENAI_MODEL,
DEFAULT_EMBEDDINGS_PROVIDER,
DEFAULT_LITELLM_API_BASE,
ENV_COHERE_API_KEY,
ENV_EMBEDDINGS_COHERE_BASE_URL,
ENV_EMBEDDINGS_COHERE_MODEL,
ENV_EMBEDDINGS_LITELLM_MODEL,
ENV_EMBEDDINGS_LOCAL_MODEL,
ENV_EMBEDDINGS_OPENAI_API_KEY,
ENV_EMBEDDINGS_OPENAI_BASE_URL,
ENV_EMBEDDINGS_OPENAI_MODEL,
ENV_EMBEDDINGS_PROVIDER,
ENV_EMBEDDINGS_TEI_URL,
ENV_LITELLM_API_BASE,
ENV_LITELLM_API_KEY,
ENV_LLM_API_KEY,
)
@@ -322,6 +329,7 @@ class OpenAIEmbeddings(Embeddings):
self,
api_key: str,
model: str = DEFAULT_EMBEDDINGS_OPENAI_MODEL,
base_url: str | None = None,
batch_size: int = 100,
max_retries: int = 3,
):
@@ -331,11 +339,13 @@ class OpenAIEmbeddings(Embeddings):
Args:
api_key: OpenAI API key
model: OpenAI embedding model name (default: text-embedding-3-small)
base_url: Custom base URL for OpenAI-compatible API (e.g., Azure OpenAI endpoint)
batch_size: Maximum batch size for embedding requests (default: 100)
max_retries: Maximum number of retries for failed requests (default: 3)
"""
self.api_key = api_key
self.model = model
self.base_url = base_url
self.batch_size = batch_size
self.max_retries = max_retries
self._client = None
@@ -361,8 +371,14 @@ class OpenAIEmbeddings(Embeddings):
except ImportError:
raise ImportError("openai is required for OpenAIEmbeddings. Install it with: pip install openai")
logger.info(f"Embeddings: initializing OpenAI provider with model {self.model}")
self._client = OpenAI(api_key=self.api_key, max_retries=self.max_retries)
base_url_msg = f" at {self.base_url}" if self.base_url else ""
logger.info(f"Embeddings: initializing OpenAI provider with model {self.model}{base_url_msg}")
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
client_kwargs = {"api_key": self.api_key, "max_retries": self.max_retries}
if self.base_url:
client_kwargs["base_url"] = self.base_url
self._client = OpenAI(**client_kwargs)
# Try to get dimension from known models, otherwise do a test embedding
if self.model in self.MODEL_DIMENSIONS:
@@ -435,6 +451,7 @@ class CohereEmbeddings(Embeddings):
self,
api_key: str,
model: str = DEFAULT_EMBEDDINGS_COHERE_MODEL,
base_url: str | None = None,
batch_size: int = 96,
timeout: float = 60.0,
input_type: str = "search_document",
@@ -445,6 +462,7 @@ class CohereEmbeddings(Embeddings):
Args:
api_key: Cohere API key
model: Cohere embedding model name (default: embed-english-v3.0)
base_url: Custom base URL for Cohere-compatible API (e.g., Azure-hosted endpoint)
batch_size: Maximum batch size for embedding requests (default: 96, Cohere's limit)
timeout: Request timeout in seconds (default: 60.0)
input_type: Input type for embeddings (default: search_document).
@@ -452,6 +470,7 @@ class CohereEmbeddings(Embeddings):
"""
self.api_key = api_key
self.model = model
self.base_url = base_url
self.batch_size = batch_size
self.timeout = timeout
self.input_type = input_type
@@ -478,8 +497,14 @@ class CohereEmbeddings(Embeddings):
except ImportError:
raise ImportError("cohere is required for CohereEmbeddings. Install it with: pip install cohere")
logger.info(f"Embeddings: initializing Cohere provider with model {self.model}")
self._client = cohere.Client(api_key=self.api_key, timeout=self.timeout)
base_url_msg = f" at {self.base_url}" if self.base_url else ""
logger.info(f"Embeddings: initializing Cohere provider with model {self.model}{base_url_msg}")
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
client_kwargs = {"api_key": self.api_key, "timeout": self.timeout}
if self.base_url:
client_kwargs["base_url"] = self.base_url
self._client = cohere.Client(**client_kwargs)
# Try to get dimension from known models, otherwise do a test embedding
if self.model in self.MODEL_DIMENSIONS:
@@ -529,6 +554,123 @@ class CohereEmbeddings(Embeddings):
return all_embeddings
class LiteLLMEmbeddings(Embeddings):
"""
LiteLLM embeddings implementation using LiteLLM proxy's /embeddings endpoint.
LiteLLM provides a unified interface for multiple embedding providers.
The proxy exposes an OpenAI-compatible /embeddings endpoint.
See: https://docs.litellm.ai/docs/embedding/supported_embedding
Supported providers via LiteLLM:
- OpenAI (text-embedding-3-small, text-embedding-ada-002, etc.)
- Cohere (embed-english-v3.0, etc.) - prefix with cohere/
- Vertex AI (textembedding-gecko, etc.) - prefix with vertex_ai/
- HuggingFace, Mistral, Voyage AI, etc.
The embedding dimension is auto-detected from the model at initialization.
"""
def __init__(
self,
api_base: str = DEFAULT_LITELLM_API_BASE,
api_key: str | None = None,
model: str = DEFAULT_EMBEDDINGS_LITELLM_MODEL,
batch_size: int = 100,
timeout: float = 60.0,
):
"""
Initialize LiteLLM embeddings client.
Args:
api_base: Base URL of the LiteLLM proxy (default: http://localhost:4000)
api_key: API key for the LiteLLM proxy (optional, depends on proxy config)
model: Embedding model name (default: text-embedding-3-small)
Use provider prefix for non-OpenAI models (e.g., cohere/embed-english-v3.0)
batch_size: Maximum batch size for embedding requests (default: 100)
timeout: Request timeout in seconds (default: 60.0)
"""
self.api_base = api_base.rstrip("/")
self.api_key = api_key
self.model = model
self.batch_size = batch_size
self.timeout = timeout
self._client: httpx.Client | None = None
self._dimension: int | None = None
@property
def provider_name(self) -> str:
return "litellm"
@property
def dimension(self) -> int:
if self._dimension is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
return self._dimension
async def initialize(self) -> None:
"""Initialize the HTTP client and detect embedding dimension."""
if self._client is not None:
return
logger.info(f"Embeddings: initializing LiteLLM provider at {self.api_base} with model {self.model}")
headers = {"Content-Type": "application/json"}
if self.api_key:
headers["Authorization"] = f"Bearer {self.api_key}"
self._client = httpx.Client(timeout=self.timeout, headers=headers)
# Do a test embedding to detect dimension
try:
response = self._client.post(
f"{self.api_base}/embeddings",
json={"model": self.model, "input": ["test"]},
)
response.raise_for_status()
result = response.json()
if result.get("data") and len(result["data"]) > 0:
self._dimension = len(result["data"][0]["embedding"])
logger.info(f"Embeddings: LiteLLM provider initialized (model: {self.model}, dim: {self._dimension})")
except httpx.HTTPError as e:
raise RuntimeError(f"Failed to connect to LiteLLM proxy at {self.api_base}: {e}")
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings using the LiteLLM proxy.
Args:
texts: List of text strings to encode
Returns:
List of embedding vectors
"""
if self._client is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
if not texts:
return []
all_embeddings = []
# Process in batches
for i in range(0, len(texts), self.batch_size):
batch = texts[i : i + self.batch_size]
response = self._client.post(
f"{self.api_base}/embeddings",
json={"model": self.model, "input": batch},
)
response.raise_for_status()
result = response.json()
# Sort by index to ensure correct order
batch_embeddings = sorted(result["data"], key=lambda x: x["index"])
all_embeddings.extend([e["embedding"] for e in batch_embeddings])
return all_embeddings
def create_embeddings_from_env() -> Embeddings:
"""
Create an Embeddings instance based on environment variables.
@@ -558,12 +700,21 @@ def create_embeddings_from_env() -> Embeddings:
f"when {ENV_EMBEDDINGS_PROVIDER} is 'openai'"
)
model = os.environ.get(ENV_EMBEDDINGS_OPENAI_MODEL, DEFAULT_EMBEDDINGS_OPENAI_MODEL)
return OpenAIEmbeddings(api_key=api_key, model=model)
base_url = os.environ.get(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None
return OpenAIEmbeddings(api_key=api_key, model=model, base_url=base_url)
elif provider == "cohere":
api_key = os.environ.get(ENV_COHERE_API_KEY)
if not api_key:
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_EMBEDDINGS_PROVIDER} is 'cohere'")
model = os.environ.get(ENV_EMBEDDINGS_COHERE_MODEL, DEFAULT_EMBEDDINGS_COHERE_MODEL)
return CohereEmbeddings(api_key=api_key, model=model)
base_url = os.environ.get(ENV_EMBEDDINGS_COHERE_BASE_URL) or None
return CohereEmbeddings(api_key=api_key, model=model, base_url=base_url)
elif provider == "litellm":
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
api_key = os.environ.get(ENV_LITELLM_API_KEY)
model = os.environ.get(ENV_EMBEDDINGS_LITELLM_MODEL, DEFAULT_EMBEDDINGS_LITELLM_MODEL)
return LiteLLMEmbeddings(api_base=api_base, api_key=api_key, model=model)
else:
raise ValueError(f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei', 'openai', 'cohere'")
raise ValueError(
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei', 'openai', 'cohere', 'litellm'"
)
@@ -19,6 +19,7 @@ from typing import TYPE_CHECKING, Any
from ..config import get_config
from ..metrics import get_metrics_collector
from .db_budget import budgeted_operation
# Context variable for current schema (async-safe, per-task isolation)
_current_schema: contextvars.ContextVar[str] = contextvars.ContextVar("current_schema", default="public")
@@ -150,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
@@ -1058,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,
):
"""
@@ -1190,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
@@ -1208,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
@@ -1242,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.
@@ -1258,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)
@@ -1282,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(
@@ -1340,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).
@@ -1365,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:
@@ -1437,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:
@@ -1555,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.
@@ -1584,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()
@@ -1594,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:
@@ -1609,38 +1628,42 @@ class MemoryEngine(MemoryEngineInterface):
tracer.record_query_embedding(query_embedding)
tracer.add_phase_metric("generate_query_embedding", step_duration)
# Step 2: N*4-Way Parallel Retrieval (N fact types × 4 retrieval methods)
# Step 2: Optimized parallel retrieval using batched queries
# - Semantic + BM25 combined in 1 CTE query for ALL fact types
# - Graph runs per fact type (complex traversal)
# - Temporal runs per fact type (if constraint detected)
step_start = time.time()
query_embedding_str = str(query_embedding)
from .search.retrieval import get_default_graph_retriever, retrieve_parallel
from .search.retrieval import (
get_default_graph_retriever,
retrieve_all_fact_types_parallel,
)
# Track each retrieval start time
retrieval_start = time.time()
# Temporal extraction now runs IN PARALLEL with other retrievals inside retrieve_parallel
# This prevents slow dateparser from blocking semantic/BM25/graph retrieval
tc_duration = 0.0 # Will be tracked inside temporal retrieval timing
# Run retrieval for each fact type in parallel
# Each retrieve_parallel uses ~4 connections, so 3 fact types = ~12 concurrent connections
retrieval_tasks = [
retrieve_parallel(
pool,
# Run optimized retrieval with connection budget
config = get_config()
async with budgeted_operation(
max_connections=config.recall_connection_budget,
operation_id=f"recall-{recall_id}",
) as op:
budgeted_pool = op.wrap_pool(pool)
parallel_start = time.time()
multi_result = await retrieve_all_fact_types_parallel(
budgeted_pool,
query,
query_embedding_str,
bank_id,
ft,
fact_type, # Pass all fact types at once
thinking_budget,
question_date,
self.query_analyzer,
temporal_constraint=None, # Extracted in parallel inside retrieve_parallel
tags=tags,
tags_match=tags_match,
)
for ft in fact_type
]
parallel_start = time.time()
all_retrievals = await asyncio.gather(*retrieval_tasks)
parallel_duration = time.time() - parallel_start
parallel_duration = time.time() - parallel_start
# Combine all results from all fact types and aggregate timings
semantic_results = []
@@ -1657,12 +1680,15 @@ class MemoryEngine(MemoryEngineInterface):
all_mpfp_timings = []
detected_temporal_constraint = None
max_conn_wait = 0.0
for idx, retrieval_result in enumerate(all_retrievals):
max_conn_wait = multi_result.max_conn_wait
for ft in fact_type:
retrieval_result = multi_result.results_by_fact_type.get(ft)
if not retrieval_result:
continue
# Log fact types in this retrieval batch
ft_name = fact_type[idx] if idx < len(fact_type) else "unknown"
logger.debug(
f"[RECALL {recall_id}] Fact type '{ft_name}': semantic={len(retrieval_result.semantic)}, bm25={len(retrieval_result.bm25)}, graph={len(retrieval_result.graph)}, temporal={len(retrieval_result.temporal) if retrieval_result.temporal else 0}"
f"[RECALL {recall_id}] Fact type '{ft}': semantic={len(retrieval_result.semantic)}, bm25={len(retrieval_result.bm25)}, graph={len(retrieval_result.graph)}, temporal={len(retrieval_result.temporal) if retrieval_result.temporal else 0}"
)
semantic_results.extend(retrieval_result.semantic)
@@ -1678,8 +1704,6 @@ class MemoryEngine(MemoryEngineInterface):
detected_temporal_constraint = retrieval_result.temporal_constraint
# Collect MPFP timings
all_mpfp_timings.extend(retrieval_result.mpfp_timings)
# Track max connection wait
max_conn_wait = max(max_conn_wait, retrieval_result.max_conn_wait)
# If no temporal results from any fact type, set to None
if not temporal_results:
@@ -1711,10 +1735,8 @@ class MemoryEngine(MemoryEngineInterface):
temporal_count = len(temporal_results) if temporal_results else 0
timing_parts.append(f"temporal={temporal_count}({aggregated_timings['temporal']:.3f}s)")
temporal_info = f" | temporal_range={start_dt.strftime('%Y-%m-%d')} to {end_dt.strftime('%Y-%m-%d')}"
# Only tc is sequential setup now (adjacency loads in parallel with retrieval)
setup_info = f", tc={tc_duration:.3f}s" if tc_duration > 0.01 else ""
log_buffer.append(
f" [2] Parallel retrieval ({len(fact_type)} fact_types): {', '.join(timing_parts)} in {parallel_duration:.3f}s{setup_info}{temporal_info}"
f" [2] Parallel retrieval ({len(fact_type)} fact_types): {', '.join(timing_parts)} in {parallel_duration:.3f}s{temporal_info}"
)
# Log graph retriever timing breakdown if available
@@ -1744,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
@@ -1751,8 +1778,10 @@ class MemoryEngine(MemoryEngineInterface):
return [(r.id, r.__dict__) for r in results]
# Add retrieval results per fact type (to show parallel execution in UI)
for idx, rr in enumerate(all_retrievals):
ft_name = fact_type[idx] if idx < len(fact_type) else "unknown"
for ft_name in fact_type:
rr = multi_result.results_by_fact_type.get(ft_name)
if not rr:
continue
# Add semantic retrieval results for this fact type
tracer.add_retrieval_results(
@@ -1784,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,
)
@@ -2051,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"),
)
)
@@ -2266,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,
@@ -2287,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(
@@ -2775,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,
@@ -3298,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.
@@ -3365,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
@@ -3678,7 +3783,7 @@ Guidelines:
SELECT id, canonical_name, mention_count, first_seen, last_seen, metadata
FROM {fq_table("entities")}
WHERE bank_id = $1
ORDER BY mention_count DESC, last_seen DESC
ORDER BY mention_count DESC, last_seen DESC, id ASC
LIMIT $2 OFFSET $3
""",
bank_id,
@@ -3717,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,
@@ -4361,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)
@@ -4384,14 +4569,16 @@ Guidelines:
)
# Submit task to background queue
await self._task_backend.submit_task(
{
"type": "batch_retain",
"operation_id": str(operation_id),
"bank_id": bank_id,
"contents": contents,
}
)
task_payload = {
"type": "batch_retain",
"operation_id": str(operation_id),
"bank_id": bank_id,
"contents": contents,
}
if document_tags:
task_payload["document_tags"] = document_tags
await self._task_backend.submit_task(task_payload)
logger.info(f"Retain task queued for bank_id={bank_id}, {len(contents)} items, operation_id={operation_id}")
@@ -85,6 +85,7 @@ class MemoryFact(BaseModel):
"metadata": {"source": "slack"},
"chunk_id": "bank123_session_abc123_0",
"activation": 0.95,
"tags": ["user_a", "session_123"],
}
}
)
@@ -102,6 +103,7 @@ class MemoryFact(BaseModel):
chunk_id: str | None = Field(
None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)"
)
tags: list[str] | None = Field(None, description="Visibility scope tags associated with this fact")
class ChunkInfo(BaseModel):
@@ -1268,6 +1268,7 @@ async def extract_facts_from_contents(
# mentioned_at: always the event_date (when the conversation/document occurred)
mentioned_at=content.event_date,
metadata=content.metadata,
tags=content.tags,
)
extracted_facts.append(extracted_fact)
@@ -45,6 +45,7 @@ async def insert_facts_batch(
metadata_jsons = []
chunk_ids = []
document_ids = []
tags_list = []
for fact in facts:
fact_texts.append(fact.fact_text)
@@ -65,16 +66,31 @@ async def insert_facts_batch(
chunk_ids.append(fact.chunk_id)
# Use per-fact document_id if available, otherwise fallback to batch-level document_id
document_ids.append(fact.document_id if fact.document_id else document_id)
# Convert tags to JSON string for proper batch insertion (PostgreSQL unnest doesn't handle 2D arrays well)
tags_list.append(json.dumps(fact.tags if fact.tags else []))
# Batch insert all facts
# Note: tags are passed as JSON strings and converted back to varchar[] via jsonb_array_elements_text + array_agg
results = await conn.fetch(
f"""
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id)
SELECT $1, * FROM unnest(
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[]
WITH input_data AS (
SELECT * FROM unnest(
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[], $15::jsonb[]
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags_json)
)
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags)
SELECT
$1,
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id,
COALESCE(
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
'{{}}'::varchar[]
)
FROM input_data
RETURNING id
""",
bank_id,
@@ -91,6 +107,7 @@ async def insert_facts_batch(
metadata_jsons,
chunk_ids,
document_ids,
tags_list,
)
unit_ids = [str(row["id"]) for row in results]
@@ -121,7 +138,13 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
async def handle_document_tracking(
conn, bank_id: str, document_id: str, combined_content: str, is_first_batch: bool, retain_params: dict | None = None
conn,
bank_id: str,
document_id: str,
combined_content: str,
is_first_batch: bool,
retain_params: dict | None = None,
document_tags: list[str] | None = None,
) -> None:
"""
Handle document tracking in the database.
@@ -133,6 +156,7 @@ async def handle_document_tracking(
combined_content: Combined content text from all content items
is_first_batch: Whether this is the first batch (for chunked operations)
retain_params: Optional parameters passed during retain (context, event_date, etc.)
document_tags: Optional list of tags to associate with the document
"""
import hashlib
@@ -149,13 +173,14 @@ async def handle_document_tracking(
# Insert document (or update if exists from concurrent operations)
await conn.execute(
f"""
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params)
VALUES ($1, $2, $3, $4, $5, $6)
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params, tags)
VALUES ($1, $2, $3, $4, $5, $6, $7)
ON CONFLICT (id, bank_id) DO UPDATE
SET original_text = EXCLUDED.original_text,
content_hash = EXCLUDED.content_hash,
metadata = EXCLUDED.metadata,
retain_params = EXCLUDED.retain_params,
tags = EXCLUDED.tags,
updated_at = NOW()
""",
document_id,
@@ -164,4 +189,5 @@ async def handle_document_tracking(
content_hash,
json.dumps({}), # Empty metadata dict
json.dumps(retain_params) if retain_params else None,
document_tags or [],
)
@@ -49,6 +49,7 @@ async def retain_batch(
is_first_batch: bool = True,
fact_type_override: str | None = None,
confidence_score: float | None = None,
document_tags: list[str] | None = None,
) -> tuple[list[list[str]], TokenUsage]:
"""
Process a batch of content through the retain pipeline.
@@ -67,6 +68,7 @@ async def retain_batch(
is_first_batch: Whether this is the first batch
fact_type_override: Override fact type for all facts
confidence_score: Confidence score for opinions
document_tags: Tags applied to all items in this batch
Returns:
Tuple of (unit ID lists, token usage for fact extraction)
@@ -88,12 +90,16 @@ async def retain_batch(
# Convert dicts to RetainContent objects
contents = []
for item in contents_dicts:
# Merge item-level tags with document-level tags
item_tags = item.get("tags", []) or []
merged_tags = list(set(item_tags + (document_tags or [])))
content = RetainContent(
content=item["content"],
context=item.get("context", ""),
event_date=item.get("event_date") or utcnow(),
metadata=item.get("metadata", {}),
entities=item.get("entities", []),
tags=merged_tags,
)
contents.append(content)
@@ -131,7 +137,7 @@ async def retain_batch(
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
)
else:
# Check for per-item document_ids
@@ -159,7 +165,7 @@ async def retain_batch(
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, document_tags
)
total_time = time.time() - start_time
@@ -225,7 +231,7 @@ async def retain_batch(
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
)
document_ids_added.append(document_id)
doc_id_mapping[None] = document_id # For backwards compatibility
@@ -269,7 +275,13 @@ async def retain_batch(
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, actual_doc_id, combined_content, is_first_batch, retain_params
conn,
bank_id,
actual_doc_id,
combined_content,
is_first_batch,
retain_params,
document_tags,
)
document_ids_added.append(actual_doc_id)
@@ -21,6 +21,7 @@ class RetainContentDict(TypedDict, total=False):
metadata: Custom key-value metadata (optional)
document_id: Document ID for this content item (optional)
entities: User-provided entities to merge with extracted entities (optional)
tags: Visibility scope tags for this content item (optional)
"""
content: str # Required
@@ -29,6 +30,7 @@ class RetainContentDict(TypedDict, total=False):
metadata: dict[str, str]
document_id: str
entities: list[dict[str, str]] # [{"text": "...", "type": "..."}]
tags: list[str] # Visibility scope tags
def _now_utc() -> datetime:
@@ -49,6 +51,7 @@ class RetainContent:
event_date: datetime = field(default_factory=_now_utc)
metadata: dict[str, str] = field(default_factory=dict)
entities: list[dict[str, str]] = field(default_factory=list) # User-provided entities
tags: list[str] = field(default_factory=list) # Visibility scope tags
@dataclass
@@ -113,6 +116,7 @@ class ExtractedFact:
context: str = ""
mentioned_at: datetime | None = None
metadata: dict[str, str] = field(default_factory=dict)
tags: list[str] = field(default_factory=list) # Visibility scope tags
@dataclass
@@ -158,6 +162,9 @@ class ProcessedFact:
# Track which content this fact came from (for user entity merging)
content_index: int = 0
# Visibility scope tags
tags: list[str] = field(default_factory=list)
@property
def is_duplicate(self) -> bool:
"""Check if this fact was marked as a duplicate."""
@@ -201,6 +208,7 @@ class ProcessedFact:
causal_relations=extracted_fact.causal_relations,
chunk_id=chunk_id,
content_index=extracted_fact.content_index,
tags=extracted_fact.tags,
)
@@ -232,6 +240,7 @@ class RetainBatch:
document_id: str | None = None
fact_type_override: str | None = None
confidence_score: float | None = None
document_tags: list[str] = field(default_factory=list) # Tags applied to all items
# Extracted data (populated during processing)
extracted_facts: list[ExtractedFact] = field(default_factory=list)
@@ -11,6 +11,7 @@ from abc import ABC, abstractmethod
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .tags import TagsMatch, filter_results_by_tags
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -43,6 +44,8 @@ class GraphRetriever(ABC):
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # TypedAdjacency, optional pre-loaded graph
tags: list[str] | None = None, # Visibility scope tags for filtering
tags_match: TagsMatch = "any", # How to match tags: 'any' (OR) or 'all' (AND)
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve relevant facts via graph traversal.
@@ -57,6 +60,7 @@ class GraphRetriever(ABC):
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
adjacency: Pre-loaded typed adjacency graph (optional, for MPFP)
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (List of RetrievalResult with activation scores, optional timing info)
@@ -114,6 +118,8 @@ class BFSGraphRetriever(GraphRetriever):
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # Not used by BFS
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts using BFS spreading activation.
@@ -129,7 +135,9 @@ class BFSGraphRetriever(GraphRetriever):
for interface compatibility but not used.
"""
async with acquire_with_retry(pool) as conn:
results = await self._retrieve_with_conn(conn, query_embedding_str, bank_id, fact_type, budget)
results = await self._retrieve_with_conn(
conn, query_embedding_str, bank_id, fact_type, budget, tags=tags, tags_match=tags_match
)
return results, None
async def _retrieve_with_conn(
@@ -139,33 +147,46 @@ class BFSGraphRetriever(GraphRetriever):
bank_id: str,
fact_type: str,
budget: int,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[RetrievalResult]:
"""Internal implementation with connection."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, self.entry_point_threshold, self.entry_point_limit]
if tags:
params.append(tags)
# Step 1: Find entry points
entry_points = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
query_embedding_str,
bank_id,
fact_type,
self.entry_point_threshold,
self.entry_point_limit,
*params,
)
if not entry_points:
logger.debug(
f"[BFS] No entry points found for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
)
return []
logger.debug(
f"[BFS] Found {len(entry_points)} entry points for fact_type={fact_type} "
f"(tags={tags}, tags_match={tags_match})"
)
# Step 2: BFS spreading activation
visited = set()
results = []
@@ -196,7 +217,7 @@ class BFSGraphRetriever(GraphRetriever):
f"""
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.occurred_end,
mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type,
mu.document_id, mu.chunk_id,
mu.document_id, mu.chunk_id, mu.tags,
ml.weight, ml.link_type, ml.from_unit_id
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
@@ -236,4 +257,8 @@ class BFSGraphRetriever(GraphRetriever):
neighbor_result = RetrievalResult.from_db_row(dict(n))
queue.append((neighbor_result, new_activation))
# Apply tags filtering (BFS may traverse into memories that don't match tags criteria)
if tags:
results = filter_results_by_tags(results, tags, match=tags_match)
return results
@@ -18,6 +18,7 @@ import time
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .graph_retrieval import GraphRetriever
from .tags import TagsMatch, filter_results_by_tags
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -30,26 +31,32 @@ async def _find_semantic_seeds(
fact_type: str,
limit: int = 20,
threshold: float = 0.3,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[RetrievalResult]:
"""Find semantic seeds via embedding search."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
if tags:
params.append(tags)
rows = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
1 - (embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_units")}
WHERE bank_id = $2
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
query_embedding_str,
bank_id,
fact_type,
threshold,
limit,
*params,
)
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
@@ -95,6 +102,8 @@ class LinkExpansionRetriever(GraphRetriever):
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts by expanding links from seeds.
@@ -109,6 +118,7 @@ class LinkExpansionRetriever(GraphRetriever):
semantic_seeds: Pre-computed semantic entry points
temporal_seeds: Pre-computed temporal entry points
adjacency: Unused, kept for interface compatibility
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (results, timings)
@@ -125,15 +135,27 @@ class LinkExpansionRetriever(GraphRetriever):
else:
seeds_start = time.time()
all_seeds = await _find_semantic_seeds(
conn, query_embedding_str, bank_id, fact_type, limit=20, threshold=0.3
conn,
query_embedding_str,
bank_id,
fact_type,
limit=20,
threshold=0.3,
tags=tags,
tags_match=tags_match,
)
timings.seeds_time = time.time() - seeds_start
logger.debug(
f"[LinkExpansion] Found {len(all_seeds)} semantic seeds for fact_type={fact_type} "
f"(tags={tags}, tags_match={tags_match})"
)
# Add temporal seeds if provided
if temporal_seeds:
all_seeds.extend(temporal_seeds)
if not all_seeds:
logger.debug("[LinkExpansion] No seeds found, returning empty results")
return [], timings
seed_ids = list({s.id for s in all_seeds})
@@ -147,7 +169,7 @@ class LinkExpansionRetriever(GraphRetriever):
SELECT
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
COUNT(*)::float AS score
FROM {fq_table("unit_entities")} seed_ue
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
@@ -172,7 +194,7 @@ class LinkExpansionRetriever(GraphRetriever):
SELECT DISTINCT ON (mu.id)
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
mu.fact_type, mu.document_id, mu.chunk_id,
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
ml.weight + 1.0 AS score
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
@@ -219,6 +241,10 @@ class LinkExpansionRetriever(GraphRetriever):
result.activation = row["score"]
results.append(result)
# Apply tags filtering (graph expansion may reach untagged memories)
if tags:
results = filter_results_by_tags(results, tags, match=tags_match)
timings.result_count = len(results)
timings.traverse = time.time() - start_time
@@ -23,6 +23,7 @@ from dataclasses import dataclass, field
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .graph_retrieval import GraphRetriever
from .tags import TagsMatch
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -448,7 +449,7 @@ async def fetch_memory_units_by_ids(
rows = await conn.fetch(
f"""
SELECT id, text, context, event_date, occurred_start, occurred_end,
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags
FROM {fq_table("memory_units")}
WHERE id = ANY($1::uuid[])
AND fact_type = $2
@@ -503,6 +504,8 @@ class MPFPGraphRetriever(GraphRetriever):
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
adjacency=None, # Ignored - kept for interface compatibility
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
"""
Retrieve facts using MPFP algorithm with lazy edge loading.
@@ -517,6 +520,7 @@ class MPFPGraphRetriever(GraphRetriever):
semantic_seeds: Pre-computed semantic entry points
temporal_seeds: Pre-computed temporal entry points
adjacency: Ignored (kept for interface compatibility)
tags: Optional list of tags for visibility filtering (OR matching)
Returns:
Tuple of (List of RetrievalResult with activation scores, MPFPTimings)
@@ -532,8 +536,13 @@ class MPFPGraphRetriever(GraphRetriever):
# If no semantic seeds provided, fall back to finding our own
if not semantic_seed_nodes:
seeds_start = time.time()
semantic_seed_nodes = await self._find_semantic_seeds(pool, query_embedding_str, bank_id, fact_type)
semantic_seed_nodes = await self._find_semantic_seeds(
pool, query_embedding_str, bank_id, fact_type, tags=tags, tags_match=tags_match
)
timings.seeds_time = time.time() - seeds_start
logger.debug(
f"[MPFP] Found {len(semantic_seed_nodes)} semantic seeds for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
)
# Collect all pattern jobs
pattern_jobs = []
@@ -549,6 +558,9 @@ class MPFPGraphRetriever(GraphRetriever):
pattern_jobs.append((temporal_seed_nodes, pattern))
if not pattern_jobs:
logger.debug(
f"[MPFP] No pattern jobs (semantic_seeds={len(semantic_seed_nodes)}, temporal_seeds={len(temporal_seed_nodes)})"
)
return [], timings
timings.pattern_count = len(pattern_jobs)
@@ -587,6 +599,7 @@ class MPFPGraphRetriever(GraphRetriever):
timings.fusion = time.time() - step_start
if not fused:
logger.debug(f"[MPFP] No fused results after RRF fusion (pattern_count={len(pattern_results)})")
return [], timings
# Get top result IDs
@@ -596,6 +609,13 @@ class MPFPGraphRetriever(GraphRetriever):
step_start = time.time()
results = await fetch_memory_units_by_ids(pool, result_ids, fact_type)
timings.fetch = time.time() - step_start
# Filter results by tags (graph traversal may have picked up unfiltered memories)
if tags:
from .tags import filter_results_by_tags
results = filter_results_by_tags(results, tags, match=tags_match)
timings.result_count = len(results)
# Add activation scores from fusion
@@ -634,8 +654,17 @@ class MPFPGraphRetriever(GraphRetriever):
fact_type: str,
limit: int = 20,
threshold: float = 0.3,
tags: list[str] | None = None,
tags_match: TagsMatch = "any",
) -> list[SeedNode]:
"""Fallback: find semantic seeds via embedding search."""
from .tags import build_tags_where_clause_simple
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
if tags:
params.append(tags)
async with acquire_with_retry(pool) as conn:
rows = await conn.fetch(
f"""
@@ -645,14 +674,11 @@ class MPFPGraphRetriever(GraphRetriever):
AND embedding IS NOT NULL
AND fact_type = $3
AND (1 - (embedding <=> $1::vector)) >= $4
{tags_clause}
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
query_embedding_str,
bank_id,
fact_type,
threshold,
limit,
*params,
)
return [SeedNode(node_id=str(r["id"]), score=r["similarity"]) for r in rows]
@@ -20,6 +20,7 @@ from ..memory_engine import fq_table
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
from .link_expansion_retrieval import LinkExpansionRetriever
from .mpfp_retrieval import MPFPGraphRetriever
from .tags import TagsMatch, build_tags_where_clause_simple
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -39,6 +40,18 @@ class ParallelRetrievalResult:
max_conn_wait: float = 0.0 # Maximum connection acquisition wait time across all methods
@dataclass
class MultiFactTypeRetrievalResult:
"""Result from retrieval across all fact types."""
# Results per fact type
results_by_fact_type: dict[str, ParallelRetrievalResult]
# Aggregate timings
timings: dict[str, float] = field(default_factory=dict)
# Max connection wait across all operations
max_conn_wait: float = 0.0
# Default graph retriever instance (can be overridden)
_default_graph_retriever: GraphRetriever | None = None
@@ -73,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.
@@ -84,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.
@@ -118,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())
@@ -139,25 +173,394 @@ 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]
async def retrieve_semantic_bm25_combined(
conn,
query_emb_str: str,
query_text: str,
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.
Uses CTEs with window functions to get top-N results per fact type per method,
all in one database round-trip.
Args:
conn: Database connection
query_emb_str: Query embedding as string
query_text: Query text for BM25
bank_id: Bank ID
fact_types: List of fact types to retrieve
limit: Maximum results per method per fact type
Returns:
Dict mapping fact_type -> (semantic_results, bm25_results)
"""
import re
# Sanitize query text for BM25 (same as retrieve_bm25)
sanitized_text = re.sub(r"[^\w\s]", " ", query_text.lower())
tokens = [token for token in sanitized_text.split() if token]
# 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, tags,
1 - (embedding <=> $1::vector) AS similarity,
NULL::float AS bm25_score,
'semantic' AS source,
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY embedding <=> $1::vector) AS rn
FROM {fq_table("memory_units")}
WHERE bank_id = $2
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, tags,
similarity, bm25_score, source
FROM semantic_ranked
WHERE rn <= $4
""",
*params,
)
# Group by fact_type
result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = {
ft: ([], []) for ft in fact_types
}
for r in results:
row = dict(r)
ft = row.get("fact_type")
row.pop("source", None)
if ft in result_dict:
result_dict[ft][0].append(RetrievalResult.from_db_row(row))
return result_dict
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, tags,
1 - (embedding <=> $1::vector) AS similarity,
NULL::float AS bm25_score,
'semantic' AS source,
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY embedding <=> $1::vector) AS rn
FROM {fq_table("memory_units")}
WHERE bank_id = $2
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, tags,
NULL::float AS similarity,
ts_rank_cd(search_vector, to_tsquery('english', $5)) AS bm25_score,
'bm25' AS source,
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY ts_rank_cd(search_vector, to_tsquery('english', $5)) DESC) AS rn
FROM {fq_table("memory_units")}
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, 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, tags,
similarity, bm25_score, source
FROM bm25_ranked WHERE rn <= $4
)
SELECT * FROM semantic
UNION ALL
SELECT * FROM bm25
""",
*params,
)
# Group results by fact_type and source
result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = {ft: ([], []) for ft in fact_types}
for r in results:
row = dict(r)
source = row.pop("source", None)
ft = row.get("fact_type")
if ft in result_dict:
if source == "semantic":
result_dict[ft][0].append(RetrievalResult.from_db_row(row))
else:
result_dict[ft][1].append(RetrievalResult.from_db_row(row))
return result_dict
async def retrieve_temporal_combined(
conn,
query_emb_str: str,
bank_id: str,
fact_types: list[str],
start_date: datetime,
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.
Batches the entry point query using window functions to get top-N per fact type,
then runs spreading for each fact type.
Args:
conn: Database connection
query_emb_str: Query embedding as string
bank_id: Bank ID
fact_types: List of fact types to retrieve
start_date: Start of time range
end_date: End of time range
budget: Node budget for spreading per fact type
semantic_threshold: Minimum semantic similarity to include
Returns:
Dict mapping fact_type -> list of RetrievalResult
"""
from ..memory_engine import fq_table
# Ensure dates are timezone-aware
if start_date.tzinfo is None:
start_date = start_date.replace(tzinfo=UTC)
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, 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")}
WHERE bank_id = $2
AND fact_type = ANY($3)
AND embedding IS NOT NULL
AND (
(occurred_start IS NOT NULL AND occurred_end IS NOT NULL
AND occurred_start <= $5 AND occurred_end >= $4)
OR
(mentioned_at IS NOT NULL AND mentioned_at BETWEEN $4 AND $5)
OR
(occurred_start IS NOT NULL AND occurred_start BETWEEN $4 AND $5)
OR
(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, tags, similarity
FROM ranked_entries
WHERE rn <= 10
""",
*params,
)
if not entry_points:
return {ft: [] for ft in fact_types}
# Group entry points by fact type
entries_by_ft: dict[str, list] = {ft: [] for ft in fact_types}
for ep in entry_points:
ft = ep["fact_type"]
if ft in entries_by_ft:
entries_by_ft[ft].append(ep)
# Calculate shared temporal parameters
total_days = (end_date - start_date).total_seconds() / 86400
mid_date = start_date + (end_date - start_date) / 2
# Process each fact type (spreading needs to stay per fact type due to link filtering)
results_by_ft: dict[str, list[RetrievalResult]] = {}
for ft in fact_types:
ft_entry_points = entries_by_ft.get(ft, [])
if not ft_entry_points:
results_by_ft[ft] = []
continue
results = []
visited = set()
node_scores = {}
# Process entry points
for ep in ft_entry_points:
unit_id = str(ep["id"])
visited.add(unit_id)
# Calculate temporal proximity
best_date = None
if ep["occurred_start"] is not None and ep["occurred_end"] is not None:
best_date = ep["occurred_start"] + (ep["occurred_end"] - ep["occurred_start"]) / 2
elif ep["occurred_start"] is not None:
best_date = ep["occurred_start"]
elif ep["occurred_end"] is not None:
best_date = ep["occurred_end"]
elif ep["mentioned_at"] is not None:
best_date = ep["mentioned_at"]
if best_date:
days_from_mid = abs((best_date - mid_date).total_seconds() / 86400)
temporal_proximity = 1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
else:
temporal_proximity = 0.5
ep_result = RetrievalResult.from_db_row(dict(ep))
ep_result.temporal_score = temporal_proximity
ep_result.temporal_proximity = temporal_proximity
results.append(ep_result)
node_scores[unit_id] = (ep["similarity"], 1.0)
# Spreading through temporal links (same as single-fact-type version)
frontier = list(node_scores.keys())
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, mu.tags,
ml.weight, ml.link_type, ml.from_unit_id,
1 - (mu.embedding <=> $1::vector) AS similarity
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
WHERE ml.from_unit_id = ANY($2::uuid[])
AND ml.link_type IN ('temporal', 'causes', 'caused_by', 'enables', 'prevents')
AND ml.weight >= 0.1
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
""",
*spreading_params,
)
for n in neighbors:
neighbor_id = str(n["id"])
if neighbor_id in visited:
continue
visited.add(neighbor_id)
budget_remaining -= 1
parent_id = str(n["from_unit_id"])
_, parent_temporal_score = node_scores.get(parent_id, (0.5, 0.5))
neighbor_best_date = None
if n["occurred_start"] is not None and n["occurred_end"] is not None:
neighbor_best_date = n["occurred_start"] + (n["occurred_end"] - n["occurred_start"]) / 2
elif n["occurred_start"] is not None:
neighbor_best_date = n["occurred_start"]
elif n["occurred_end"] is not None:
neighbor_best_date = n["occurred_end"]
elif n["mentioned_at"] is not None:
neighbor_best_date = n["mentioned_at"]
if neighbor_best_date:
days_from_mid = abs((neighbor_best_date - mid_date).total_seconds() / 86400)
neighbor_temporal_proximity = (
1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
)
else:
neighbor_temporal_proximity = 0.3
link_type = n["link_type"]
if link_type in ("causes", "caused_by"):
causal_boost = 2.0
elif link_type in ("enables", "prevents"):
causal_boost = 1.5
else:
causal_boost = 1.0
propagated_temporal = parent_temporal_score * n["weight"] * causal_boost * 0.7
combined_temporal = max(neighbor_temporal_proximity, propagated_temporal)
neighbor_result = RetrievalResult.from_db_row(dict(n))
neighbor_result.temporal_score = combined_temporal
neighbor_result.temporal_proximity = neighbor_temporal_proximity
results.append(neighbor_result)
if budget_remaining > 0 and combined_temporal > 0.2:
node_scores[neighbor_id] = (n["similarity"], combined_temporal)
frontier.append(neighbor_id)
if budget_remaining <= 0:
break
results_by_ft[ft] = results
return results_by_ft
async def retrieve_temporal(
conn,
query_emb_str: str,
@@ -167,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.
@@ -185,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
@@ -196,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
@@ -218,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:
@@ -378,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).
@@ -393,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
@@ -413,6 +823,7 @@ async def retrieve_parallel(
retriever,
question_date,
query_analyzer,
tags=tags,
)
else:
# For BFS, extract temporal constraint upfront (legacy path)
@@ -423,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,
)
@@ -447,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.
@@ -468,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:
@@ -477,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]:
@@ -495,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
@@ -666,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
@@ -673,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:
@@ -691,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)
@@ -706,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)
@@ -748,3 +1176,171 @@ async def _retrieve_parallel_bfs(
},
temporal_constraint=None,
)
async def retrieve_all_fact_types_parallel(
pool,
query_text: str,
query_embedding_str: str,
bank_id: str,
fact_types: list[str],
thinking_budget: int,
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.
This reduces database round-trips by:
1. Combining semantic + BM25 into one CTE query for ALL fact types (1 query instead of 2N)
2. Running graph retrieval per fact type in parallel (N parallel tasks)
3. Running temporal retrieval per fact type in parallel (N parallel tasks)
Args:
pool: Database connection pool
query_text: Query text
query_embedding_str: Query embedding as string
bank_id: Bank ID
fact_types: List of fact types to retrieve
thinking_budget: Budget for graph traversal and retrieval limits
question_date: Optional date when question was asked (for temporal filtering)
query_analyzer: Query analyzer to use (defaults to TransformerQueryAnalyzer)
graph_retriever: Graph retrieval strategy (defaults to configured retriever)
Returns:
MultiFactTypeRetrievalResult with results organized by fact type
"""
import time
retriever = graph_retriever or get_default_graph_retriever()
start_time = time.time()
timings: dict[str, float] = {}
# Step 1: Extract temporal constraint first (CPU work, no DB)
# Do this before DB queries so we know if we need temporal retrieval
temporal_extraction_start = time.time()
from .temporal_extraction import extract_temporal_constraint
temporal_constraint = extract_temporal_constraint(query_text, reference_date=question_date, analyzer=query_analyzer)
temporal_extraction_time = time.time() - temporal_extraction_start
timings["temporal_extraction"] = temporal_extraction_time
# Step 2: Run semantic + BM25 + temporal combined in ONE connection!
# This reduces connection usage from 2 to 1 for these operations
semantic_bm25_start = time.time()
temporal_results_by_ft: dict[str, list[RetrievalResult]] = {}
temporal_time = 0.0
async with acquire_with_retry(pool) as conn:
conn_wait = time.time() - semantic_bm25_start
# Semantic + BM25 combined
semantic_bm25_results = await retrieve_semantic_bm25_combined(
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
# Temporal combined (if constraint detected) - same connection!
if temporal_constraint:
tc_start, tc_end = temporal_constraint
temporal_start = time.time()
temporal_results_by_ft = await retrieve_temporal_combined(
conn,
query_embedding_str,
bank_id,
fact_types,
tc_start,
tc_end,
budget=thinking_budget,
semantic_threshold=0.1,
tags=tags,
tags_match=tags_match,
)
temporal_time = time.time() - temporal_start
timings["semantic_bm25_combined"] = semantic_bm25_time
timings["temporal_combined"] = temporal_time
# Step 3: Run graph retrieval for each fact type in parallel
async def run_graph_for_fact_type(ft: str) -> tuple[str, list[RetrievalResult], float, MPFPTimings | None]:
graph_start = time.time()
results, mpfp_timing = await retriever.retrieve(
pool=pool,
query_embedding_str=query_embedding_str,
bank_id=bank_id,
fact_type=ft,
budget=thinking_budget,
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
# Run graph for all fact types in parallel
graph_tasks = [run_graph_for_fact_type(ft) for ft in fact_types]
graph_results_list = await asyncio.gather(*graph_tasks)
# Organize results by fact type
results_by_fact_type: dict[str, ParallelRetrievalResult] = {}
max_conn_wait = conn_wait # Single connection for semantic+bm25+temporal
all_mpfp_timings: list[MPFPTimings] = []
for ft in fact_types:
# Get semantic + bm25 results for this fact type
semantic_results, bm25_results = semantic_bm25_results.get(ft, ([], []))
# Find graph results for this fact type
graph_results = []
graph_time = 0.0
mpfp_timing = None
for gr in graph_results_list:
if gr[0] == ft:
graph_results = gr[1]
graph_time = gr[2]
mpfp_timing = gr[3]
if mpfp_timing:
all_mpfp_timings.append(mpfp_timing)
break
# Get temporal results for this fact type from combined result
temporal_results = temporal_results_by_ft.get(ft) if temporal_constraint else None
if temporal_results is not None and len(temporal_results) == 0:
temporal_results = None
results_by_fact_type[ft] = ParallelRetrievalResult(
semantic=semantic_results,
bm25=bm25_results,
graph=graph_results,
temporal=temporal_results,
timings={
"semantic": semantic_bm25_time / 2, # Approximate split
"bm25": semantic_bm25_time / 2,
"graph": graph_time,
"temporal": temporal_time, # Same for all fact types (single query)
"temporal_extraction": temporal_extraction_time,
},
temporal_constraint=temporal_constraint,
mpfp_timings=[mpfp_timing] if mpfp_timing else [],
max_conn_wait=max_conn_wait,
)
total_time = time.time() - start_time
timings["total"] = total_time
return MultiFactTypeRetrievalResult(
results_by_fact_type=results_by_fact_type,
timings=timings,
max_conn_wait=max_conn_wait,
)
@@ -0,0 +1,172 @@
"""
Tags filtering utilities for retrieval.
Provides SQL building functions for filtering memories by tags.
Supports four matching modes via TagsMatch enum:
- "any": OR matching, includes untagged memories (default, backward compatible)
- "all": AND matching, includes untagged memories
- "any_strict": OR matching, excludes untagged memories
- "all_strict": AND matching, excludes untagged memories
OR matching (any/any_strict): Memory matches if ANY of its tags overlap with request tags
AND matching (all/all_strict): Memory matches if ALL request tags are present in its tags
"""
from typing import Literal
TagsMatch = Literal["any", "all", "any_strict", "all_strict"]
def _parse_tags_match(match: TagsMatch) -> tuple[str, bool]:
"""
Parse TagsMatch into operator and include_untagged flag.
Returns:
Tuple of (operator, include_untagged)
- operator: "&&" for any/any_strict, "@>" for all/all_strict
- include_untagged: True for any/all, False for any_strict/all_strict
"""
if match == "any":
return "&&", True
elif match == "all":
return "@>", True
elif match == "any_strict":
return "&&", False
elif match == "all_strict":
return "@>", False
else:
# Default to "any" behavior
return "&&", True
def build_tags_where_clause(
tags: list[str] | None,
param_offset: int = 1,
table_alias: str = "",
match: TagsMatch = "any",
) -> tuple[str, list, int]:
"""
Build a SQL WHERE clause for filtering by tags.
Supports four matching modes:
- "any" (default): OR matching, includes untagged memories
- "all": AND matching, includes untagged memories
- "any_strict": OR matching, excludes untagged memories
- "all_strict": AND matching, excludes untagged memories
Args:
tags: List of tags to filter by. If None or empty, returns empty clause (no filtering).
param_offset: Starting parameter number for SQL placeholders (default 1).
table_alias: Optional table alias prefix (e.g., "mu." for "memory_units mu").
match: Matching mode. Defaults to "any".
Returns:
Tuple of (sql_clause, params, next_param_offset):
- sql_clause: SQL WHERE clause string
- params: List of parameter values to bind
- next_param_offset: Next available parameter number
Example:
>>> clause, params, next_offset = build_tags_where_clause(['user_a'], 3, 'mu.', 'any_strict')
>>> print(clause) # "AND mu.tags IS NOT NULL AND mu.tags != '{}' AND mu.tags && $3"
"""
if not tags:
return "", [], param_offset
column = f"{table_alias}tags" if table_alias else "tags"
operator, include_untagged = _parse_tags_match(match)
if include_untagged:
# Include untagged memories (NULL or empty array) OR matching tags
clause = f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_offset})"
else:
# Strict: only memories with matching tags (exclude NULL and empty)
clause = f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_offset}"
return clause, [tags], param_offset + 1
def build_tags_where_clause_simple(
tags: list[str] | None,
param_num: int,
table_alias: str = "",
match: TagsMatch = "any",
) -> str:
"""
Build a simple SQL WHERE clause for tags filtering.
This is a convenience version that returns just the clause string,
assuming the caller will add the tags array to their params list.
Args:
tags: List of tags to filter by. If None or empty, returns empty string.
param_num: Parameter number to use in the clause.
table_alias: Optional table alias prefix.
match: Matching mode. Defaults to "any".
Returns:
SQL clause string or empty string.
"""
if not tags:
return ""
column = f"{table_alias}tags" if table_alias else "tags"
operator, include_untagged = _parse_tags_match(match)
if include_untagged:
# Include untagged memories (NULL or empty array) OR matching tags
return f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_num})"
else:
# Strict: only memories with matching tags (exclude NULL and empty)
return f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_num}"
def filter_results_by_tags(
results: list,
tags: list[str] | None,
match: TagsMatch = "any",
) -> list:
"""
Filter retrieval results by tags in Python (for post-processing).
Used when SQL filtering isn't possible (e.g., graph traversal results).
Args:
results: List of RetrievalResult objects with a 'tags' attribute.
tags: List of tags to filter by. If None or empty, returns all results.
match: Matching mode. Defaults to "any".
Returns:
Filtered list of results.
"""
if not tags:
return results
_, include_untagged = _parse_tags_match(match)
is_any_match = match in ("any", "any_strict")
tags_set = set(tags)
filtered = []
for result in results:
result_tags = getattr(result, "tags", None)
# Check if untagged
is_untagged = result_tags is None or len(result_tags) == 0
if is_untagged:
if include_untagged:
filtered.append(result)
# else: skip untagged
else:
result_tags_set = set(result_tags)
if is_any_match:
# Any overlap
if result_tags_set & tags_set:
filtered.append(result)
else:
# All tags must be present
if tags_set <= result_tags_set:
filtered.append(result)
return filtered
@@ -11,6 +11,13 @@ from typing import Any, Literal
from pydantic import BaseModel, Field
class TemporalConstraint(BaseModel):
"""Detected temporal constraint from query analysis."""
start: datetime | None = Field(default=None, description="Start of temporal range")
end: datetime | None = Field(default=None, description="End of temporal range")
class QueryInfo(BaseModel):
"""Information about the search query."""
@@ -19,6 +26,11 @@ class QueryInfo(BaseModel):
timestamp: datetime = Field(description="When the query was executed")
budget: int = Field(description="Maximum nodes to explore")
max_tokens: int = Field(description="Maximum tokens to return in results")
tags: list[str] | None = Field(default=None, description="Tags filter applied to recall")
tags_match: str | None = Field(default=None, description="Tags matching mode: any, all, any_strict, all_strict")
temporal_constraint: TemporalConstraint | None = Field(
default=None, description="Detected temporal range from query"
)
class EntryPoint(BaseModel):
@@ -22,6 +22,7 @@ from .trace import (
SearchPhaseMetrics,
SearchSummary,
SearchTrace,
TemporalConstraint,
WeightComponents,
)
@@ -45,7 +46,14 @@ class SearchTracer:
json_output = trace.to_json()
"""
def __init__(self, query: str, budget: int, max_tokens: int):
def __init__(
self,
query: str,
budget: int,
max_tokens: int,
tags: list[str] | None = None,
tags_match: str | None = None,
):
"""
Initialize tracer.
@@ -53,10 +61,14 @@ class SearchTracer:
query: Search query text
budget: Maximum nodes to explore
max_tokens: Maximum tokens to return in results
tags: Tags filter applied to recall
tags_match: Tags matching mode (any, all, any_strict, all_strict)
"""
self.query_text = query
self.budget = budget
self.max_tokens = max_tokens
self.tags = tags
self.tags_match = tags_match
# Trace data
self.query_embedding: list[float] | None = None
@@ -66,6 +78,9 @@ class SearchTracer:
self.pruned: list[PruningDecision] = []
self.phase_metrics: list[SearchPhaseMetrics] = []
# Temporal constraint detected from query
self.temporal_constraint: TemporalConstraint | None = None
# New 4-way retrieval tracking
self.retrieval_results: list[RetrievalMethodResults] = []
self.rrf_merged: list[RRFMergeResult] = []
@@ -88,6 +103,11 @@ class SearchTracer:
"""Record the query embedding."""
self.query_embedding = embedding
def record_temporal_constraint(self, start: datetime | None, end: datetime | None):
"""Record the detected temporal constraint from query analysis."""
if start is not None or end is not None:
self.temporal_constraint = TemporalConstraint(start=start, end=end)
def add_entry_point(self, node_id: str, text: str, similarity: float, rank: int):
"""
Record an entry point.
@@ -428,6 +448,9 @@ class SearchTracer:
timestamp=datetime.now(UTC),
budget=self.budget,
max_tokens=self.max_tokens,
tags=self.tags,
tags_match=self.tags_match,
temporal_constraint=self.temporal_constraint,
)
# Create summary
@@ -48,6 +48,7 @@ class RetrievalResult:
chunk_id: str | None = None
access_count: int = 0
embedding: list[float] | None = None
tags: list[str] | None = None # Visibility scope tags
# Retrieval-specific scores (only one will be set depending on retrieval method)
similarity: float | None = None # Semantic retrieval
@@ -72,6 +73,7 @@ class RetrievalResult:
chunk_id=row.get("chunk_id"),
access_count=row.get("access_count", 0),
embedding=row.get("embedding"),
tags=row.get("tags"),
similarity=row.get("similarity"),
bm25_score=row.get("bm25_score"),
activation=row.get("activation"),
@@ -156,6 +158,7 @@ class ScoredResult:
"chunk_id": self.retrieval.chunk_id,
"access_count": self.retrieval.access_count,
"embedding": self.retrieval.embedding,
"tags": self.retrieval.tags,
"semantic_similarity": self.retrieval.similarity,
"bm25_score": self.retrieval.bm25_score,
}
+4
View File
@@ -187,12 +187,15 @@ def main():
embeddings_provider=config.embeddings_provider,
embeddings_local_model=config.embeddings_local_model,
embeddings_tei_url=config.embeddings_tei_url,
embeddings_openai_base_url=config.embeddings_openai_base_url,
embeddings_cohere_base_url=config.embeddings_cohere_base_url,
reranker_provider=config.reranker_provider,
reranker_local_model=config.reranker_local_model,
reranker_tei_url=config.reranker_tei_url,
reranker_tei_batch_size=config.reranker_tei_batch_size,
reranker_tei_max_concurrent=config.reranker_tei_max_concurrent,
reranker_max_candidates=config.reranker_max_candidates,
reranker_cohere_base_url=config.reranker_cohere_base_url,
host=args.host,
port=args.port,
log_level=args.log_level,
@@ -200,6 +203,7 @@ def main():
graph_retriever=config.graph_retriever,
mpfp_top_k_neighbors=config.mpfp_top_k_neighbors,
recall_max_concurrent=config.recall_max_concurrent,
recall_connection_budget=config.recall_connection_budget,
observation_min_facts=config.observation_min_facts,
observation_top_entities=config.observation_top_entities,
retain_max_completion_tokens=config.retain_max_completion_tokens,
+15
View File
@@ -28,6 +28,15 @@ from opentelemetry.sdk.resources import Resource
if TYPE_CHECKING:
import asyncpg
def _get_tenant() -> str:
"""Get current tenant (schema) from context for metrics labeling."""
# Import here to avoid circular imports
from hindsight_api.engine.memory_engine import get_current_schema
return get_current_schema()
# Custom bucket boundaries for operation duration (in seconds)
# Fine granularity in 0-30s range where most operations complete
DURATION_BUCKETS = (0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 20.0, 30.0, 60.0, 120.0)
@@ -323,6 +332,7 @@ class MetricsCollector(MetricsCollectorBase):
"operation": operation,
"bank_id": bank_id,
"source": source,
"tenant": _get_tenant(),
}
if budget:
attributes["budget"] = budget
@@ -373,6 +383,7 @@ class MetricsCollector(MetricsCollectorBase):
"model": model,
"scope": scope,
"success": str(success).lower(),
"tenant": _get_tenant(),
}
# Record duration
@@ -425,10 +436,14 @@ class MetricsCollector(MetricsCollectorBase):
status_code = status_code_getter()
status_class = f"{status_code // 100}xx"
# Get tenant from context (may be set during request processing)
tenant = _get_tenant()
attributes = {
**base_attributes,
"status_code": str(status_code),
"status_class": status_class,
"tenant": tenant,
}
# Record duration and count
+13 -1
View File
@@ -22,6 +22,7 @@ from pathlib import Path
from alembic import command
from alembic.config import Config
from alembic.script.revision import ResolutionError
from sqlalchemy import create_engine, text
logger = logging.getLogger(__name__)
@@ -78,7 +79,18 @@ def _run_migrations_internal(database_url: str, script_location: str, schema: st
alembic_cfg.set_main_option("target_schema", schema)
# Run migrations
command.upgrade(alembic_cfg, "head")
try:
command.upgrade(alembic_cfg, "head")
except ResolutionError as e:
# This happens during rolling deployments when a newer version of the code
# has already run migrations, and this older replica doesn't have the new
# migration files. The database is already at a newer revision than we know.
# This is safe to ignore - the newer code has already applied its migrations.
logger.warning(
f"Database is at a newer migration revision than this code version knows about. "
f"This is expected during rolling deployments. Skipping migrations. Error: {e}"
)
return
logger.info(f"Database migrations completed successfully for schema '{schema_name}'")
+31 -1
View File
@@ -7,6 +7,7 @@ This module provides the ASGI app for uvicorn import string usage:
For CLI usage, use the hindsight-api command instead.
"""
import logging
import os
import warnings
@@ -17,6 +18,12 @@ warnings.filterwarnings("ignore", message="websockets.server.WebSocketServerProt
from hindsight_api import MemoryEngine
from hindsight_api.api import create_app
from hindsight_api.config import get_config
from hindsight_api.extensions import (
DefaultExtensionContext,
OperationValidatorExtension,
TenantExtension,
load_extension,
)
# Disable tokenizers parallelism to avoid warnings
os.environ["TOKENIZERS_PARALLELISM"] = "false"
@@ -25,10 +32,33 @@ os.environ["TOKENIZERS_PARALLELISM"] = "false"
config = get_config()
config.configure_logging()
# Load operation validator extension if configured
operation_validator = load_extension("OPERATION_VALIDATOR", OperationValidatorExtension)
if operation_validator:
logging.info(f"Loaded operation validator: {operation_validator.__class__.__name__}")
# Load tenant extension if configured
tenant_extension = load_extension("TENANT", TenantExtension)
if tenant_extension:
logging.info(f"Loaded tenant extension: {tenant_extension.__class__.__name__}")
# Create app at module level (required for uvicorn import string)
# MemoryEngine reads configuration from environment variables automatically
# Note: run_migrations=True by default, but migrations are idempotent so safe with workers
_memory = MemoryEngine(run_migrations=config.run_migrations_on_startup)
_memory = MemoryEngine(
operation_validator=operation_validator,
tenant_extension=tenant_extension,
run_migrations=config.run_migrations_on_startup,
)
# Set extension context on tenant extension (needed for schema provisioning)
if tenant_extension:
extension_context = DefaultExtensionContext(
database_url=config.database_url,
memory_engine=_memory,
)
tenant_extension.set_context(extension_context)
logging.info("Extension context set on tenant extension")
# Create unified app with both HTTP and optionally MCP
app = create_app(
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "hindsight-api"
version = "0.2.1"
version = "0.3.0"
description = "Hindsight: Agent Memory That Works Like Human Memory"
readme = "README.md"
requires-python = ">=3.11"
+396
View File
@@ -0,0 +1,396 @@
"""
Tests for hindsight_api.main module (single-worker code path).
The main.py module is used when running with a single worker:
hindsight-api (or hindsight-api --workers 1)
When workers=1, main.py creates the app directly and passes it to uvicorn.
These tests ensure that extensions are properly loaded in this code path.
Compare with test_server_module.py which tests the multi-worker path (workers > 1).
"""
import sys
from unittest.mock import MagicMock, patch
class TestMainModuleExtensionLoading:
"""Tests that main.py correctly loads extensions when configured via environment."""
def test_main_loads_tenant_extension_when_configured(self, monkeypatch):
"""
Verify that main.py loads tenant extension from HINDSIGHT_API_TENANT_EXTENSION.
This ensures extension loading works in the single-worker code path.
"""
# Set up environment to configure a tenant extension
monkeypatch.setenv(
"HINDSIGHT_API_TENANT_EXTENSION",
"tests.test_main_module:MockTenantExtension",
)
# Ensure single worker mode
monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1")
# Track what extensions were loaded via load_extension
loaded_extensions = {}
# Get the real load_extension function
from hindsight_api.extensions.loader import load_extension as real_load_extension
def tracking_load_extension(name, base_class):
"""Track calls to load_extension and delegate to original."""
result = real_load_extension(name, base_class)
loaded_extensions[name] = result
return result
with patch("hindsight_api.main.MemoryEngine") as mock_engine, \
patch("hindsight_api.main.create_app") as mock_create_app, \
patch("hindsight_api.main.get_config") as mock_get_config, \
patch("hindsight_api.main.load_extension", side_effect=tracking_load_extension), \
patch("hindsight_api.main.DefaultExtensionContext"), \
patch("hindsight_api.main.print_banner"), \
patch("uvicorn.run"): # Don't actually start uvicorn
mock_config = MagicMock()
mock_config.host = "0.0.0.0"
mock_config.port = 8888
mock_config.log_level = "info"
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_engine.return_value = MagicMock()
mock_create_app.return_value = MagicMock()
# Mock sys.argv to simulate CLI invocation
with patch.object(sys, 'argv', ['hindsight-api']):
from hindsight_api.main import main
main()
# Verify TENANT extension was loaded
assert "TENANT" in loaded_extensions, \
"main.py did not call load_extension('TENANT', ...) - extensions not loaded!"
assert loaded_extensions["TENANT"] is not None, \
"load_extension('TENANT', ...) returned None despite env var being set"
assert isinstance(loaded_extensions["TENANT"], MockTenantExtension), \
f"Expected MockTenantExtension, got {type(loaded_extensions['TENANT'])}"
def test_main_loads_operation_validator_when_configured(self, monkeypatch):
"""
Verify that main.py loads operation validator from HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION.
"""
monkeypatch.setenv(
"HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION",
"tests.test_main_module:MockOperationValidator",
)
monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1")
loaded_extensions = {}
from hindsight_api.extensions.loader import load_extension as real_load_extension
def tracking_load_extension(name, base_class):
result = real_load_extension(name, base_class)
loaded_extensions[name] = result
return result
with patch("hindsight_api.main.MemoryEngine") as mock_engine, \
patch("hindsight_api.main.create_app") as mock_create_app, \
patch("hindsight_api.main.get_config") as mock_get_config, \
patch("hindsight_api.main.load_extension", side_effect=tracking_load_extension), \
patch("hindsight_api.main.DefaultExtensionContext"), \
patch("hindsight_api.main.print_banner"), \
patch("uvicorn.run"):
mock_config = MagicMock()
mock_config.host = "0.0.0.0"
mock_config.port = 8888
mock_config.log_level = "info"
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_engine.return_value = MagicMock()
mock_create_app.return_value = MagicMock()
with patch.object(sys, 'argv', ['hindsight-api']):
from hindsight_api.main import main
main()
assert "OPERATION_VALIDATOR" in loaded_extensions, \
"main.py did not call load_extension('OPERATION_VALIDATOR', ...)"
assert loaded_extensions["OPERATION_VALIDATOR"] is not None
assert isinstance(loaded_extensions["OPERATION_VALIDATOR"], MockOperationValidator)
def test_main_passes_extensions_to_memory_engine(self, monkeypatch):
"""
Verify that main.py passes loaded extensions to MemoryEngine constructor.
This is the critical test - even if extensions are loaded, they must be
passed to MemoryEngine for authentication to work.
"""
monkeypatch.setenv(
"HINDSIGHT_API_TENANT_EXTENSION",
"tests.test_main_module:MockTenantExtension",
)
monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1")
memory_engine_calls = []
def capture_memory_engine(*args, **kwargs):
memory_engine_calls.append({"args": args, "kwargs": kwargs})
return MagicMock()
with patch("hindsight_api.main.MemoryEngine", side_effect=capture_memory_engine), \
patch("hindsight_api.main.create_app") as mock_create_app, \
patch("hindsight_api.main.get_config") as mock_get_config, \
patch("hindsight_api.main.DefaultExtensionContext"), \
patch("hindsight_api.main.print_banner"), \
patch("uvicorn.run"):
mock_config = MagicMock()
mock_config.host = "0.0.0.0"
mock_config.port = 8888
mock_config.log_level = "info"
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_create_app.return_value = MagicMock()
with patch.object(sys, 'argv', ['hindsight-api']):
from hindsight_api.main import main
main()
# Verify MemoryEngine was called
assert len(memory_engine_calls) == 1, "MemoryEngine should be called exactly once"
call_kwargs = memory_engine_calls[0]["kwargs"]
# THE CRITICAL ASSERTION: tenant_extension must be passed and not None
assert "tenant_extension" in call_kwargs, \
"MemoryEngine was not called with tenant_extension parameter!"
assert call_kwargs["tenant_extension"] is not None, \
"tenant_extension was None - main.py did not pass loaded extension to MemoryEngine!"
def test_main_sets_extension_context_on_tenant_extension(self, monkeypatch):
"""
Verify that main.py sets the extension context on tenant extension.
This is required for tenant extensions that need to provision schemas.
"""
monkeypatch.setenv(
"HINDSIGHT_API_TENANT_EXTENSION",
"tests.test_main_module:MockTenantExtension",
)
monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1")
captured_tenant_ext = [None]
def capture_memory_engine(*args, **kwargs):
captured_tenant_ext[0] = kwargs.get("tenant_extension")
return MagicMock()
context_created = []
def capture_context(*args, **kwargs):
ctx = MagicMock()
context_created.append(ctx)
return ctx
with patch("hindsight_api.main.MemoryEngine", side_effect=capture_memory_engine), \
patch("hindsight_api.main.create_app") as mock_create_app, \
patch("hindsight_api.main.get_config") as mock_get_config, \
patch("hindsight_api.main.DefaultExtensionContext", side_effect=capture_context), \
patch("hindsight_api.main.print_banner"), \
patch("uvicorn.run"):
mock_config = MagicMock()
mock_config.host = "0.0.0.0"
mock_config.port = 8888
mock_config.log_level = "info"
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_create_app.return_value = MagicMock()
with patch.object(sys, 'argv', ['hindsight-api']):
from hindsight_api.main import main
main()
# Verify context was created and set
assert len(context_created) == 1, "DefaultExtensionContext should be created"
assert captured_tenant_ext[0] is not None, "Tenant extension should be captured"
assert captured_tenant_ext[0]._context_set, \
"set_context was not called on tenant extension"
def test_main_works_without_extensions(self, monkeypatch):
"""
Verify that main.py works correctly when no extensions are configured.
"""
# Ensure no extension env vars are set
monkeypatch.delenv("HINDSIGHT_API_TENANT_EXTENSION", raising=False)
monkeypatch.delenv("HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION", raising=False)
monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1")
memory_engine_calls = []
def capture_memory_engine(*args, **kwargs):
memory_engine_calls.append({"args": args, "kwargs": kwargs})
return MagicMock()
with patch("hindsight_api.main.MemoryEngine", side_effect=capture_memory_engine), \
patch("hindsight_api.main.create_app") as mock_create_app, \
patch("hindsight_api.main.get_config") as mock_get_config, \
patch("hindsight_api.main.print_banner"), \
patch("uvicorn.run"):
mock_config = MagicMock()
mock_config.host = "0.0.0.0"
mock_config.port = 8888
mock_config.log_level = "info"
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_create_app.return_value = MagicMock()
with patch.object(sys, 'argv', ['hindsight-api']):
from hindsight_api.main import main
main()
# Should work without extensions
assert len(memory_engine_calls) == 1
call_kwargs = memory_engine_calls[0]["kwargs"]
# Extensions should be None when not configured
assert call_kwargs.get("tenant_extension") is None
assert call_kwargs.get("operation_validator") is None
def test_main_uses_app_object_for_single_worker(self, monkeypatch):
"""
Verify that main.py passes the app object (not import string) when workers=1.
This is important because it means single-worker mode uses the app created
in main.py (with extensions loaded), not server.py.
"""
monkeypatch.setenv("HINDSIGHT_API_WORKERS", "1")
monkeypatch.delenv("HINDSIGHT_API_TENANT_EXTENSION", raising=False)
uvicorn_calls = []
def capture_uvicorn_run(**kwargs):
uvicorn_calls.append(kwargs)
mock_app = MagicMock()
with patch("hindsight_api.main.MemoryEngine") as mock_engine, \
patch("hindsight_api.main.create_app", return_value=mock_app), \
patch("hindsight_api.main.get_config") as mock_get_config, \
patch("hindsight_api.main.print_banner"), \
patch("uvicorn.run", side_effect=capture_uvicorn_run):
mock_config = MagicMock()
mock_config.host = "0.0.0.0"
mock_config.port = 8888
mock_config.log_level = "info"
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_engine.return_value = MagicMock()
with patch.object(sys, 'argv', ['hindsight-api', '--workers', '1']):
from hindsight_api.main import main
main()
assert len(uvicorn_calls) == 1
# With workers=1, should pass app object, not import string
assert uvicorn_calls[0]["app"] is mock_app, \
"main.py should pass app object (not import string) when workers=1"
def test_main_uses_import_string_for_multiple_workers(self, monkeypatch):
"""
Verify that main.py uses import string when workers > 1.
This is important because multi-worker mode requires server.py to be imported
by each worker process.
"""
monkeypatch.setenv("HINDSIGHT_API_WORKERS", "2")
monkeypatch.delenv("HINDSIGHT_API_TENANT_EXTENSION", raising=False)
uvicorn_calls = []
def capture_uvicorn_run(**kwargs):
uvicorn_calls.append(kwargs)
with patch("hindsight_api.main.MemoryEngine") as mock_engine, \
patch("hindsight_api.main.create_app") as mock_create_app, \
patch("hindsight_api.main.get_config") as mock_get_config, \
patch("hindsight_api.main.print_banner"), \
patch("uvicorn.run", side_effect=capture_uvicorn_run):
mock_config = MagicMock()
mock_config.host = "0.0.0.0"
mock_config.port = 8888
mock_config.log_level = "info"
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_engine.return_value = MagicMock()
mock_create_app.return_value = MagicMock()
with patch.object(sys, 'argv', ['hindsight-api', '--workers', '2']):
from hindsight_api.main import main
main()
assert len(uvicorn_calls) == 1
# With workers > 1, should use import string
assert uvicorn_calls[0]["app"] == "hindsight_api.server:app", \
"main.py should use import string when workers > 1"
assert uvicorn_calls[0]["workers"] == 2
# Mock extensions for testing
from hindsight_api.extensions import (
TenantExtension,
TenantContext,
RequestContext,
OperationValidatorExtension,
ValidationResult,
RetainContext,
RecallContext,
ReflectContext,
)
class MockTenantExtension(TenantExtension):
"""Mock tenant extension for testing main.py extension loading."""
def __init__(self, config: dict):
super().__init__(config)
self._context_set = False
async def authenticate(self, request_context: RequestContext) -> TenantContext:
return TenantContext(schema_name="public")
def set_context(self, context) -> None:
self._context_set = True
class MockOperationValidator(OperationValidatorExtension):
"""Mock operation validator for testing main.py extension loading."""
def __init__(self, config: dict):
super().__init__(config)
async def validate_retain(self, ctx: RetainContext) -> ValidationResult:
return ValidationResult.accept()
async def validate_recall(self, ctx: RecallContext) -> ValidationResult:
return ValidationResult.accept()
async def validate_reflect(self, ctx: ReflectContext) -> ValidationResult:
return ValidationResult.accept()
+290
View File
@@ -0,0 +1,290 @@
"""
Tests for hindsight_api.server module (multi-worker code path).
The server.py module is used when running with multiple workers:
uvicorn hindsight_api.server:app --workers 2
This module executes code at import time, creating the app at module level.
These tests ensure that extensions are properly loaded in this code path,
which was previously a regression that caused authentication bypass in production.
"""
import importlib
import sys
from unittest.mock import MagicMock, patch
def _clean_server_module():
"""Remove hindsight_api.server from sys.modules for fresh import."""
modules_to_remove = [k for k in sys.modules.keys() if k.startswith("hindsight_api.server")]
for mod in modules_to_remove:
del sys.modules[mod]
class TestServerModuleExtensionLoading:
"""Tests that server.py correctly loads extensions when configured via environment."""
def test_server_loads_tenant_extension_when_configured(self, monkeypatch):
"""
Verify that server.py loads tenant extension from HINDSIGHT_API_TENANT_EXTENSION.
This test catches the regression where server.py didn't call load_extension(),
causing authentication to be bypassed in multi-worker deployments.
"""
# Set up environment to configure a tenant extension
monkeypatch.setenv(
"HINDSIGHT_API_TENANT_EXTENSION",
"tests.test_server_module:MockTenantExtension",
)
_clean_server_module()
# Track what extensions were loaded via load_extension
loaded_extensions = {}
# Get the real load_extension function
from hindsight_api.extensions.loader import load_extension as real_load_extension
def tracking_load_extension(name, base_class):
"""Track calls to load_extension and delegate to original."""
result = real_load_extension(name, base_class)
loaded_extensions[name] = result
return result
# Patch at source level BEFORE importing server
# Note: We patch the entire hindsight_api module namespace
with patch("hindsight_api.MemoryEngine") as mock_engine, \
patch("hindsight_api.api.create_app") as mock_create_app, \
patch("hindsight_api.config.get_config") as mock_get_config, \
patch("hindsight_api.extensions.load_extension", side_effect=tracking_load_extension), \
patch("hindsight_api.extensions.DefaultExtensionContext"):
mock_config = MagicMock()
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_engine.return_value = MagicMock()
mock_create_app.return_value = MagicMock()
# Now import server - this triggers module-level code
import hindsight_api.server
# Verify TENANT extension was loaded
assert "TENANT" in loaded_extensions, \
"server.py did not call load_extension('TENANT', ...) - extensions not loaded!"
assert loaded_extensions["TENANT"] is not None, \
"load_extension('TENANT', ...) returned None despite env var being set"
assert isinstance(loaded_extensions["TENANT"], MockTenantExtension), \
f"Expected MockTenantExtension, got {type(loaded_extensions['TENANT'])}"
def test_server_loads_operation_validator_when_configured(self, monkeypatch):
"""
Verify that server.py loads operation validator from HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION.
"""
monkeypatch.setenv(
"HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION",
"tests.test_server_module:MockOperationValidator",
)
_clean_server_module()
loaded_extensions = {}
from hindsight_api.extensions.loader import load_extension as real_load_extension
def tracking_load_extension(name, base_class):
result = real_load_extension(name, base_class)
loaded_extensions[name] = result
return result
with patch("hindsight_api.MemoryEngine") as mock_engine, \
patch("hindsight_api.api.create_app") as mock_create_app, \
patch("hindsight_api.config.get_config") as mock_get_config, \
patch("hindsight_api.extensions.load_extension", side_effect=tracking_load_extension), \
patch("hindsight_api.extensions.DefaultExtensionContext"):
mock_config = MagicMock()
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_engine.return_value = MagicMock()
mock_create_app.return_value = MagicMock()
import hindsight_api.server
assert "OPERATION_VALIDATOR" in loaded_extensions, \
"server.py did not call load_extension('OPERATION_VALIDATOR', ...)"
assert loaded_extensions["OPERATION_VALIDATOR"] is not None
assert isinstance(loaded_extensions["OPERATION_VALIDATOR"], MockOperationValidator)
def test_server_passes_extensions_to_memory_engine(self, monkeypatch):
"""
Verify that server.py passes loaded extensions to MemoryEngine constructor.
This is the critical test - even if extensions are loaded, they must be
passed to MemoryEngine for authentication to work.
"""
monkeypatch.setenv(
"HINDSIGHT_API_TENANT_EXTENSION",
"tests.test_server_module:MockTenantExtension",
)
_clean_server_module()
memory_engine_calls = []
def capture_memory_engine(*args, **kwargs):
memory_engine_calls.append({"args": args, "kwargs": kwargs})
return MagicMock()
with patch("hindsight_api.MemoryEngine", side_effect=capture_memory_engine), \
patch("hindsight_api.api.create_app") as mock_create_app, \
patch("hindsight_api.config.get_config") as mock_get_config, \
patch("hindsight_api.extensions.DefaultExtensionContext"):
mock_config = MagicMock()
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_create_app.return_value = MagicMock()
import hindsight_api.server
# Verify MemoryEngine was called
assert len(memory_engine_calls) == 1, "MemoryEngine should be called exactly once"
call_kwargs = memory_engine_calls[0]["kwargs"]
# THE CRITICAL ASSERTION: tenant_extension must be passed and not None
assert "tenant_extension" in call_kwargs, \
"MemoryEngine was not called with tenant_extension parameter!"
assert call_kwargs["tenant_extension"] is not None, \
"tenant_extension was None - server.py did not pass loaded extension to MemoryEngine!"
def test_server_sets_extension_context_on_tenant_extension(self, monkeypatch):
"""
Verify that server.py sets the extension context on tenant extension.
This is required for tenant extensions that need to provision schemas.
"""
monkeypatch.setenv(
"HINDSIGHT_API_TENANT_EXTENSION",
"tests.test_server_module:MockTenantExtension",
)
_clean_server_module()
context_set_calls = []
captured_tenant_ext = [None]
def capture_memory_engine(*args, **kwargs):
captured_tenant_ext[0] = kwargs.get("tenant_extension")
return MagicMock()
def capture_context(*args, **kwargs):
ctx = MagicMock()
context_set_calls.append(ctx)
return ctx
with patch("hindsight_api.MemoryEngine", side_effect=capture_memory_engine), \
patch("hindsight_api.api.create_app") as mock_create_app, \
patch("hindsight_api.config.get_config") as mock_get_config, \
patch("hindsight_api.extensions.DefaultExtensionContext", side_effect=capture_context):
mock_config = MagicMock()
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_create_app.return_value = MagicMock()
import hindsight_api.server
# Verify context was created and set
assert len(context_set_calls) == 1, "DefaultExtensionContext should be created"
assert captured_tenant_ext[0] is not None, "Tenant extension should be captured"
assert captured_tenant_ext[0]._context_set, \
"set_context was not called on tenant extension"
def test_server_works_without_extensions(self, monkeypatch):
"""
Verify that server.py works correctly when no extensions are configured.
"""
# Ensure no extension env vars are set
monkeypatch.delenv("HINDSIGHT_API_TENANT_EXTENSION", raising=False)
monkeypatch.delenv("HINDSIGHT_API_OPERATION_VALIDATOR_EXTENSION", raising=False)
_clean_server_module()
memory_engine_calls = []
def capture_memory_engine(*args, **kwargs):
memory_engine_calls.append({"args": args, "kwargs": kwargs})
return MagicMock()
with patch("hindsight_api.MemoryEngine", side_effect=capture_memory_engine), \
patch("hindsight_api.api.create_app") as mock_create_app, \
patch("hindsight_api.config.get_config") as mock_get_config:
mock_config = MagicMock()
mock_config.mcp_enabled = False
mock_config.run_migrations_on_startup = False
mock_config.database_url = "postgresql://test:test@localhost/test"
mock_get_config.return_value = mock_config
mock_create_app.return_value = MagicMock()
import hindsight_api.server
# Should work without extensions
assert len(memory_engine_calls) == 1
call_kwargs = memory_engine_calls[0]["kwargs"]
# Extensions should be None when not configured
assert call_kwargs.get("tenant_extension") is None
assert call_kwargs.get("operation_validator") is None
# Mock extensions for testing
from hindsight_api.extensions import (
TenantExtension,
TenantContext,
RequestContext,
OperationValidatorExtension,
ValidationResult,
RetainContext,
RecallContext,
ReflectContext,
)
class MockTenantExtension(TenantExtension):
"""Mock tenant extension for testing server.py extension loading."""
def __init__(self, config: dict):
super().__init__(config)
self._context_set = False
async def authenticate(self, request_context: RequestContext) -> TenantContext:
return TenantContext(schema_name="public")
def set_context(self, context) -> None:
self._context_set = True
class MockOperationValidator(OperationValidatorExtension):
"""Mock operation validator for testing server.py extension loading."""
def __init__(self, config: dict):
super().__init__(config)
async def validate_retain(self, ctx: RetainContext) -> ValidationResult:
return ValidationResult.accept()
async def validate_recall(self, ctx: RecallContext) -> ValidationResult:
return ValidationResult.accept()
async def validate_reflect(self, ctx: ReflectContext) -> ValidationResult:
return ValidationResult.accept()
+883
View File
@@ -0,0 +1,883 @@
"""
Tests for tags-based visibility scoping.
This module tests the tags feature which allows filtering memories by visibility tags.
Use cases:
- Multi-user agent: Agent has a single memory bank, users should only see memories from
conversations they participated in
- Student tracking: Teacher tracks students, students should only see their own data
The tags use OR-based matching: a memory matches if ANY of its tags overlap with the request tags.
"""
from datetime import datetime
import httpx
import pytest
import pytest_asyncio
from hindsight_api.api import create_app
from hindsight_api.engine.search.tags import build_tags_where_clause_simple, filter_results_by_tags
# ============================================================================
# Unit Tests for tags SQL builder
# ============================================================================
class TestTagsWhereClauseBuilder:
"""Unit tests for the tags WHERE clause SQL builder."""
def test_no_tags_returns_empty_string(self):
"""When tags is None, should return empty string (no filtering)."""
result = build_tags_where_clause_simple(None, 5)
assert result == ""
def test_empty_tags_list_returns_empty_string(self):
"""When tags is an empty list, should return empty string (no filtering)."""
result = build_tags_where_clause_simple([], 5)
assert result == ""
def test_tags_with_different_param_num(self):
"""Should use the provided parameter number."""
result = build_tags_where_clause_simple(["user_a", "user_b"], 3)
# Default is "any" which includes untagged
assert "$3" in result
def test_tags_with_table_alias(self):
"""Should include table alias when provided."""
result = build_tags_where_clause_simple(["user_a"], 5, table_alias="mu.")
assert "mu.tags" in result
# ---- Test "any" mode (OR, includes untagged - default) ----
def test_tags_match_any_includes_untagged(self):
"""When match='any', should include untagged memories (NULL or empty)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="any")
# Should use OR with NULL/empty check
assert "IS NULL" in result
assert "= '{}'" in result
assert "&&" in result # overlap operator
def test_tags_match_any_uses_overlap(self):
"""When match='any', should use overlap operator (&&)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="any")
assert "&&" in result
# ---- Test "all" mode (AND, includes untagged) ----
def test_tags_match_all_includes_untagged(self):
"""When match='all', should include untagged memories (NULL or empty)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="all")
# Should use OR with NULL/empty check
assert "IS NULL" in result
assert "= '{}'" in result
assert "@>" in result # contains operator
def test_tags_match_all_uses_contains(self):
"""When match='all', should use contains operator (@>)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="all")
assert "@>" in result
# ---- Test "any_strict" mode (OR, excludes untagged) ----
def test_tags_match_any_strict_excludes_untagged(self):
"""When match='any_strict', should exclude untagged memories."""
result = build_tags_where_clause_simple(["user_a"], 5, match="any_strict")
# Should require tags to be NOT NULL and not empty
assert "IS NOT NULL" in result
assert "!= '{}'" in result
assert "&&" in result # overlap operator
def test_tags_match_any_strict_uses_overlap(self):
"""When match='any_strict', should use overlap operator (&&)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="any_strict")
assert "&&" in result
# Should NOT include untagged
assert "IS NULL" not in result or "IS NOT NULL" in result
# ---- Test "all_strict" mode (AND, excludes untagged) ----
def test_tags_match_all_strict_excludes_untagged(self):
"""When match='all_strict', should exclude untagged memories."""
result = build_tags_where_clause_simple(["user_a"], 5, match="all_strict")
# Should require tags to be NOT NULL and not empty
assert "IS NOT NULL" in result
assert "!= '{}'" in result
assert "@>" in result # contains operator
def test_tags_match_all_strict_uses_contains(self):
"""When match='all_strict', should use contains operator (@>)."""
result = build_tags_where_clause_simple(["user_a"], 5, match="all_strict")
assert "@>" in result
# ---- Test table alias with all modes ----
def test_tags_match_any_with_table_alias(self):
"""Should include table alias with any mode."""
result = build_tags_where_clause_simple(["user_a"], 3, table_alias="mu.", match="any")
assert "mu.tags" in result
def test_tags_match_all_strict_with_table_alias(self):
"""Should include table alias with all_strict mode."""
result = build_tags_where_clause_simple(["user_a", "user_b"], 3, table_alias="mu.", match="all_strict")
assert "mu.tags" in result
assert "@>" in result
assert "IS NOT NULL" in result
# ============================================================================
# Unit Tests for filter_results_by_tags (Python-side filtering)
# ============================================================================
class MockResult:
"""Mock result object for testing filter_results_by_tags."""
def __init__(self, tags):
self.tags = tags
class TestFilterResultsByTags:
"""Unit tests for the Python-side tags filter function."""
def test_no_tags_returns_all(self):
"""When tags is None, should return all results."""
results = [MockResult(["a"]), MockResult(["b"]), MockResult(None)]
filtered = filter_results_by_tags(results, None)
assert len(filtered) == 3
def test_empty_tags_returns_all(self):
"""When tags is empty list, should return all results."""
results = [MockResult(["a"]), MockResult(["b"]), MockResult(None)]
filtered = filter_results_by_tags(results, [])
assert len(filtered) == 3
# ---- Test "any" mode (OR, includes untagged) ----
def test_any_mode_includes_matching_tags(self):
"""'any' mode should include results with matching tags."""
results = [MockResult(["a"]), MockResult(["b"]), MockResult(["c"])]
filtered = filter_results_by_tags(results, ["a", "b"], match="any")
# "a" and "b" match, "c" doesn't match and isn't untagged, so excluded
assert len(filtered) == 2
tags_found = [r.tags[0] for r in filtered if r.tags]
assert "a" in tags_found
assert "b" in tags_found
assert "c" not in tags_found
def test_any_mode_includes_untagged(self):
"""'any' mode should include untagged results."""
results = [MockResult(["a"]), MockResult(None), MockResult([])]
filtered = filter_results_by_tags(results, ["a"], match="any")
assert len(filtered) == 3 # a matches, None is untagged, [] is untagged
def test_any_mode_includes_partial_overlap(self):
"""'any' mode should include results with ANY overlapping tag."""
results = [MockResult(["a", "x"]), MockResult(["b", "y"])]
filtered = filter_results_by_tags(results, ["a"], match="any")
# ["a", "x"] matches, ["b", "y"] doesn't, but untagged would be included
tags_found = [r.tags for r in filtered]
assert ["a", "x"] in tags_found
# ---- Test "any_strict" mode (OR, excludes untagged) ----
def test_any_strict_excludes_untagged(self):
"""'any_strict' mode should exclude untagged results."""
results = [MockResult(["a"]), MockResult(None), MockResult([])]
filtered = filter_results_by_tags(results, ["a"], match="any_strict")
assert len(filtered) == 1 # Only ["a"] matches
assert filtered[0].tags == ["a"]
def test_any_strict_excludes_non_matching(self):
"""'any_strict' mode should exclude non-matching tagged results."""
results = [MockResult(["a"]), MockResult(["b"]), MockResult(["c"])]
filtered = filter_results_by_tags(results, ["a"], match="any_strict")
assert len(filtered) == 1
assert filtered[0].tags == ["a"]
# ---- Test "all" mode (AND, includes untagged) ----
def test_all_mode_requires_all_tags(self):
"""'all' mode should require ALL requested tags to be present."""
results = [MockResult(["a", "b"]), MockResult(["a"]), MockResult(["b"])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all")
# Only ["a", "b"] has both tags, but untagged would also be included
tags_found = [r.tags for r in filtered]
assert ["a", "b"] in tags_found
def test_all_mode_includes_untagged(self):
"""'all' mode should include untagged results."""
results = [MockResult(["a", "b"]), MockResult(None), MockResult([])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all")
assert len(filtered) == 3 # ["a", "b"] matches, None is untagged, [] is untagged
# ---- Test "all_strict" mode (AND, excludes untagged) ----
def test_all_strict_requires_all_tags(self):
"""'all_strict' mode should require ALL requested tags."""
results = [MockResult(["a", "b"]), MockResult(["a"]), MockResult(["b"])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
assert len(filtered) == 1
assert filtered[0].tags == ["a", "b"]
def test_all_strict_excludes_untagged(self):
"""'all_strict' mode should exclude untagged results."""
results = [MockResult(["a", "b"]), MockResult(None), MockResult([])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
assert len(filtered) == 1
assert filtered[0].tags == ["a", "b"]
def test_all_strict_allows_superset(self):
"""'all_strict' mode should allow results with MORE tags than requested."""
results = [MockResult(["a", "b", "c"]), MockResult(["a"])]
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
assert len(filtered) == 1
assert filtered[0].tags == ["a", "b", "c"] # Has a, b, AND c
# ============================================================================
# Integration Tests for tags in retain/recall/reflect
# ============================================================================
@pytest_asyncio.fixture
async def api_client(memory):
"""Create an async test client for the FastAPI app."""
app = create_app(memory, initialize_memory=False)
transport = httpx.ASGITransport(app=app)
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
yield client
@pytest.fixture
def test_bank_id():
"""Provide a unique bank ID for this test run."""
return f"tags_test_{datetime.now().timestamp()}"
@pytest.mark.asyncio
async def test_retain_with_tags(api_client, test_bank_id):
"""Test that memories can be stored with tags."""
# Store memory with tags
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{
"content": "Alice loves hiking in the mountains.",
"tags": ["user_alice"]
}
]
}
)
assert response.status_code == 200
result = response.json()
assert result["success"] is True
assert result["items_count"] == 1
@pytest.mark.asyncio
async def test_retain_with_document_tags(api_client, test_bank_id):
"""Test that document-level tags are applied to all items."""
# Store memories with document-level tags
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"document_tags": ["session_123"],
"items": [
{"content": "Bob discussed the quarterly report."},
{"content": "Charlie mentioned the new product launch."}
]
}
)
assert response.status_code == 200
result = response.json()
assert result["success"] is True
assert result["items_count"] == 2
@pytest.mark.asyncio
async def test_retain_merges_document_and_item_tags(api_client, test_bank_id):
"""Test that document tags and item tags are merged."""
# Store memory with both document and item tags
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"document_tags": ["session_abc"],
"items": [
{
"content": "Dave talked about machine learning.",
"tags": ["user_dave"]
}
]
}
)
assert response.status_code == 200
result = response.json()
assert result["success"] is True
@pytest.mark.asyncio
async def test_recall_without_tags_returns_all_memories(api_client, test_bank_id):
"""Test that recall without tags returns all memories (no filtering)."""
# Store memories for different users
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Eve works on natural language processing.", "tags": ["user_eve"]},
{"content": "Frank specializes in computer vision.", "tags": ["user_frank"]},
]
}
)
assert response.status_code == 200
# Recall without tags - should return all
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "Who works on what?", "budget": "low"}
)
assert response.status_code == 200
results = response.json()["results"]
# Should find both Eve and Frank
texts = [r["text"] for r in results]
assert any("Eve" in t for t in texts), "Should find Eve"
assert any("Frank" in t for t in texts), "Should find Frank"
@pytest.mark.asyncio
async def test_recall_with_tags_filters_memories(api_client, test_bank_id):
"""Test that recall with tags only returns matching memories."""
# Store memories for different users
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Grace is a data scientist at Google.", "tags": ["user_grace"]},
{"content": "Henry is a software engineer at Meta.", "tags": ["user_henry"]},
]
}
)
assert response.status_code == 200
# Recall with user_grace tag - should only return Grace's memory
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "Who works at which company?", "budget": "low", "tags": ["user_grace"]}
)
assert response.status_code == 200
results = response.json()["results"]
# Should find Grace but not Henry
texts = [r["text"] for r in results]
assert any("Grace" in t for t in texts), "Should find Grace with user_grace tag"
# Henry should NOT be found since he has user_henry tag
assert not any("Henry" in t for t in texts), "Should NOT find Henry (different tag)"
@pytest.mark.asyncio
async def test_recall_with_multiple_tags_uses_or_matching(api_client, test_bank_id):
"""Test that multiple tags use OR matching (any match returns the memory)."""
# Store memories for different users
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Ivan leads the security team.", "tags": ["user_ivan"]},
{"content": "Julia manages the design team.", "tags": ["user_julia"]},
{"content": "Karl oversees the marketing team.", "tags": ["user_karl"]},
]
}
)
assert response.status_code == 200
# Recall with user_ivan OR user_julia - should return both Ivan and Julia, but not Karl
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "Who leads which team?", "budget": "low", "tags": ["user_ivan", "user_julia"]}
)
assert response.status_code == 200
results = response.json()["results"]
texts = [r["text"] for r in results]
assert any("Ivan" in t for t in texts), "Should find Ivan (tag matches)"
assert any("Julia" in t for t in texts), "Should find Julia (tag matches)"
assert not any("Karl" in t for t in texts), "Should NOT find Karl (tag doesn't match)"
@pytest.mark.asyncio
async def test_recall_returns_memories_with_any_overlapping_tag(api_client, test_bank_id):
"""Test that memories with multiple tags are returned if ANY tag matches."""
# Store memory with multiple tags
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{
"content": "Lisa and Mike discussed the budget in a group chat.",
"tags": ["user_lisa", "user_mike"] # Memory visible to both
},
{"content": "Nancy reviewed the budget alone.", "tags": ["user_nancy"]},
]
}
)
assert response.status_code == 200
# Recall with user_lisa - should return the group chat memory
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "What was discussed about the budget?", "budget": "low", "tags": ["user_lisa"]}
)
assert response.status_code == 200
results = response.json()["results"]
texts = [r["text"] for r in results]
assert any("Lisa" in t and "Mike" in t for t in texts), "Should find group chat (Lisa is in tags)"
assert not any("Nancy" in t for t in texts), "Should NOT find Nancy's memory"
@pytest.mark.asyncio
async def test_reflect_with_tags_filters_memories(api_client, test_bank_id):
"""Test that reflect with tags only uses matching memories for reasoning."""
# Store different memories for different users
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Oscar's favorite color is blue.", "tags": ["user_oscar"]},
{"content": "Peter's favorite color is red.", "tags": ["user_peter"]},
]
}
)
assert response.status_code == 200
# Reflect with user_oscar tag - should only use Oscar's memories
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/reflect",
json={
"query": "What is the favorite color?",
"budget": "low",
"tags": ["user_oscar"],
"include": {"facts": {}} # Request facts to verify what was used
}
)
assert response.status_code == 200
result = response.json()
# The response should mention Oscar's color (blue), not Peter's (red)
# Note: We can check based_on facts if they're returned
if result.get("based_on"):
fact_texts = [f["text"] for f in result["based_on"]]
# Should use Oscar's memory
assert any("Oscar" in t or "blue" in t for t in fact_texts), "Should use Oscar's memory"
@pytest.mark.asyncio
async def test_recall_with_empty_tags_returns_all(api_client, test_bank_id):
"""Test that empty tags list behaves same as no tags (returns all)."""
# Store memories
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories",
json={
"items": [
{"content": "Quinn studies mathematics.", "tags": ["user_quinn"]},
{"content": "Rachel studies physics.", "tags": ["user_rachel"]},
]
}
)
assert response.status_code == 200
# Recall with empty tags list - should return all
response = await api_client.post(
f"/v1/default/banks/{test_bank_id}/memories/recall",
json={"query": "Who studies what?", "budget": "low", "tags": []}
)
assert response.status_code == 200
results = response.json()["results"]
texts = [r["text"] for r in results]
assert any("Quinn" in t for t in texts), "Should find Quinn"
assert any("Rachel" in t for t in texts), "Should find Rachel"
@pytest.mark.asyncio
async def test_multi_user_agent_visibility(api_client):
"""
Test multi-user agent visibility scoping.
Scenario:
- Agent has one memory bank
- Agent chats with User A (room 1) and User B (room 2) separately
- Agent also hosts a group chat with both users (room 3)
- User A should only see memories from rooms 1 and 3
- User B should only see memories from rooms 2 and 3
- Agent (no filter) should see all memories
"""
bank_id = f"multi_user_test_{datetime.now().timestamp()}"
# Store memories from different chat rooms
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
# Room 1: Agent + User A private chat
{"content": "User A said they prefer morning meetings.", "tags": ["user_a"]},
# Room 2: Agent + User B private chat
{"content": "User B mentioned they like afternoon meetings.", "tags": ["user_b"]},
# Room 3: Group chat with both users
{"content": "In the group meeting, they agreed to meet at noon.", "tags": ["user_a", "user_b"]},
]
}
)
assert response.status_code == 200
# User A queries - should see their private chat and group chat
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "What meeting time preferences were discussed?", "budget": "low", "tags": ["user_a"]}
)
assert response.status_code == 200
user_a_results = response.json()["results"]
user_a_texts = [r["text"] for r in user_a_results]
assert any("morning" in t for t in user_a_texts), "User A should see their own preference (morning)"
assert any("noon" in t for t in user_a_texts), "User A should see group chat (noon)"
assert not any("afternoon" in t for t in user_a_texts), "User A should NOT see User B's private preference"
# User B queries - should see their private chat and group chat
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "What meeting time preferences were discussed?", "budget": "low", "tags": ["user_b"]}
)
assert response.status_code == 200
user_b_results = response.json()["results"]
user_b_texts = [r["text"] for r in user_b_results]
assert any("afternoon" in t for t in user_b_texts), "User B should see their own preference (afternoon)"
assert any("noon" in t for t in user_b_texts), "User B should see group chat (noon)"
assert not any("morning" in t for t in user_b_texts), "User B should NOT see User A's private preference"
# Agent queries (no filter) - should see everything
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "What meeting time preferences were discussed?", "budget": "low"} # No tags
)
assert response.status_code == 200
agent_results = response.json()["results"]
agent_texts = [r["text"] for r in agent_results]
assert any("morning" in t for t in agent_texts), "Agent should see User A's preference"
assert any("afternoon" in t for t in agent_texts), "Agent should see User B's preference"
assert any("noon" in t for t in agent_texts), "Agent should see group chat"
@pytest.mark.asyncio
async def test_student_tracking_visibility(api_client):
"""
Test student tracking visibility scoping.
Scenario:
- Teacher bot has one memory bank
- Teacher records observations for Student A, Student B
- Student A should only see their own data
- Teacher (no filter) should see all student data
"""
bank_id = f"student_test_{datetime.now().timestamp()}"
# Store memories for different students
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Student A showed improvement in algebra today.", "tags": ["student_a"]},
{"content": "Student B struggled with geometry concepts.", "tags": ["student_b"]},
{"content": "Student A participated actively in class discussion.", "tags": ["student_a"]},
]
}
)
assert response.status_code == 200
# Student A queries - should only see their own data
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "How am I doing in class?", "budget": "low", "tags": ["student_a"]}
)
assert response.status_code == 200
student_a_results = response.json()["results"]
student_a_texts = [r["text"] for r in student_a_results]
assert any("algebra" in t for t in student_a_texts), "Student A should see their algebra progress"
assert any("participated" in t for t in student_a_texts), "Student A should see their participation"
assert not any("Student B" in t or "geometry" in t for t in student_a_texts), "Student A should NOT see Student B's data"
# Teacher queries (no filter) - should see all students
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories/recall",
json={"query": "Which students need help?", "budget": "low"} # No tags
)
assert response.status_code == 200
teacher_results = response.json()["results"]
teacher_texts = [r["text"] for r in teacher_results]
assert any("Student A" in t for t in teacher_texts), "Teacher should see Student A's data"
assert any("Student B" in t for t in teacher_texts), "Teacher should see Student B's data"
# ============================================================================
# Tests for list_tags API endpoint
# ============================================================================
@pytest.mark.asyncio
async def test_list_tags_returns_all_tags(api_client):
"""Test that list_tags returns all unique tags with counts."""
bank_id = f"list_tags_test_{datetime.now().timestamp()}"
# Store memories with various tags
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Memory 1 for user alice.", "tags": ["user:alice"]},
{"content": "Memory 2 for user alice.", "tags": ["user:alice"]},
{"content": "Memory 3 for user bob.", "tags": ["user:bob"]},
{"content": "Memory 4 in session 123.", "tags": ["session:123"]},
{"content": "Memory 5 for alice in session 456.", "tags": ["user:alice", "session:456"]},
]
}
)
assert response.status_code == 200
# List all tags
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
assert response.status_code == 200
result = response.json()
# Verify structure
assert "items" in result
assert "total" in result
assert "limit" in result
assert "offset" in result
# Verify tags and counts
tags_map = {item["tag"]: item["count"] for item in result["items"]}
assert "user:alice" in tags_map
assert tags_map["user:alice"] == 3 # 3 memories have this tag
assert "user:bob" in tags_map
assert tags_map["user:bob"] == 1
assert "session:123" in tags_map
assert tags_map["session:123"] == 1
assert "session:456" in tags_map
assert tags_map["session:456"] == 1
assert result["total"] == 4 # 4 unique tags
@pytest.mark.asyncio
async def test_list_tags_with_wildcard_prefix(api_client):
"""Test that list_tags filters with prefix wildcard pattern (user:*)."""
bank_id = f"list_tags_wildcard_test_{datetime.now().timestamp()}"
# Store memories with various tags
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Memory for alice who works at tech.", "tags": ["user:alice"]},
{"content": "Memory for bob who is an engineer.", "tags": ["user:bob"]},
{"content": "Memory for charlie the designer.", "tags": ["user:charlie"]},
{"content": "Session memory about the meeting.", "tags": ["session:abc"]},
{"content": "Room memory for conference room.", "tags": ["room:123"]},
]
}
)
assert response.status_code == 200
# List tags with 'user:*' wildcard pattern
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "user:*"})
assert response.status_code == 200
result = response.json()
# Should only return user:* tags
tags = [item["tag"] for item in result["items"]]
assert "user:alice" in tags
assert "user:bob" in tags
assert "user:charlie" in tags
assert "session:abc" not in tags
assert "room:123" not in tags
assert result["total"] == 3
@pytest.mark.asyncio
async def test_list_tags_with_wildcard_suffix(api_client):
"""Test that list_tags filters with suffix wildcard pattern (*-admin)."""
bank_id = f"list_tags_suffix_test_{datetime.now().timestamp()}"
# Store memories with various tags
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Admin role memory for super admin.", "tags": ["role-admin"]},
{"content": "Super admin memory about permissions.", "tags": ["super-admin"]},
{"content": "User memory for standard users.", "tags": ["role-user"]},
{"content": "Guest memory for visitors.", "tags": ["role-guest"]},
]
}
)
assert response.status_code == 200
# List tags with '*-admin' wildcard pattern (suffix match)
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "*-admin"})
assert response.status_code == 200
result = response.json()
# Should only return *-admin tags
tags = [item["tag"] for item in result["items"]]
assert "role-admin" in tags
assert "super-admin" in tags
assert "role-user" not in tags
assert "role-guest" not in tags
assert result["total"] == 2
@pytest.mark.asyncio
async def test_list_tags_with_wildcard_middle(api_client):
"""Test that list_tags filters with middle wildcard pattern (env*-prod)."""
bank_id = f"list_tags_middle_test_{datetime.now().timestamp()}"
# Store memories with various tags - use meaningful content for fact extraction
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "The production environment is configured with high availability and uses AWS infrastructure.", "tags": ["env-prod"]},
{"content": "The enterprise environment for production runs on dedicated servers with 24/7 monitoring.", "tags": ["environment-prod"]},
{"content": "The staging environment mirrors production but uses smaller instance sizes.", "tags": ["env-staging"]},
{"content": "The development environment allows developers to test their code locally.", "tags": ["env-dev"]},
]
}
)
assert response.status_code == 200
# List tags with 'env*-prod' wildcard pattern (middle match)
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "env*-prod"})
assert response.status_code == 200
result = response.json()
# Should only return env*-prod tags
tags = [item["tag"] for item in result["items"]]
assert "env-prod" in tags
assert "environment-prod" in tags
assert "env-staging" not in tags
assert "env-dev" not in tags
assert result["total"] == 2
@pytest.mark.asyncio
async def test_list_tags_case_insensitive(api_client):
"""Test that list_tags wildcard matching is case-insensitive."""
bank_id = f"list_tags_case_test_{datetime.now().timestamp()}"
# Store memories with mixed case tags - use meaningful content
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Alice is a software engineer who specializes in machine learning algorithms.", "tags": ["User:Alice"]},
{"content": "Bob works as a data scientist at a large technology company.", "tags": ["user:bob"]},
{"content": "Charlie is the lead designer responsible for the user interface.", "tags": ["USER:CHARLIE"]},
]
}
)
assert response.status_code == 200
# List tags with lowercase pattern - should match all cases
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "user:*"})
assert response.status_code == 200
result = response.json()
# Should match all user tags regardless of case
tags = [item["tag"] for item in result["items"]]
assert len(tags) == 3
assert result["total"] == 3
@pytest.mark.asyncio
async def test_list_tags_pagination(api_client):
"""Test that list_tags supports pagination."""
bank_id = f"list_tags_pagination_test_{datetime.now().timestamp()}"
# Store memories with many tags - use meaningful content for fact extraction
names = ["Alice", "Bob", "Charlie", "Diana", "Eve", "Frank", "Grace", "Henry", "Ivan", "Julia"]
items = [
{"content": f"{name} works as a software engineer at company {i}.", "tags": [f"tag:{i:03d}"]}
for i, name in enumerate(names)
]
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={"items": items}
)
assert response.status_code == 200
# Get first page (limit 3)
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"limit": 3, "offset": 0})
assert response.status_code == 200
result = response.json()
assert len(result["items"]) == 3
assert result["total"] == 10
assert result["limit"] == 3
assert result["offset"] == 0
# Get second page
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"limit": 3, "offset": 3})
assert response.status_code == 200
result = response.json()
assert len(result["items"]) == 3
assert result["offset"] == 3
@pytest.mark.asyncio
async def test_list_tags_empty_bank(api_client):
"""Test that list_tags returns empty for bank with no tags."""
bank_id = f"list_tags_empty_test_{datetime.now().timestamp()}"
# List tags without storing anything
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
assert response.status_code == 200
result = response.json()
assert result["items"] == []
assert result["total"] == 0
@pytest.mark.asyncio
async def test_list_tags_ordered_by_count(api_client):
"""Test that list_tags returns tags ordered by frequency (most used first)."""
bank_id = f"list_tags_order_test_{datetime.now().timestamp()}"
# Store memories with tags having different frequencies - use meaningful content
response = await api_client.post(
f"/v1/default/banks/{bank_id}/memories",
json={
"items": [
{"content": "Alice works at a startup company as a developer.", "tags": ["rare"]},
{"content": "Bob is a senior engineer at Google.", "tags": ["common"]},
{"content": "Charlie manages the marketing team at Microsoft.", "tags": ["common"]},
{"content": "Diana leads the design department at Apple.", "tags": ["common"]},
{"content": "Eve is a data scientist at Amazon.", "tags": ["medium"]},
{"content": "Frank handles customer support at Meta.", "tags": ["medium"]},
]
}
)
assert response.status_code == 200
# List tags - should be ordered by count descending
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
assert response.status_code == 200
result = response.json()
tags = [item["tag"] for item in result["items"]]
# common (3) should come before medium (2) which should come before rare (1)
assert tags.index("common") < tags.index("medium")
assert tags.index("medium") < tags.index("rare")
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "hindsight-cli"
version = "0.2.1"
version = "0.3.0"
edition = "2021"
authors = ["Hindsight Team"]
description = "A beautiful CLI for Hindsight - semantic memory system"
+5 -1
View File
@@ -5,7 +5,7 @@ use crossterm::{
execute,
terminal::{disable_raw_mode, enable_raw_mode, EnterAlternateScreen, LeaveAlternateScreen},
};
use hindsight_client::types::{BankListItem, RecallResult, EntityListItem, Budget};
use hindsight_client::types::{BankListItem, RecallResult, EntityListItem, Budget, TagsMatch};
use serde_json::{Map, Value};
use ratatui::{
backend::{Backend, CrosstermBackend},
@@ -341,6 +341,8 @@ impl App {
trace: false,
query_timestamp: None,
include: None,
tags: None,
tags_match: TagsMatch::Any,
};
let result = client.recall(&bank_id, &request, false)
@@ -357,6 +359,8 @@ impl App {
max_tokens: 4096,
include: None,
response_schema: None,
tags: None,
tags_match: TagsMatch::Any,
};
let result = client.reflect(&bank_id, &request, false)
+9 -1
View File
@@ -9,7 +9,7 @@ use crate::output::{self, OutputFormat};
use crate::ui;
// Import types from generated client
use hindsight_client::types::{Budget, ChunkIncludeOptions, IncludeOptions};
use hindsight_client::types::{Budget, ChunkIncludeOptions, IncludeOptions, TagsMatch};
use serde_json;
// Helper function to parse budget string to Budget enum
@@ -60,6 +60,8 @@ pub fn recall(
trace,
query_timestamp: None,
include,
tags: None,
tags_match: TagsMatch::Any,
};
let response = client.recall(agent_id, &request, verbose);
@@ -116,6 +118,8 @@ pub fn reflect(
max_tokens: max_tokens.unwrap_or(4096),
include: None,
response_schema,
tags: None,
tags_match: TagsMatch::Any,
};
let response = client.reflect(agent_id, &request, verbose);
@@ -162,11 +166,13 @@ pub fn retain(
timestamp: None,
document_id: Some(doc_id.clone()),
entities: None,
tags: None,
};
let request = RetainRequest {
items: vec![item],
async_: r#async,
document_tags: None,
};
let response = client.retain(agent_id, &request, r#async, verbose);
@@ -272,6 +278,7 @@ pub fn retain_files(
timestamp: None,
document_id: Some(doc_id),
entities: None,
tags: None,
});
pb.inc(1);
@@ -288,6 +295,7 @@ pub fn retain_files(
let request = RetainRequest {
items,
async_: r#async,
document_tags: None,
};
let response = client.retain(agent_id, &request, r#async, verbose);
@@ -39,6 +39,7 @@ hindsight_client_api/models/http_validation_error.py
hindsight_client_api/models/include_options.py
hindsight_client_api/models/list_documents_response.py
hindsight_client_api/models/list_memory_units_response.py
hindsight_client_api/models/list_tags_response.py
hindsight_client_api/models/memory_item.py
hindsight_client_api/models/operation_response.py
hindsight_client_api/models/operations_list_response.py
@@ -51,6 +52,7 @@ hindsight_client_api/models/reflect_request.py
hindsight_client_api/models/reflect_response.py
hindsight_client_api/models/retain_request.py
hindsight_client_api/models/retain_response.py
hindsight_client_api/models/tag_item.py
hindsight_client_api/models/token_usage.py
hindsight_client_api/models/update_disposition_request.py
hindsight_client_api/models/validation_error.py
@@ -115,6 +115,7 @@ class Hindsight:
document_id: Optional[str] = None,
metadata: Optional[Dict[str, str]] = None,
entities: Optional[List[Dict[str, str]]] = None,
tags: Optional[List[str]] = None,
) -> RetainResponse:
"""
Store a single memory (simplified interface).
@@ -127,13 +128,14 @@ class Hindsight:
document_id: Optional document ID for grouping
metadata: Optional user-defined metadata
entities: Optional list of entities [{"text": "...", "type": "..."}]
tags: Optional list of tags for this memory
Returns:
RetainResponse with success status
"""
return self.retain_batch(
bank_id=bank_id,
items=[{"content": content, "timestamp": timestamp, "context": context, "metadata": metadata, "entities": entities}],
items=[{"content": content, "timestamp": timestamp, "context": context, "metadata": metadata, "entities": entities, "tags": tags}],
document_id=document_id,
)
@@ -143,15 +145,17 @@ class Hindsight:
items: List[Dict[str, Any]],
document_id: Optional[str] = None,
retain_async: bool = False,
document_tags: Optional[List[str]] = None,
) -> RetainResponse:
"""
Store multiple memories in batch.
Args:
bank_id: The memory bank ID
items: List of memory items with 'content' and optional 'timestamp', 'context', 'metadata', 'document_id', 'entities'
items: List of memory items with 'content' and optional 'timestamp', 'context', 'metadata', 'document_id', 'entities', 'tags'
document_id: Optional document ID for grouping memories (applied to items that don't have their own)
retain_async: If True, process asynchronously in background (default: False)
document_tags: Optional list of tags to apply to all memories in this batch
Returns:
RetainResponse with success status and item count
@@ -175,12 +179,14 @@ class Hindsight:
# Use item's document_id if provided, otherwise fall back to batch-level document_id
document_id=item.get("document_id") or document_id,
entities=entities,
tags=item.get("tags"),
)
)
request_obj = retain_request.RetainRequest(
items=memory_items,
async_=retain_async,
document_tags=document_tags,
)
return _run_async(self._memory_api.retain_memories(bank_id, request_obj))
@@ -198,6 +204,8 @@ class Hindsight:
max_entity_tokens: int = 500,
include_chunks: bool = False,
max_chunk_tokens: int = 8192,
tags: Optional[List[str]] = None,
tags_match: str = "any",
) -> RecallResponse:
"""
Recall memories using semantic similarity.
@@ -214,6 +222,9 @@ class Hindsight:
max_entity_tokens: Maximum tokens for entity observations (default: 500)
include_chunks: Include raw text chunks in results (default: False)
max_chunk_tokens: Maximum tokens for chunks (default: 8192)
tags: Optional list of tags to filter memories by
tags_match: How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged),
'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged). Default: 'any'
Returns:
RecallResponse with results, optional entities, optional chunks, and optional trace
@@ -233,6 +244,8 @@ class Hindsight:
trace=trace,
query_timestamp=query_timestamp,
include=include_opts,
tags=tags,
tags_match=tags_match,
)
return _run_async(self._memory_api.recall_memories(bank_id, request_obj))
@@ -245,6 +258,8 @@ class Hindsight:
context: Optional[str] = None,
max_tokens: Optional[int] = None,
response_schema: Optional[Dict[str, Any]] = None,
tags: Optional[List[str]] = None,
tags_match: str = "any",
) -> ReflectResponse:
"""
Generate a contextual answer based on bank identity and memories.
@@ -258,6 +273,9 @@ class Hindsight:
response_schema: 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: Optional list of tags to filter memories by
tags_match: How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged),
'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged). Default: 'any'
Returns:
ReflectResponse with answer text, optionally facts used, and optionally
@@ -269,6 +287,8 @@ class Hindsight:
context=context,
max_tokens=max_tokens,
response_schema=response_schema,
tags=tags,
tags_match=tags_match,
)
return _run_async(self._memory_api.reflect(bank_id, request_obj))
@@ -64,6 +64,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
from hindsight_client_api.models.include_options import IncludeOptions
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
from hindsight_client_api.models.list_tags_response import ListTagsResponse
from hindsight_client_api.models.memory_item import MemoryItem
from hindsight_client_api.models.operation_response import OperationResponse
from hindsight_client_api.models.operations_list_response import OperationsListResponse
@@ -76,6 +77,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
from hindsight_client_api.models.reflect_response import ReflectResponse
from hindsight_client_api.models.retain_request import RetainRequest
from hindsight_client_api.models.retain_response import RetainResponse
from hindsight_client_api.models.tag_item import TagItem
from hindsight_client_api.models.token_usage import TokenUsage
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
from hindsight_client_api.models.validation_error import ValidationError
@@ -17,11 +17,12 @@ from typing import Any, Dict, List, Optional, Tuple, Union
from typing_extensions import Annotated
from pydantic import Field, StrictInt, StrictStr
from typing import Optional
from typing import Any, Optional
from typing_extensions import Annotated
from hindsight_client_api.models.delete_response import DeleteResponse
from hindsight_client_api.models.graph_data_response import GraphDataResponse
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
from hindsight_client_api.models.list_tags_response import ListTagsResponse
from hindsight_client_api.models.recall_request import RecallRequest
from hindsight_client_api.models.recall_response import RecallResponse
from hindsight_client_api.models.reflect_request import ReflectRequest
@@ -654,6 +655,299 @@ class MemoryApi:
@validate_call
async def get_memory(
self,
bank_id: StrictStr,
memory_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> object:
"""Get memory unit
Get a single memory unit by ID with all its metadata including entities and tags.
:param bank_id: (required)
:type bank_id: str
:param memory_id: (required)
:type memory_id: str
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._get_memory_serialize(
bank_id=bank_id,
memory_id=memory_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "object",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
await response_data.read()
return self.api_client.response_deserialize(
response_data=response_data,
response_types_map=_response_types_map,
).data
@validate_call
async def get_memory_with_http_info(
self,
bank_id: StrictStr,
memory_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> ApiResponse[object]:
"""Get memory unit
Get a single memory unit by ID with all its metadata including entities and tags.
:param bank_id: (required)
:type bank_id: str
:param memory_id: (required)
:type memory_id: str
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._get_memory_serialize(
bank_id=bank_id,
memory_id=memory_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "object",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
await response_data.read()
return self.api_client.response_deserialize(
response_data=response_data,
response_types_map=_response_types_map,
)
@validate_call
async def get_memory_without_preload_content(
self,
bank_id: StrictStr,
memory_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> RESTResponseType:
"""Get memory unit
Get a single memory unit by ID with all its metadata including entities and tags.
:param bank_id: (required)
:type bank_id: str
:param memory_id: (required)
:type memory_id: str
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._get_memory_serialize(
bank_id=bank_id,
memory_id=memory_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "object",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
return response_data.response
def _get_memory_serialize(
self,
bank_id,
memory_id,
authorization,
_request_auth,
_content_type,
_headers,
_host_index,
) -> RequestSerialized:
_host = None
_collection_formats: Dict[str, str] = {
}
_path_params: Dict[str, str] = {}
_query_params: List[Tuple[str, str]] = []
_header_params: Dict[str, Optional[str]] = _headers or {}
_form_params: List[Tuple[str, str]] = []
_files: Dict[
str, Union[str, bytes, List[str], List[bytes], List[Tuple[str, bytes]]]
] = {}
_body_params: Optional[bytes] = None
# process the path parameters
if bank_id is not None:
_path_params['bank_id'] = bank_id
if memory_id is not None:
_path_params['memory_id'] = memory_id
# process the query parameters
# process the header parameters
if authorization is not None:
_header_params['authorization'] = authorization
# process the form parameters
# process the body parameter
# set the HTTP header `Accept`
if 'Accept' not in _header_params:
_header_params['Accept'] = self.api_client.select_header_accept(
[
'application/json'
]
)
# authentication setting
_auth_settings: List[str] = [
]
return self.api_client.param_serialize(
method='GET',
resource_path='/v1/default/banks/{bank_id}/memories/{memory_id}',
path_params=_path_params,
query_params=_query_params,
header_params=_header_params,
body=_body_params,
post_params=_form_params,
files=_files,
auth_settings=_auth_settings,
collection_formats=_collection_formats,
_host=_host,
_request_auth=_request_auth
)
@validate_call
async def list_memories(
self,
@@ -1000,6 +1294,335 @@ class MemoryApi:
@validate_call
async def list_tags(
self,
bank_id: StrictStr,
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> ListTagsResponse:
"""List tags
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
:param bank_id: (required)
:type bank_id: str
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
:type q: str
:param limit: Maximum number of tags to return
:type limit: int
:param offset: Offset for pagination
:type offset: int
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._list_tags_serialize(
bank_id=bank_id,
q=q,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "ListTagsResponse",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
await response_data.read()
return self.api_client.response_deserialize(
response_data=response_data,
response_types_map=_response_types_map,
).data
@validate_call
async def list_tags_with_http_info(
self,
bank_id: StrictStr,
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> ApiResponse[ListTagsResponse]:
"""List tags
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
:param bank_id: (required)
:type bank_id: str
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
:type q: str
:param limit: Maximum number of tags to return
:type limit: int
:param offset: Offset for pagination
:type offset: int
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._list_tags_serialize(
bank_id=bank_id,
q=q,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "ListTagsResponse",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
await response_data.read()
return self.api_client.response_deserialize(
response_data=response_data,
response_types_map=_response_types_map,
)
@validate_call
async def list_tags_without_preload_content(
self,
bank_id: StrictStr,
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
Tuple[
Annotated[StrictFloat, Field(gt=0)],
Annotated[StrictFloat, Field(gt=0)]
]
] = None,
_request_auth: Optional[Dict[StrictStr, Any]] = None,
_content_type: Optional[StrictStr] = None,
_headers: Optional[Dict[StrictStr, Any]] = None,
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
) -> RESTResponseType:
"""List tags
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
:param bank_id: (required)
:type bank_id: str
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
:type q: str
:param limit: Maximum number of tags to return
:type limit: int
:param offset: Offset for pagination
:type offset: int
:param authorization:
:type authorization: str
:param _request_timeout: timeout setting for this request. If one
number provided, it will be total request
timeout. It can also be a pair (tuple) of
(connection, read) timeouts.
:type _request_timeout: int, tuple(int, int), optional
:param _request_auth: set to override the auth_settings for an a single
request; this effectively ignores the
authentication in the spec for a single request.
:type _request_auth: dict, optional
:param _content_type: force content-type for the request.
:type _content_type: str, Optional
:param _headers: set to override the headers for a single
request; this effectively ignores the headers
in the spec for a single request.
:type _headers: dict, optional
:param _host_index: set to override the host_index for a single
request; this effectively ignores the host_index
in the spec for a single request.
:type _host_index: int, optional
:return: Returns the result object.
""" # noqa: E501
_param = self._list_tags_serialize(
bank_id=bank_id,
q=q,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
_host_index=_host_index
)
_response_types_map: Dict[str, Optional[str]] = {
'200': "ListTagsResponse",
'422': "HTTPValidationError",
}
response_data = await self.api_client.call_api(
*_param,
_request_timeout=_request_timeout
)
return response_data.response
def _list_tags_serialize(
self,
bank_id,
q,
limit,
offset,
authorization,
_request_auth,
_content_type,
_headers,
_host_index,
) -> RequestSerialized:
_host = None
_collection_formats: Dict[str, str] = {
}
_path_params: Dict[str, str] = {}
_query_params: List[Tuple[str, str]] = []
_header_params: Dict[str, Optional[str]] = _headers or {}
_form_params: List[Tuple[str, str]] = []
_files: Dict[
str, Union[str, bytes, List[str], List[bytes], List[Tuple[str, bytes]]]
] = {}
_body_params: Optional[bytes] = None
# process the path parameters
if bank_id is not None:
_path_params['bank_id'] = bank_id
# process the query parameters
if q is not None:
_query_params.append(('q', q))
if limit is not None:
_query_params.append(('limit', limit))
if offset is not None:
_query_params.append(('offset', offset))
# process the header parameters
if authorization is not None:
_header_params['authorization'] = authorization
# process the form parameters
# process the body parameter
# set the HTTP header `Accept`
if 'Accept' not in _header_params:
_header_params['Accept'] = self.api_client.select_header_accept(
[
'application/json'
]
)
# authentication setting
_auth_settings: List[str] = [
]
return self.api_client.param_serialize(
method='GET',
resource_path='/v1/default/banks/{bank_id}/tags',
path_params=_path_params,
query_params=_query_params,
header_params=_header_params,
body=_body_params,
post_params=_form_params,
files=_files,
auth_settings=_auth_settings,
collection_formats=_collection_formats,
_host=_host,
_request_auth=_request_auth
)
@validate_call
async def recall_memories(
self,
@@ -42,6 +42,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
from hindsight_client_api.models.include_options import IncludeOptions
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
from hindsight_client_api.models.list_tags_response import ListTagsResponse
from hindsight_client_api.models.memory_item import MemoryItem
from hindsight_client_api.models.operation_response import OperationResponse
from hindsight_client_api.models.operations_list_response import OperationsListResponse
@@ -54,6 +55,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
from hindsight_client_api.models.reflect_response import ReflectResponse
from hindsight_client_api.models.retain_request import RetainRequest
from hindsight_client_api.models.retain_response import RetainResponse
from hindsight_client_api.models.tag_item import TagItem
from hindsight_client_api.models.token_usage import TokenUsage
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
from hindsight_client_api.models.validation_error import ValidationError
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, StrictInt, StrictStr
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
from typing import Any, ClassVar, Dict, List, Optional
from typing import Optional, Set
from typing_extensions import Self
@@ -33,7 +33,8 @@ class DocumentResponse(BaseModel):
created_at: StrictStr
updated_at: StrictStr
memory_unit_count: StrictInt
__properties: ClassVar[List[str]] = ["id", "bank_id", "original_text", "content_hash", "created_at", "updated_at", "memory_unit_count"]
tags: Optional[List[StrictStr]] = Field(default=None, description="Tags associated with this document")
__properties: ClassVar[List[str]] = ["id", "bank_id", "original_text", "content_hash", "created_at", "updated_at", "memory_unit_count", "tags"]
model_config = ConfigDict(
populate_by_name=True,
@@ -97,7 +98,8 @@ class DocumentResponse(BaseModel):
"content_hash": obj.get("content_hash"),
"created_at": obj.get("created_at"),
"updated_at": obj.get("updated_at"),
"memory_unit_count": obj.get("memory_unit_count")
"memory_unit_count": obj.get("memory_unit_count"),
"tags": obj.get("tags")
})
return _obj
@@ -0,0 +1,101 @@
# coding: utf-8
"""
Hindsight HTTP API
HTTP API for Hindsight
The version of the OpenAPI document: 0.1.0
Generated by OpenAPI Generator (https://openapi-generator.tech)
Do not edit the class manually.
""" # noqa: E501
from __future__ import annotations
import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, StrictInt
from typing import Any, ClassVar, Dict, List
from hindsight_client_api.models.tag_item import TagItem
from typing import Optional, Set
from typing_extensions import Self
class ListTagsResponse(BaseModel):
"""
Response model for list tags endpoint.
""" # noqa: E501
items: List[TagItem]
total: StrictInt
limit: StrictInt
offset: StrictInt
__properties: ClassVar[List[str]] = ["items", "total", "limit", "offset"]
model_config = ConfigDict(
populate_by_name=True,
validate_assignment=True,
protected_namespaces=(),
)
def to_str(self) -> str:
"""Returns the string representation of the model using alias"""
return pprint.pformat(self.model_dump(by_alias=True))
def to_json(self) -> str:
"""Returns the JSON representation of the model using alias"""
# TODO: pydantic v2: use .model_dump_json(by_alias=True, exclude_unset=True) instead
return json.dumps(self.to_dict())
@classmethod
def from_json(cls, json_str: str) -> Optional[Self]:
"""Create an instance of ListTagsResponse from a JSON string"""
return cls.from_dict(json.loads(json_str))
def to_dict(self) -> Dict[str, Any]:
"""Return the dictionary representation of the model using alias.
This has the following differences from calling pydantic's
`self.model_dump(by_alias=True)`:
* `None` is only added to the output dict for nullable fields that
were set at model initialization. Other fields with value `None`
are ignored.
"""
excluded_fields: Set[str] = set([
])
_dict = self.model_dump(
by_alias=True,
exclude=excluded_fields,
exclude_none=True,
)
# override the default output from pydantic by calling `to_dict()` of each item in items (list)
_items = []
if self.items:
for _item_items in self.items:
if _item_items:
_items.append(_item_items.to_dict())
_dict['items'] = _items
return _dict
@classmethod
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional[Self]:
"""Create an instance of ListTagsResponse from a dict"""
if obj is None:
return None
if not isinstance(obj, dict):
return cls.model_validate(obj)
_obj = cls.model_validate({
"items": [TagItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None,
"total": obj.get("total"),
"limit": obj.get("limit"),
"offset": obj.get("offset")
})
return _obj
@@ -34,7 +34,8 @@ class MemoryItem(BaseModel):
metadata: Optional[Dict[str, StrictStr]] = None
document_id: Optional[StrictStr] = None
entities: Optional[List[EntityInput]] = None
__properties: ClassVar[List[str]] = ["content", "timestamp", "context", "metadata", "document_id", "entities"]
tags: Optional[List[StrictStr]] = None
__properties: ClassVar[List[str]] = ["content", "timestamp", "context", "metadata", "document_id", "entities", "tags"]
model_config = ConfigDict(
populate_by_name=True,
@@ -107,6 +108,11 @@ class MemoryItem(BaseModel):
if self.entities is None and "entities" in self.model_fields_set:
_dict['entities'] = None
# set to None if tags (nullable) is None
# and model_fields_set contains the field
if self.tags is None and "tags" in self.model_fields_set:
_dict['tags'] = None
return _dict
@classmethod
@@ -124,7 +130,8 @@ class MemoryItem(BaseModel):
"context": obj.get("context"),
"metadata": obj.get("metadata"),
"document_id": obj.get("document_id"),
"entities": [EntityInput.from_dict(_item) for _item in obj["entities"]] if obj.get("entities") is not None else None
"entities": [EntityInput.from_dict(_item) for _item in obj["entities"]] if obj.get("entities") is not None else None,
"tags": obj.get("tags")
})
return _obj
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictInt, StrictStr
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictInt, StrictStr, field_validator
from typing import Any, ClassVar, Dict, List, Optional
from hindsight_client_api.models.budget import Budget
from hindsight_client_api.models.include_options import IncludeOptions
@@ -35,7 +35,19 @@ class RecallRequest(BaseModel):
trace: Optional[StrictBool] = False
query_timestamp: Optional[StrictStr] = None
include: Optional[IncludeOptions] = Field(default=None, description="Options for including additional data (entities are included by default)")
__properties: ClassVar[List[str]] = ["query", "types", "budget", "max_tokens", "trace", "query_timestamp", "include"]
tags: Optional[List[StrictStr]] = None
tags_match: Optional[StrictStr] = Field(default='any', description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).")
__properties: ClassVar[List[str]] = ["query", "types", "budget", "max_tokens", "trace", "query_timestamp", "include", "tags", "tags_match"]
@field_validator('tags_match')
def tags_match_validate_enum(cls, value):
"""Validates the enum"""
if value is None:
return value
if value not in set(['any', 'all', 'any_strict', 'all_strict']):
raise ValueError("must be one of enum values ('any', 'all', 'any_strict', 'all_strict')")
return value
model_config = ConfigDict(
populate_by_name=True,
@@ -89,6 +101,11 @@ class RecallRequest(BaseModel):
if self.query_timestamp is None and "query_timestamp" in self.model_fields_set:
_dict['query_timestamp'] = None
# set to None if tags (nullable) is None
# and model_fields_set contains the field
if self.tags is None and "tags" in self.model_fields_set:
_dict['tags'] = None
return _dict
@classmethod
@@ -107,7 +124,9 @@ class RecallRequest(BaseModel):
"max_tokens": obj.get("max_tokens") if obj.get("max_tokens") is not None else 4096,
"trace": obj.get("trace") if obj.get("trace") is not None else False,
"query_timestamp": obj.get("query_timestamp"),
"include": IncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None
"include": IncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None,
"tags": obj.get("tags"),
"tags_match": obj.get("tags_match") if obj.get("tags_match") is not None else 'any'
})
return _obj
@@ -37,7 +37,8 @@ class RecallResult(BaseModel):
document_id: Optional[StrictStr] = None
metadata: Optional[Dict[str, StrictStr]] = None
chunk_id: Optional[StrictStr] = None
__properties: ClassVar[List[str]] = ["id", "text", "type", "entities", "context", "occurred_start", "occurred_end", "mentioned_at", "document_id", "metadata", "chunk_id"]
tags: Optional[List[StrictStr]] = None
__properties: ClassVar[List[str]] = ["id", "text", "type", "entities", "context", "occurred_start", "occurred_end", "mentioned_at", "document_id", "metadata", "chunk_id", "tags"]
model_config = ConfigDict(
populate_by_name=True,
@@ -123,6 +124,11 @@ class RecallResult(BaseModel):
if self.chunk_id is None and "chunk_id" in self.model_fields_set:
_dict['chunk_id'] = None
# set to None if tags (nullable) is None
# and model_fields_set contains the field
if self.tags is None and "tags" in self.model_fields_set:
_dict['tags'] = None
return _dict
@classmethod
@@ -145,7 +151,8 @@ class RecallResult(BaseModel):
"mentioned_at": obj.get("mentioned_at"),
"document_id": obj.get("document_id"),
"metadata": obj.get("metadata"),
"chunk_id": obj.get("chunk_id")
"chunk_id": obj.get("chunk_id"),
"tags": obj.get("tags")
})
return _obj
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr, field_validator
from typing import Any, ClassVar, Dict, List, Optional
from hindsight_client_api.models.budget import Budget
from hindsight_client_api.models.reflect_include_options import ReflectIncludeOptions
@@ -34,7 +34,19 @@ class ReflectRequest(BaseModel):
max_tokens: Optional[StrictInt] = Field(default=4096, description="Maximum tokens for the response")
include: Optional[ReflectIncludeOptions] = Field(default=None, description="Options for including additional data (disabled by default)")
response_schema: Optional[Dict[str, Any]] = None
__properties: ClassVar[List[str]] = ["query", "budget", "context", "max_tokens", "include", "response_schema"]
tags: Optional[List[StrictStr]] = None
tags_match: Optional[StrictStr] = Field(default='any', description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).")
__properties: ClassVar[List[str]] = ["query", "budget", "context", "max_tokens", "include", "response_schema", "tags", "tags_match"]
@field_validator('tags_match')
def tags_match_validate_enum(cls, value):
"""Validates the enum"""
if value is None:
return value
if value not in set(['any', 'all', 'any_strict', 'all_strict']):
raise ValueError("must be one of enum values ('any', 'all', 'any_strict', 'all_strict')")
return value
model_config = ConfigDict(
populate_by_name=True,
@@ -88,6 +100,11 @@ class ReflectRequest(BaseModel):
if self.response_schema is None and "response_schema" in self.model_fields_set:
_dict['response_schema'] = None
# set to None if tags (nullable) is None
# and model_fields_set contains the field
if self.tags is None and "tags" in self.model_fields_set:
_dict['tags'] = None
return _dict
@classmethod
@@ -105,7 +122,9 @@ class ReflectRequest(BaseModel):
"context": obj.get("context"),
"max_tokens": obj.get("max_tokens") if obj.get("max_tokens") is not None else 4096,
"include": ReflectIncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None,
"response_schema": obj.get("response_schema")
"response_schema": obj.get("response_schema"),
"tags": obj.get("tags"),
"tags_match": obj.get("tags_match") if obj.get("tags_match") is not None else 'any'
})
return _obj
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, Field, StrictBool
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictStr
from typing import Any, ClassVar, Dict, List, Optional
from hindsight_client_api.models.memory_item import MemoryItem
from typing import Optional, Set
@@ -29,7 +29,8 @@ class RetainRequest(BaseModel):
""" # noqa: E501
items: List[MemoryItem]
var_async: Optional[StrictBool] = Field(default=False, description="If true, process asynchronously in background. If false, wait for completion (default: false)", alias="async")
__properties: ClassVar[List[str]] = ["items", "async"]
document_tags: Optional[List[StrictStr]] = None
__properties: ClassVar[List[str]] = ["items", "async", "document_tags"]
model_config = ConfigDict(
populate_by_name=True,
@@ -77,6 +78,11 @@ class RetainRequest(BaseModel):
if _item_items:
_items.append(_item_items.to_dict())
_dict['items'] = _items
# set to None if document_tags (nullable) is None
# and model_fields_set contains the field
if self.document_tags is None and "document_tags" in self.model_fields_set:
_dict['document_tags'] = None
return _dict
@classmethod
@@ -90,7 +96,8 @@ class RetainRequest(BaseModel):
_obj = cls.model_validate({
"items": [MemoryItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None,
"async": obj.get("async") if obj.get("async") is not None else False
"async": obj.get("async") if obj.get("async") is not None else False,
"document_tags": obj.get("document_tags")
})
return _obj
@@ -0,0 +1,89 @@
# coding: utf-8
"""
Hindsight HTTP API
HTTP API for Hindsight
The version of the OpenAPI document: 0.1.0
Generated by OpenAPI Generator (https://openapi-generator.tech)
Do not edit the class manually.
""" # noqa: E501
from __future__ import annotations
import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
from typing import Any, ClassVar, Dict, List
from typing import Optional, Set
from typing_extensions import Self
class TagItem(BaseModel):
"""
Single tag with usage count.
""" # noqa: E501
tag: StrictStr = Field(description="The tag value")
count: StrictInt = Field(description="Number of memories with this tag")
__properties: ClassVar[List[str]] = ["tag", "count"]
model_config = ConfigDict(
populate_by_name=True,
validate_assignment=True,
protected_namespaces=(),
)
def to_str(self) -> str:
"""Returns the string representation of the model using alias"""
return pprint.pformat(self.model_dump(by_alias=True))
def to_json(self) -> str:
"""Returns the JSON representation of the model using alias"""
# TODO: pydantic v2: use .model_dump_json(by_alias=True, exclude_unset=True) instead
return json.dumps(self.to_dict())
@classmethod
def from_json(cls, json_str: str) -> Optional[Self]:
"""Create an instance of TagItem from a JSON string"""
return cls.from_dict(json.loads(json_str))
def to_dict(self) -> Dict[str, Any]:
"""Return the dictionary representation of the model using alias.
This has the following differences from calling pydantic's
`self.model_dump(by_alias=True)`:
* `None` is only added to the output dict for nullable fields that
were set at model initialization. Other fields with value `None`
are ignored.
"""
excluded_fields: Set[str] = set([
])
_dict = self.model_dump(
by_alias=True,
exclude=excluded_fields,
exclude_none=True,
)
return _dict
@classmethod
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional[Self]:
"""Create an instance of TagItem from a dict"""
if obj is None:
return None
if not isinstance(obj, dict):
return cls.model_validate(obj)
_obj = cls.model_validate({
"tag": obj.get("tag"),
"count": obj.get("count")
})
return _obj
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "hindsight-client"
version = "0.2.1"
version = "0.3.0"
description = "Python client for Hindsight - Semantic memory system with personality-driven thinking"
authors = [
{name = "Hindsight Team"}
+7
View File
@@ -64,6 +64,7 @@ mod tests {
metadata: None,
timestamp: None,
entities: None,
tags: None,
},
types::MemoryItem {
content: "Bob works with Alice on the search team".to_string(),
@@ -72,8 +73,10 @@ mod tests {
metadata: None,
timestamp: None,
entities: None,
tags: None,
},
],
document_tags: None,
};
let retain_response = client
.retain_memories(&bank_id, None, &retain_request)
@@ -90,6 +93,8 @@ mod tests {
include: None,
query_timestamp: None,
types: None,
tags: None,
tags_match: types::TagsMatch::Any,
};
let recall_response = client
.recall_memories(&bank_id, None, &recall_request)
@@ -106,6 +111,8 @@ mod tests {
max_tokens: 4096,
include: None,
response_schema: None,
tags: None,
tags_match: types::TagsMatch::Any,
};
let reflect_response = client
.reflect(&bank_id, None, &reflect_request)
@@ -39,6 +39,9 @@ import type {
GetGraphData,
GetGraphErrors,
GetGraphResponses,
GetMemoryData,
GetMemoryErrors,
GetMemoryResponses,
HealthEndpointHealthGetData,
HealthEndpointHealthGetResponses,
ListBanksData,
@@ -56,6 +59,9 @@ import type {
ListOperationsData,
ListOperationsErrors,
ListOperationsResponses,
ListTagsData,
ListTagsErrors,
ListTagsResponses,
MetricsEndpointMetricsGetData,
MetricsEndpointMetricsGetResponses,
RecallMemoriesData,
@@ -148,6 +154,20 @@ export const listMemories = <ThrowOnError extends boolean = false>(
ThrowOnError
>({ url: "/v1/default/banks/{bank_id}/memories/list", ...options });
/**
* Get memory unit
*
* Get a single memory unit by ID with all its metadata including entities and tags.
*/
export const getMemory = <ThrowOnError extends boolean = false>(
options: Options<GetMemoryData, ThrowOnError>,
) =>
(options.client ?? client).get<
GetMemoryResponses,
GetMemoryErrors,
ThrowOnError
>({ url: "/v1/default/banks/{bank_id}/memories/{memory_id}", ...options });
/**
* Recall memory
*
@@ -329,6 +349,20 @@ export const getDocument = <ThrowOnError extends boolean = false>(
ThrowOnError
>({ url: "/v1/default/banks/{bank_id}/documents/{document_id}", ...options });
/**
* List tags
*
* List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
*/
export const listTags = <ThrowOnError extends boolean = false>(
options: Options<ListTagsData, ThrowOnError>,
) =>
(options.client ?? client).get<
ListTagsResponses,
ListTagsErrors,
ThrowOnError
>({ url: "/v1/default/banks/{bank_id}/tags", ...options });
/**
* Get chunk details
*
@@ -377,6 +377,12 @@ export type DocumentResponse = {
* Memory Unit Count
*/
memory_unit_count: number;
/**
* Tags
*
* Tags associated with this document
*/
tags?: Array<string>;
};
/**
@@ -666,6 +672,30 @@ export type ListMemoryUnitsResponse = {
offset: number;
};
/**
* ListTagsResponse
*
* Response model for list tags endpoint.
*/
export type ListTagsResponse = {
/**
* Items
*/
items: Array<TagItem>;
/**
* Total
*/
total: number;
/**
* Limit
*/
limit: number;
/**
* Offset
*/
offset: number;
};
/**
* MemoryItem
*
@@ -702,6 +732,12 @@ export type MemoryItem = {
* Optional entities to combine with auto-extracted entities.
*/
entities?: Array<EntityInput> | null;
/**
* Tags
*
* Optional tags for visibility scoping. Memories with tags can be filtered during recall.
*/
tags?: Array<string> | null;
};
/**
@@ -791,6 +827,18 @@ export type RecallRequest = {
* Options for including additional data (entities are included by default)
*/
include?: IncludeOptions;
/**
* Tags
*
* Filter memories by tags. If not specified, all memories are returned.
*/
tags?: Array<string> | null;
/**
* Tags Match
*
* How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).
*/
tags_match?: "any" | "all" | "any_strict" | "all_strict";
};
/**
@@ -879,6 +927,10 @@ export type RecallResult = {
* Chunk Id
*/
chunk_id?: string | null;
/**
* Tags
*/
tags?: Array<string> | null;
};
/**
@@ -958,6 +1010,18 @@ export type ReflectRequest = {
response_schema?: {
[key: string]: unknown;
} | null;
/**
* Tags
*
* Filter memories by tags during reflection. If not specified, all memories are considered.
*/
tags?: Array<string> | null;
/**
* Tags Match
*
* How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).
*/
tags_match?: "any" | "all" | "any_strict" | "all_strict";
};
/**
@@ -1004,6 +1068,12 @@ export type RetainRequest = {
* If true, process asynchronously in background. If false, wait for completion (default: false)
*/
async?: boolean;
/**
* Document Tags
*
* Tags applied to all items in this request. These are merged with any item-level tags.
*/
document_tags?: Array<string> | null;
};
/**
@@ -1042,6 +1112,26 @@ export type RetainResponse = {
usage?: TokenUsage | null;
};
/**
* TagItem
*
* Single tag with usage count.
*/
export type TagItem = {
/**
* Tag
*
* The tag value
*/
tag: string;
/**
* Count
*
* Number of memories with this tag
*/
count: number;
};
/**
* TokenUsage
*
@@ -1225,6 +1315,44 @@ export type ListMemoriesResponses = {
export type ListMemoriesResponse =
ListMemoriesResponses[keyof ListMemoriesResponses];
export type GetMemoryData = {
body?: never;
headers?: {
/**
* Authorization
*/
authorization?: string | null;
};
path: {
/**
* Bank Id
*/
bank_id: string;
/**
* Memory Id
*/
memory_id: string;
};
query?: never;
url: "/v1/default/banks/{bank_id}/memories/{memory_id}";
};
export type GetMemoryErrors = {
/**
* Validation Error
*/
422: HttpValidationError;
};
export type GetMemoryError = GetMemoryErrors[keyof GetMemoryErrors];
export type GetMemoryResponses = {
/**
* Successful Response
*/
200: unknown;
};
export type RecallMemoriesData = {
body: RecallRequest;
headers?: {
@@ -1632,6 +1760,61 @@ export type GetDocumentResponses = {
export type GetDocumentResponse =
GetDocumentResponses[keyof GetDocumentResponses];
export type ListTagsData = {
body?: never;
headers?: {
/**
* Authorization
*/
authorization?: string | null;
};
path: {
/**
* Bank Id
*/
bank_id: string;
};
query?: {
/**
* Q
*
* Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
*/
q?: string | null;
/**
* Limit
*
* Maximum number of tags to return
*/
limit?: number;
/**
* Offset
*
* Offset for pagination
*/
offset?: number;
};
url: "/v1/default/banks/{bank_id}/tags";
};
export type ListTagsErrors = {
/**
* Validation Error
*/
422: HttpValidationError;
};
export type ListTagsError = ListTagsErrors[keyof ListTagsErrors];
export type ListTagsResponses = {
/**
* Successful Response
*/
200: ListTagsResponse;
};
export type ListTagsResponse2 = ListTagsResponses[keyof ListTagsResponses];
export type GetChunkData = {
body?: never;
headers?: {
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@vectorize-io/hindsight-client",
"version": "0.2.1",
"version": "0.3.0",
"description": "TypeScript client for Hindsight - Semantic memory system with personality-driven thinking",
"main": "./dist/src/index.js",
"types": "./dist/src/index.d.ts",
+26 -2
View File
@@ -62,6 +62,7 @@ export interface MemoryItemInput {
metadata?: Record<string, string>;
document_id?: string;
entities?: EntityInput[];
tags?: string[];
}
export class HindsightClient {
@@ -101,6 +102,8 @@ export class HindsightClient {
documentId?: string;
async?: boolean;
entities?: EntityInput[];
/** Optional list of tags for this memory */
tags?: string[];
}
): Promise<RetainResponse> {
const item: {
@@ -110,6 +113,7 @@ export class HindsightClient {
metadata?: Record<string, string>;
document_id?: string;
entities?: EntityInput[];
tags?: string[];
} = { content };
if (options?.timestamp) {
item.timestamp =
@@ -129,6 +133,9 @@ export class HindsightClient {
if (options?.entities) {
item.entities = options.entities;
}
if (options?.tags) {
item.tags = options.tags;
}
const response = await sdk.retainMemories({
client: this.client,
@@ -142,13 +149,14 @@ export class HindsightClient {
/**
* Retain multiple memories in batch.
*/
async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; async?: boolean }): Promise<RetainResponse> {
async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; documentTags?: string[]; async?: boolean }): Promise<RetainResponse> {
const processedItems = items.map((item) => ({
content: item.content,
context: item.context,
metadata: item.metadata,
document_id: item.document_id,
entities: item.entities,
tags: item.tags,
timestamp:
item.timestamp instanceof Date
? item.timestamp.toISOString()
@@ -166,6 +174,7 @@ export class HindsightClient {
path: { bank_id: bankId },
body: {
items: itemsWithDocId,
document_tags: options?.documentTags,
async: options?.async,
},
});
@@ -189,6 +198,10 @@ export class HindsightClient {
maxEntityTokens?: number;
includeChunks?: boolean;
maxChunkTokens?: number;
/** Optional list of tags to filter memories by */
tags?: string[];
/** How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged). Default: 'any' */
tagsMatch?: 'any' | 'all' | 'any_strict' | 'all_strict';
}
): Promise<RecallResponse> {
const response = await sdk.recallMemories({
@@ -205,6 +218,8 @@ export class HindsightClient {
entities: options?.includeEntities ? { max_tokens: options?.maxEntityTokens ?? 500 } : undefined,
chunks: options?.includeChunks ? { max_tokens: options?.maxChunkTokens ?? 8192 } : undefined,
},
tags: options?.tags,
tags_match: options?.tagsMatch,
},
});
@@ -217,7 +232,14 @@ export class HindsightClient {
async reflect(
bankId: string,
query: string,
options?: { context?: string; budget?: Budget }
options?: {
context?: string;
budget?: Budget;
/** Optional list of tags to filter memories by */
tags?: string[];
/** How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged). Default: 'any' */
tagsMatch?: 'any' | 'all' | 'any_strict' | 'all_strict';
}
): Promise<ReflectResponse> {
const response = await sdk.reflect({
client: this.client,
@@ -226,6 +248,8 @@ export class HindsightClient {
query,
context: options?.context,
budget: options?.budget || 'low',
tags: options?.tags,
tags_match: options?.tagsMatch,
},
});
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@vectorize-io/hindsight-control-plane",
"version": "0.2.1",
"version": "0.3.0",
"description": "Control plane for Hindsight - Semantic memory system",
"bin": {
"hindsight-control-plane": "./bin/cli.js"
@@ -0,0 +1,41 @@
import { NextRequest, NextResponse } from "next/server";
const DATAPLANE_URL = process.env.HINDSIGHT_CP_DATAPLANE_API_URL || "http://localhost:8888";
export async function GET(
request: NextRequest,
{ params }: { params: Promise<{ memoryId: string }> }
) {
try {
const { memoryId } = await params;
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get("bank_id");
if (!bankId) {
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const response = await fetch(
`${DATAPLANE_URL}/v1/default/banks/${bankId}/memories/${memoryId}`,
{
method: "GET",
headers: {
"Content-Type": "application/json",
},
}
);
if (!response.ok) {
if (response.status === 404) {
return NextResponse.json({ error: "Memory not found" }, { status: 404 });
}
throw new Error(`API returned ${response.status}`);
}
const data = await response.json();
return NextResponse.json(data, { status: 200 });
} catch (error) {
console.error("Error fetching memory:", error);
return NextResponse.json({ error: "Failed to fetch memory" }, { status: 500 });
}
}
@@ -10,9 +10,12 @@ export async function POST(request: NextRequest) {
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const { items, document_id } = body;
const { items, document_id, document_tags } = body;
const response = await hindsightClient.retainBatch(bankId, items, { documentId: document_id });
const response = await hindsightClient.retainBatch(bankId, items, {
documentId: document_id,
documentTags: document_tags,
});
return NextResponse.json(response, { status: 200 });
} catch (error) {
@@ -5,7 +5,18 @@ export async function POST(request: NextRequest) {
try {
const body = await request.json();
const bankId = body.bank_id || body.agent_id || "default";
const { query, types, fact_type, max_tokens, trace, budget, include, query_timestamp } = body;
const {
query,
types,
fact_type,
max_tokens,
trace,
budget,
include,
query_timestamp,
tags,
tags_match,
} = body;
const response = await sdk.recallMemories({
client: lowLevelClient,
@@ -18,6 +29,8 @@ export async function POST(request: NextRequest) {
budget: budget || "mid",
include,
query_timestamp,
tags,
tags_match,
},
});
@@ -5,12 +5,14 @@ export async function POST(request: NextRequest) {
try {
const body = await request.json();
const bankId = body.bank_id || body.agent_id || "default";
const { query, context, budget, thinking_budget, include_facts } = body;
const { query, context, budget, thinking_budget, include_facts, tags, tags_match } = body;
const requestBody: any = {
query,
budget: budget || (thinking_budget ? "mid" : "low"),
context: context || undefined,
tags,
tags_match,
};
// Add include options if specified
@@ -7,6 +7,7 @@ import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
import { Checkbox } from "@/components/ui/checkbox";
import { Tag } from "lucide-react";
export function AddMemoryView() {
const { currentBank } = useBank();
@@ -14,6 +15,7 @@ export function AddMemoryView() {
const [context, setContext] = useState("");
const [eventDate, setEventDate] = useState("");
const [documentId, setDocumentId] = useState("");
const [tags, setTags] = useState("");
const [async, setAsync] = useState(false);
const [loading, setLoading] = useState(false);
const [result, setResult] = useState<string | null>(null);
@@ -23,6 +25,7 @@ export function AddMemoryView() {
setContext("");
setEventDate("");
setDocumentId("");
setTags("");
setAsync(false);
setResult(null);
};
@@ -37,16 +40,24 @@ export function AddMemoryView() {
setResult(null);
try {
// Parse tags from comma-separated string
const parsedTags = tags
.split(",")
.map((t) => t.trim())
.filter((t) => t.length > 0);
const item: any = { content };
if (context) item.context = context;
// datetime-local gives "2024-01-15T10:30", add seconds for proper ISO format
if (eventDate) item.timestamp = eventDate + ":00";
if (parsedTags.length > 0) item.tags = parsedTags;
const data: any = await client.retain({
bank_id: currentBank,
items: [item],
document_id: documentId,
async,
...(parsedTags.length > 0 && { document_tags: parsedTags }),
});
setResult(data.message as string);
@@ -112,6 +123,22 @@ export function AddMemoryView() {
</small>
</div>
<div className="mb-4">
<label className="font-bold block mb-1 text-card-foreground flex items-center gap-2">
<Tag className="h-4 w-4" />
Tags
</label>
<Input
type="text"
value={tags}
onChange={(e) => setTags(e.target.value)}
placeholder="user_alice, session_123, project_x"
/>
<small className="text-muted-foreground text-xs mt-1 block">
Comma-separated tags for filtering during recall/reflect. Tags cannot contain commas.
</small>
</div>
<div className="mb-4">
<div className="flex items-center gap-2">
<Checkbox
@@ -23,7 +23,7 @@ import {
DialogFooter,
} from "@/components/ui/dialog";
import { Input } from "@/components/ui/input";
import { Check, ChevronsUpDown, Plus, FileText, Moon, Sun, Github } from "lucide-react";
import { Check, ChevronsUpDown, Plus, FileText, Moon, Sun, Github, Tag } from "lucide-react";
import { useTheme } from "@/lib/theme-context";
import Image from "next/image";
import { Textarea } from "@/components/ui/textarea";
@@ -47,6 +47,7 @@ function BankSelectorInner() {
const [docContext, setDocContext] = React.useState("");
const [docEventDate, setDocEventDate] = React.useState("");
const [docDocumentId, setDocDocumentId] = React.useState("");
const [docTags, setDocTags] = React.useState("");
const [docAsync, setDocAsync] = React.useState(false);
const [isCreatingDoc, setIsCreatingDoc] = React.useState(false);
const [docError, setDocError] = React.useState<string | null>(null);
@@ -83,10 +84,17 @@ function BankSelectorInner() {
setDocError(null);
try {
// Parse tags from comma-separated string
const parsedTags = docTags
.split(",")
.map((t) => t.trim())
.filter((t) => t.length > 0);
const item: any = { content: docContent };
if (docContext) item.context = docContext;
// datetime-local gives "2024-01-15T10:30", add seconds for proper ISO format
if (docEventDate) item.timestamp = docEventDate + ":00";
if (parsedTags.length > 0) item.tags = parsedTags;
const params: any = {
bank_id: currentBank,
@@ -94,6 +102,7 @@ function BankSelectorInner() {
};
if (docDocumentId) params.document_id = docDocumentId;
if (parsedTags.length > 0) params.document_tags = parsedTags;
if (docAsync) {
await client.retain({ ...params, async: true });
@@ -107,6 +116,7 @@ function BankSelectorInner() {
setDocContext("");
setDocEventDate("");
setDocDocumentId("");
setDocTags("");
setDocAsync(false);
// Navigate to documents view to see the new document
@@ -335,6 +345,22 @@ function BankSelectorInner() {
</div>
</div>
<div>
<label className="font-bold block mb-1 text-sm text-foreground flex items-center gap-2">
<Tag className="h-4 w-4" />
Tags
</label>
<Input
type="text"
value={docTags}
onChange={(e) => setDocTags(e.target.value)}
placeholder="user_alice, session_123, project_x"
/>
<p className="text-xs text-muted-foreground mt-1">
Comma-separated tags for filtering during recall/reflect
</p>
</div>
<div className="flex items-center gap-2">
<Checkbox
id="async-doc"
@@ -357,6 +383,7 @@ function BankSelectorInner() {
setDocContext("");
setDocEventDate("");
setDocDocumentId("");
setDocTags("");
setDocAsync(false);
setDocError(null);
}}
@@ -362,6 +362,7 @@ export function DataView({ factType }: DataViewProps) {
memory={selectedGraphNode}
onClose={() => setSelectedGraphNode(null)}
inPanel
bankId={currentBank || undefined}
/>
) : (
/* Legend & Controls View */
@@ -738,13 +739,20 @@ export function DataView({ factType }: DataViewProps) {
memory={selectedTableMemory}
onClose={() => setSelectedTableMemory(null)}
inPanel
bankId={currentBank || undefined}
/>
</div>
)}
</div>
)}
{viewMode === "timeline" && <TimelineView data={data} filteredRows={filteredTableRows} />}
{viewMode === "timeline" && (
<TimelineView
data={data}
filteredRows={filteredTableRows}
bankId={currentBank || undefined}
/>
)}
</>
) : (
<div className="flex items-center justify-center py-20">
@@ -761,7 +769,15 @@ export function DataView({ factType }: DataViewProps) {
// Timeline View Component - Custom compact timeline with zoom and navigation
type Granularity = "year" | "month" | "week" | "day";
function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }) {
function TimelineView({
data,
filteredRows,
bankId,
}: {
data: any;
filteredRows: any[];
bankId?: string;
}) {
const [selectedItem, setSelectedItem] = useState<any>(null);
const [granularity, setGranularity] = useState<Granularity>("month");
const [currentIndex, setCurrentIndex] = useState(0);
@@ -1114,7 +1130,12 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
{/* Detail Panel - Fixed on Right */}
{selectedItem && (
<div className="fixed right-0 top-0 h-screen w-[420px] bg-card border-l-2 border-primary shadow-2xl z-50 overflow-y-auto animate-in slide-in-from-right duration-300 ease-out">
<MemoryDetailPanel memory={selectedItem} onClose={() => setSelectedItem(null)} inPanel />
<MemoryDetailPanel
memory={selectedItem}
onClose={() => setSelectedItem(null)}
inPanel
bankId={bankId}
/>
</div>
)}
</div>
@@ -318,6 +318,25 @@ export function DocumentsView() {
</div>
)}
{/* Tags */}
{selectedDocument.tags && selectedDocument.tags.length > 0 && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Tags
</div>
<div className="flex flex-wrap gap-2">
{selectedDocument.tags.map((tag: string, i: number) => (
<span
key={i}
className="text-sm px-3 py-1.5 rounded-full bg-amber-500/10 text-amber-600 dark:text-amber-400 font-medium"
>
{tag}
</span>
))}
</div>
</div>
)}
{/* Delete Button */}
<div className="pt-2 border-t border-border">
<Button
@@ -1,15 +1,17 @@
"use client";
import { useState } from "react";
import { useState, useEffect } from "react";
import { Button } from "@/components/ui/button";
import { Copy, Check, X } from "lucide-react";
import { Copy, Check, X, Loader2 } from "lucide-react";
import { DocumentChunkModal } from "./document-chunk-modal";
import { client } from "@/lib/api";
interface MemoryDetailPanelProps {
memory: any;
onClose: () => void;
compact?: boolean;
inPanel?: boolean;
bankId?: string;
}
export function MemoryDetailPanel({
@@ -17,10 +19,40 @@ export function MemoryDetailPanel({
onClose,
compact = false,
inPanel = false,
bankId,
}: MemoryDetailPanelProps) {
const [copiedId, setCopiedId] = useState<string | null>(null);
const [modalType, setModalType] = useState<"document" | "chunk" | null>(null);
const [modalId, setModalId] = useState<string | null>(null);
const [fullMemory, setFullMemory] = useState<any>(null);
const [loading, setLoading] = useState(false);
// Fetch full memory data when panel opens
useEffect(() => {
const memoryId = memory?.id || memory?.node_id;
if (!memoryId || !bankId) {
setFullMemory(null);
return;
}
setLoading(true);
client
.getMemory(memoryId, bankId)
.then((data) => {
setFullMemory(data);
})
.catch((err) => {
console.error("Failed to fetch memory details:", err);
// Fall back to showing the partial data we have
setFullMemory(null);
})
.finally(() => {
setLoading(false);
});
}, [memory?.id, memory?.node_id, bankId]);
// Use full memory data if available, otherwise fall back to the partial data passed in
const displayMemory = fullMemory || memory;
const copyToClipboard = async (text: string) => {
try {
@@ -50,7 +82,7 @@ export function MemoryDetailPanel({
if (!memory) return null;
// Handle both 'id' and 'node_id' (trace results use node_id)
const memoryId = memory.id || memory.node_id;
const memoryId = displayMemory.id || displayMemory.node_id;
const labelSize = compact ? "text-[10px]" : "text-xs";
const textSize = compact ? "text-xs" : "text-sm";
@@ -71,123 +103,156 @@ export function MemoryDetailPanel({
</Button>
</div>
<div className="space-y-5">
{/* Full Text */}
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Full Text
</div>
<div className="text-sm whitespace-pre-wrap leading-relaxed text-foreground">
{memory.text}
</div>
{loading ? (
<div className="flex items-center justify-center py-12">
<Loader2 className="h-6 w-6 animate-spin text-muted-foreground" />
<span className="ml-2 text-muted-foreground">Loading memory details...</span>
</div>
{/* Context */}
{memory.context && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Context
</div>
<div className="text-sm text-foreground">{memory.context}</div>
</div>
)}
{/* Dates */}
<div className="grid grid-cols-2 gap-4">
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Occurred
</div>
<div className="text-sm font-medium text-foreground">
{memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"}
</div>
</div>
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Mentioned
</div>
<div className="text-sm font-medium text-foreground">
{memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"}
</div>
</div>
</div>
{/* Entities */}
{memory.entities && (
) : (
<div className="space-y-5">
{/* Full Text */}
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">
Entities
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Full Text
</div>
<div className="flex flex-wrap gap-2">
{(Array.isArray(memory.entities)
? memory.entities
: String(memory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string" ? entity : entity?.name || JSON.stringify(entity);
return (
<div className="text-sm whitespace-pre-wrap leading-relaxed text-foreground">
{displayMemory.text}
</div>
</div>
{/* Context */}
{displayMemory.context && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Context
</div>
<div className="text-sm text-foreground">{displayMemory.context}</div>
</div>
)}
{/* Dates */}
<div className="grid grid-cols-2 gap-4">
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Occurred
</div>
<div className="text-sm font-medium text-foreground">
{displayMemory.occurred_start
? new Date(displayMemory.occurred_start).toLocaleString()
: "N/A"}
</div>
</div>
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Mentioned
</div>
<div className="text-sm font-medium text-foreground">
{displayMemory.mentioned_at
? new Date(displayMemory.mentioned_at).toLocaleString()
: "N/A"}
</div>
</div>
</div>
{/* Entities */}
{displayMemory.entities &&
(Array.isArray(displayMemory.entities)
? displayMemory.entities.length > 0
: displayMemory.entities) && (
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">
Entities
</div>
<div className="flex flex-wrap gap-2">
{(Array.isArray(displayMemory.entities)
? displayMemory.entities
: String(displayMemory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string"
? entity
: entity?.name || JSON.stringify(entity);
return (
<span
key={i}
className="text-sm px-3 py-1.5 rounded-full bg-primary/10 text-primary font-medium"
>
{entityText}
</span>
);
})}
</div>
</div>
)}
{/* Tags */}
{displayMemory.tags && displayMemory.tags.length > 0 && (
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">Tags</div>
<div className="flex flex-wrap gap-2">
{displayMemory.tags.map((tag: string, i: number) => (
<span
key={i}
className="text-sm px-3 py-1.5 rounded-full bg-primary/10 text-primary font-medium"
className="text-sm px-3 py-1.5 rounded-full bg-amber-500/10 text-amber-600 dark:text-amber-400 font-medium"
>
{entityText}
{tag}
</span>
);
})}
))}
</div>
</div>
</div>
)}
)}
{/* ID */}
{memoryId && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Memory ID
{/* ID */}
{memoryId && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Memory ID
</div>
<div className="flex items-center gap-2">
<code className="text-xs font-mono break-all flex-1 text-muted-foreground">
{memoryId}
</code>
<Button
variant="ghost"
size="sm"
className="h-8 w-8 p-0 flex-shrink-0"
onClick={() => copyToClipboard(memoryId)}
>
{copiedId === memoryId ? (
<Check className="h-4 w-4 text-green-600" />
) : (
<Copy className="h-4 w-4" />
)}
</Button>
</div>
</div>
<div className="flex items-center gap-2">
<code className="text-xs font-mono break-all flex-1 text-muted-foreground">
{memoryId}
</code>
<Button
variant="ghost"
size="sm"
className="h-8 w-8 p-0 flex-shrink-0"
onClick={() => copyToClipboard(memoryId)}
>
{copiedId === memoryId ? (
<Check className="h-4 w-4 text-green-600" />
) : (
<Copy className="h-4 w-4" />
)}
</Button>
</div>
</div>
)}
)}
{/* Document/Chunk buttons */}
{(memory.document_id || memory.chunk_id) && (
<div className="flex gap-3 pt-2">
{memory.document_id && (
<Button
onClick={() => openDocumentModal(memory.document_id)}
variant="secondary"
className="flex-1"
>
View Document
</Button>
)}
{memory.chunk_id && (
<Button
onClick={() => openChunkModal(memory.chunk_id)}
variant="secondary"
className="flex-1"
>
View Chunk
</Button>
)}
</div>
)}
</div>
{/* Document/Chunk buttons */}
{(displayMemory.document_id || displayMemory.chunk_id) && (
<div className="flex gap-3 pt-2">
{displayMemory.document_id && (
<Button
onClick={() => openDocumentModal(displayMemory.document_id)}
variant="secondary"
className="flex-1"
>
View Document
</Button>
)}
{displayMemory.chunk_id && (
<Button
onClick={() => openChunkModal(displayMemory.chunk_id)}
variant="secondary"
className="flex-1"
>
View Chunk
</Button>
)}
</div>
)}
</div>
)}
</div>
{/* Document/Chunk Modal */}
@@ -225,123 +290,158 @@ export function MemoryDetailPanel({
</Button>
</div>
<div className={gap}>
{/* Full Text */}
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Full Text
</div>
<div className={`${textSize} whitespace-pre-wrap`}>{memory.text}</div>
{loading ? (
<div className="flex items-center justify-center py-8">
<Loader2 className="h-5 w-5 animate-spin text-muted-foreground" />
<span className="ml-2 text-sm text-muted-foreground">Loading...</span>
</div>
{/* Context */}
{memory.context && (
) : (
<div className={gap}>
{/* Full Text */}
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Context
Full Text
</div>
<div className={textSize}>{memory.context}</div>
<div className={`${textSize} whitespace-pre-wrap`}>{displayMemory.text}</div>
</div>
)}
{/* Dates */}
<div className="grid grid-cols-2 gap-2">
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Occurred
{/* Context */}
{displayMemory.context && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Context
</div>
<div className={textSize}>{displayMemory.context}</div>
</div>
<div className={textSize}>
{memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"}
</div>
</div>
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Mentioned
</div>
<div className={textSize}>
{memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"}
</div>
</div>
</div>
)}
{/* Entities */}
{memory.entities && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
Entities
{/* Dates */}
<div className="grid grid-cols-2 gap-2">
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Occurred
</div>
<div className={textSize}>
{displayMemory.occurred_start
? new Date(displayMemory.occurred_start).toLocaleString()
: "N/A"}
</div>
</div>
<div className="flex flex-wrap gap-1">
{(Array.isArray(memory.entities)
? memory.entities
: String(memory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string" ? entity : entity?.name || JSON.stringify(entity);
return (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Mentioned
</div>
<div className={textSize}>
{displayMemory.mentioned_at
? new Date(displayMemory.mentioned_at).toLocaleString()
: "N/A"}
</div>
</div>
</div>
{/* Entities */}
{displayMemory.entities &&
(Array.isArray(displayMemory.entities)
? displayMemory.entities.length > 0
: displayMemory.entities) && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
Entities
</div>
<div className="flex flex-wrap gap-1">
{(Array.isArray(displayMemory.entities)
? displayMemory.entities
: String(displayMemory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string"
? entity
: entity?.name || JSON.stringify(entity);
return (
<span
key={i}
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-secondary text-secondary-foreground`}
>
{entityText}
</span>
);
})}
</div>
</div>
)}
{/* Tags */}
{displayMemory.tags && displayMemory.tags.length > 0 && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
Tags
</div>
<div className="flex flex-wrap gap-1">
{displayMemory.tags.map((tag: string, i: number) => (
<span
key={i}
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-secondary text-secondary-foreground`}
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-amber-500/10 text-amber-600 dark:text-amber-400`}
>
{entityText}
{tag}
</span>
);
})}
))}
</div>
</div>
</div>
)}
)}
{/* ID */}
{memoryId && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Memory ID
{/* ID */}
{memoryId && (
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Memory ID
</div>
<div className="flex items-center gap-2">
<span className={`${compact ? "text-[10px]" : "text-sm"} font-mono break-all`}>
{memoryId}
</span>
<Button
variant="ghost"
size="sm"
className="h-6 w-6 p-0 flex-shrink-0"
onClick={() => copyToClipboard(memoryId)}
>
{copiedId === memoryId ? (
<Check className="h-3 w-3 text-green-600" />
) : (
<Copy className="h-3 w-3" />
)}
</Button>
</div>
</div>
<div className="flex items-center gap-2">
<span className={`${compact ? "text-[10px]" : "text-sm"} font-mono break-all`}>
{memoryId}
</span>
<Button
variant="ghost"
size="sm"
className="h-6 w-6 p-0 flex-shrink-0"
onClick={() => copyToClipboard(memoryId)}
>
{copiedId === memoryId ? (
<Check className="h-3 w-3 text-green-600" />
) : (
<Copy className="h-3 w-3" />
)}
</Button>
</div>
</div>
)}
)}
{/* Document/Chunk buttons */}
{(memory.document_id || memory.chunk_id) && (
<div className={`flex gap-2 ${compact ? "pt-1" : ""}`}>
{memory.document_id && (
<Button
onClick={() => openDocumentModal(memory.document_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Document
</Button>
)}
{memory.chunk_id && (
<Button
onClick={() => openChunkModal(memory.chunk_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Chunk
</Button>
)}
</div>
)}
</div>
{/* Document/Chunk buttons */}
{(displayMemory.document_id || displayMemory.chunk_id) && (
<div className={`flex gap-2 ${compact ? "pt-1" : ""}`}>
{displayMemory.document_id && (
<Button
onClick={() => openDocumentModal(displayMemory.document_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Document
</Button>
)}
{displayMemory.chunk_id && (
<Button
onClick={() => openChunkModal(displayMemory.chunk_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Chunk
</Button>
)}
</div>
)}
</div>
)}
</div>
{/* Document/Chunk Modal */}
@@ -25,6 +25,8 @@ import {
FileText,
Users,
ArrowDown,
Tag,
Calendar,
} from "lucide-react";
import JsonView from "react18-json-view";
import "react18-json-view/src/style.css";
@@ -32,6 +34,7 @@ import { MemoryDetailPanel } from "./memory-detail-panel";
type FactType = "world" | "experience" | "opinion";
type Budget = "low" | "mid" | "high";
type TagsMatch = "any" | "all" | "any_strict" | "all_strict";
type ViewMode = "results" | "trace" | "json";
export function SearchDebugView() {
@@ -45,6 +48,8 @@ export function SearchDebugView() {
const [queryDate, setQueryDate] = useState("");
const [includeChunks, setIncludeChunks] = useState(false);
const [includeEntities, setIncludeEntities] = useState(false);
const [tags, setTags] = useState("");
const [tagsMatch, setTagsMatch] = useState<TagsMatch>("any");
// Results state
const [results, setResults] = useState<any[] | null>(null);
@@ -83,6 +88,14 @@ export function SearchDebugView() {
const INITIAL_RESULTS_COUNT = 5;
// Helper to find full memory data from results when clicking trace items
const selectMemoryFromTrace = (traceResult: any) => {
const nodeId = traceResult.id || traceResult.node_id;
// Try to find the full result with all metadata
const fullResult = results?.find((r: any) => r.id === nodeId || r.node_id === nodeId);
setSelectedMemory(fullResult || traceResult);
};
const runSearch = async () => {
if (!currentBank) {
alert("Please select a memory bank first");
@@ -99,6 +112,12 @@ export function SearchDebugView() {
setLoading(true);
try {
// Parse tags from comma-separated string
const parsedTags = tags
.split(",")
.map((t) => t.trim())
.filter((t) => t.length > 0);
const requestBody: any = {
bank_id: currentBank,
query: query,
@@ -111,6 +130,7 @@ export function SearchDebugView() {
chunks: includeChunks ? { max_tokens: 8192 } : null,
},
...(queryDate && { query_timestamp: queryDate }),
...(parsedTags.length > 0 && { tags: parsedTags, tags_match: tagsMatch }),
};
const data: any = await client.recall(requestBody);
@@ -246,6 +266,31 @@ export function SearchDebugView() {
</label>
</div>
</div>
{/* Tags Filter */}
<div className="flex items-center gap-4 mt-4 pt-4 border-t">
<Tag className="h-4 w-4 text-muted-foreground" />
<div className="flex-1 max-w-md">
<Input
type="text"
value={tags}
onChange={(e) => setTags(e.target.value)}
placeholder="Filter by tags (comma-separated)"
className="h-8"
/>
</div>
<Select value={tagsMatch} onValueChange={(v) => setTagsMatch(v as TagsMatch)}>
<SelectTrigger className="w-40 h-8">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="any">Any (incl. untagged)</SelectItem>
<SelectItem value="all">All (incl. untagged)</SelectItem>
<SelectItem value="any_strict">Any (strict)</SelectItem>
<SelectItem value="all_strict">All (strict)</SelectItem>
</SelectContent>
</Select>
</div>
</CardContent>
</Card>
@@ -507,9 +552,29 @@ export function SearchDebugView() {
}}
>
<div className="flex items-center justify-between mb-1">
<span className="font-medium text-sm text-foreground capitalize">
{method.method_name}
</span>
<div className="flex items-center gap-2">
<span className="font-medium text-sm text-foreground capitalize">
{method.method_name}
</span>
{/* Show temporal range inline */}
{method.method_name === "temporal" &&
method.metadata?.constraint && (
<span className="flex items-center gap-1 text-[10px] text-muted-foreground">
<Calendar className="h-3 w-3" />
{method.metadata.constraint.start
? new Date(
method.metadata.constraint.start
).toLocaleDateString()
: "any"}
{" → "}
{method.metadata.constraint.end
? new Date(
method.metadata.constraint.end
).toLocaleDateString()
: "any"}
</span>
)}
</div>
{isMethodExpanded ? (
<ChevronDown className="h-3 w-3 text-muted-foreground" />
) : (
@@ -546,7 +611,7 @@ export function SearchDebugView() {
className="p-2 bg-background rounded cursor-pointer hover:bg-muted/50 transition-colors border border-border"
onClick={(e) => {
e.stopPropagation();
setSelectedMemory(r);
selectMemoryFromTrace(r);
}}
>
<div className="flex items-start gap-2">
@@ -684,7 +749,7 @@ export function SearchDebugView() {
className="p-3 bg-muted/30 rounded-lg cursor-pointer hover:bg-muted/50 transition-colors"
onClick={(e) => {
e.stopPropagation();
setSelectedMemory(r);
selectMemoryFromTrace(r);
}}
>
<div className="flex items-start gap-3">
@@ -792,7 +857,7 @@ export function SearchDebugView() {
className="p-3 bg-muted/30 rounded-lg cursor-pointer hover:bg-muted/50 transition-colors"
onClick={(e) => {
e.stopPropagation();
setSelectedMemory(r);
selectMemoryFromTrace(r);
}}
>
<div className="flex items-start gap-3">
@@ -931,6 +996,7 @@ export function SearchDebugView() {
memory={selectedMemory}
onClose={() => setSelectedMemory(null)}
inPanel
bankId={currentBank || undefined}
/>
</div>
)}
@@ -15,10 +15,12 @@ import {
} from "@/components/ui/select";
import { Checkbox } from "@/components/ui/checkbox";
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card";
import { Sparkles, Info } from "lucide-react";
import { Sparkles, Info, Tag } from "lucide-react";
import JsonView from "react18-json-view";
import "react18-json-view/src/style.css";
type TagsMatch = "any" | "all" | "any_strict" | "all_strict";
export function ThinkView() {
const { currentBank } = useBank();
const [query, setQuery] = useState("");
@@ -28,6 +30,8 @@ export function ThinkView() {
const [showRawJson, setShowRawJson] = useState(false);
const [result, setResult] = useState<any>(null);
const [loading, setLoading] = useState(false);
const [tags, setTags] = useState("");
const [tagsMatch, setTagsMatch] = useState<TagsMatch>("any");
const runReflect = async () => {
if (!currentBank || !query) return;
@@ -35,12 +39,19 @@ export function ThinkView() {
setLoading(true);
setShowRawJson(false);
try {
// Parse tags from comma-separated string
const parsedTags = tags
.split(",")
.map((t) => t.trim())
.filter((t) => t.length > 0);
const data: any = await client.reflect({
bank_id: currentBank,
query,
budget,
context: context || undefined,
include_facts: includeFacts,
...(parsedTags.length > 0 && { tags: parsedTags, tags_match: tagsMatch }),
});
setResult(data);
} catch (error) {
@@ -103,6 +114,29 @@ export function ThinkView() {
rows={3}
/>
</div>
<div className="flex items-center gap-4 mt-4 pt-4 border-t">
<Tag className="h-4 w-4 text-muted-foreground" />
<div className="flex-1 max-w-md">
<Input
type="text"
value={tags}
onChange={(e) => setTags(e.target.value)}
placeholder="Filter by tags (comma-separated)"
className="h-8"
/>
</div>
<Select value={tagsMatch} onValueChange={(v) => setTagsMatch(v as TagsMatch)}>
<SelectTrigger className="w-40 h-8">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="any">Any (incl. untagged)</SelectItem>
<SelectItem value="all">All (incl. untagged)</SelectItem>
<SelectItem value="any_strict">Any (strict)</SelectItem>
<SelectItem value="all_strict">All (strict)</SelectItem>
</SelectContent>
</Select>
</div>
</CardContent>
</Card>
+24
View File
@@ -53,6 +53,8 @@ export class ControlPlaneClient {
chunks?: { max_tokens: number } | null;
};
query_timestamp?: string;
tags?: string[];
tags_match?: "any" | "all" | "any_strict" | "all_strict";
}) {
return this.fetchApi("/api/recall", {
method: "POST",
@@ -69,6 +71,8 @@ export class ControlPlaneClient {
budget?: string;
context?: string;
include_facts?: boolean;
tags?: string[];
tags_match?: "any" | "all" | "any_strict" | "all_strict";
}) {
return this.fetchApi("/api/reflect", {
method: "POST",
@@ -209,6 +213,26 @@ export class ControlPlaneClient {
return this.fetchApi(`/api/chunks/${chunkId}`);
}
/**
* Get a single memory by ID
*/
async getMemory(memoryId: string, bankId: string) {
return this.fetchApi<{
id: string;
text: string;
context: string;
date: string;
type: string;
mentioned_at: string | null;
occurred_start: string | null;
occurred_end: string | null;
entities: string[];
document_id: string | null;
chunk_id: string | null;
tags: string[];
}>(`/api/memories/${memoryId}?bank_id=${bankId}`);
}
/**
* Get bank profile
*/
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "hindsight-dev"
version = "0.2.1"
version = "0.3.0"
description = "Development utilities for Hindsight"
requires-python = ">=3.11"
dependencies = [
+39 -2
View File
@@ -8,11 +8,48 @@ This changelog highlights user-facing changes only. Internal maintenance, CI/CD,
For full release details, see [GitHub Releases](https://github.com/vectorize-io/hindsight/releases).
## [Unreleased]
## [0.3.0](https://github.com/vectorize-io/hindsight/releases/tag/v0.3.0)
**Features**
- Add per-request token usage tracking to retain and reflect endpoints for cost monitoring and billing integration.
- Add memory tags so you can label and filter memories during recall/reflect. ([`20c8f8b`](https://github.com/vectorize-io/hindsight/commit/20c8f8b))
- Allow choosing different AI providers/models per operation. ([`e6709d5`](https://github.com/vectorize-io/hindsight/commit/e6709d5))
- Add Cohere support for embeddings and reranking. ([`4de0730`](https://github.com/vectorize-io/hindsight/commit/4de0730))
- Add configurable embedding dimensions and OpenAI embeddings support. ([`70de23e`](https://github.com/vectorize-io/hindsight/commit/70de23e))
- Support custom base URLs for OpenAI-style embeddings and Cohere endpoints. ([`fa53917`](https://github.com/vectorize-io/hindsight/commit/fa53917))
- Add LiteLLM gateway support for routing LLM/embedding requests. ([`d47c8a2`](https://github.com/vectorize-io/hindsight/commit/d47c8a2))
- Add multilingual content support to improve handling and retrieval across languages. ([`c65c6a9`](https://github.com/vectorize-io/hindsight/commit/c65c6a9))
- Add delete memory bank capability. ([`4b82d2d`](https://github.com/vectorize-io/hindsight/commit/4b82d2d))
- Add backup/restore tooling for memory banks. ([`67b273d`](https://github.com/vectorize-io/hindsight/commit/67b273d))
**Improvements**
- Add retention modes to control how memories are extracted and stored. ([`fb31a35`](https://github.com/vectorize-io/hindsight/commit/fb31a35))
- Add offline (optional) database migrations to support restricted/air-gapped deployments. ([`233bd2e`](https://github.com/vectorize-io/hindsight/commit/233bd2e))
- Add database connection configuration options for more flexible deployments. ([`33fac2c`](https://github.com/vectorize-io/hindsight/commit/33fac2c))
- Load .env automatically on startup to simplify configuration. ([`c06d9b4`](https://github.com/vectorize-io/hindsight/commit/c06d9b4))
- Expose an operation ID from retain requests so async/background processing can be tracked. ([`1dacd0e`](https://github.com/vectorize-io/hindsight/commit/1dacd0e))
- Add per-request LLM token usage metrics for monitoring and cost tracking. ([`29a542d`](https://github.com/vectorize-io/hindsight/commit/29a542d))
- Add LLM call latency metrics for performance monitoring. ([`5e1f13e`](https://github.com/vectorize-io/hindsight/commit/5e1f13e))
- Include tenant in metrics labels for better multi-tenant observability. ([`1ffc2a4`](https://github.com/vectorize-io/hindsight/commit/1ffc2a4))
- Add async processing option to MCP retain tool for background retention workflows. ([`37fc7fb`](https://github.com/vectorize-io/hindsight/commit/37fc7fb))
**Bug Fixes**
- Fix extension loading in multi-worker deployments so all workers load extensions correctly. ([`f5f3fca`](https://github.com/vectorize-io/hindsight/commit/f5f3fca))
- Improve recall performance by batching recall queries. ([`5991308`](https://github.com/vectorize-io/hindsight/commit/5991308))
- Improve retrieval quality and stability for large memory banks (graph/MPFP retrieval fixes). ([`6232e69`](https://github.com/vectorize-io/hindsight/commit/6232e69))
- Fix entities list being limited to 100 entities. ([`26bf571`](https://github.com/vectorize-io/hindsight/commit/26bf571))
- Fix UI only showing the first 1000 memories. ([`67c1a42`](https://github.com/vectorize-io/hindsight/commit/67c1a42))
- Fix duplicated causal relationships and improve token usage during processing. ([`49e233c`](https://github.com/vectorize-io/hindsight/commit/49e233c))
- Improve causal link detection accuracy. ([`2a00df0`](https://github.com/vectorize-io/hindsight/commit/2a00df0))
- Make retain max completion tokens configurable to prevent truncation issues. ([`7715a51`](https://github.com/vectorize-io/hindsight/commit/7715a51))
- Fix Python SDK not sending the Authorization header, preventing authenticated requests. ([`39e3f7c`](https://github.com/vectorize-io/hindsight/commit/39e3f7c))
- Fix stats endpoint missing tenant authentication in multi-tenant setups. ([`d6ff191`](https://github.com/vectorize-io/hindsight/commit/d6ff191))
- Fix embedding dimension handling for tenant schemas in multi-tenant databases. ([`6fe9314`](https://github.com/vectorize-io/hindsight/commit/6fe9314))
- Fix Groq free-tier compatibility so requests work correctly. ([`d899d18`](https://github.com/vectorize-io/hindsight/commit/d899d18))
- Fix security vulnerability (qs / CVE-2025-15284). ([`b3becb6`](https://github.com/vectorize-io/hindsight/commit/b3becb6))
- Restore MCP tools for listing and creating memory banks. ([`9fd5679`](https://github.com/vectorize-io/hindsight/commit/9fd5679))
## [0.2.0](https://github.com/vectorize-io/hindsight/releases/tag/v0.2.0)
+51 -1
View File
@@ -43,11 +43,13 @@ Make sure you've completed the [Quick Start](./quickstart) to install the client
|-----------|------|---------|-------------|
| `query` | string | required | Natural language query |
| `types` | list | all | Filter: `world`, `experience`, `opinion` |
| `budget` | string | "mid" | Budget level: "low", "mid", "high" |
| `budget` | string | "mid" | Budget level: `low`, `mid`, `high` |
| `max_tokens` | int | 4096 | Token budget for results |
| `trace` | bool | false | Enable trace output for debugging |
| `include_entities` | bool | false | Include entity observations |
| `max_entity_tokens` | int | 500 | Token budget for entity observations |
| `tags` | list | None | Filter memories by tags (see [Tag Filtering](#filter-by-tags)) |
| `tags_match` | string | "any" | How to match tags: `any`, `all`, `any_strict`, `all_strict` |
<Tabs>
<TabItem value="python" label="Python">
@@ -127,3 +129,51 @@ The `budget` parameter controls graph traversal depth:
<CodeSnippet code={recallMjs} section="recall-budget-levels" language="javascript" />
</TabItem>
</Tabs>
## Filter by Tags
Tags enable **visibility scoping**—filter memories based on tags assigned during [retain](./retain#tagging-memories). This is essential for multi-user agents where each user should only see their own memories.
### Basic Tag Filtering
<Tabs>
<TabItem value="python" label="Python">
<CodeSnippet code={recallPy} section="recall-with-tags" language="python" />
</TabItem>
</Tabs>
### Tag Match Modes
The `tags_match` parameter controls how tags are matched:
| Mode | Behavior | Untagged Memories |
|------|----------|-------------------|
| `any` | OR: memory has ANY of the specified tags | **Included** |
| `all` | AND: memory has ALL of the specified tags | **Included** |
| `any_strict` | OR: memory has ANY of the specified tags | **Excluded** |
| `all_strict` | AND: memory has ALL of the specified tags | **Excluded** |
**Strict modes** are useful when you want to ensure only tagged memories are returned:
<Tabs>
<TabItem value="python" label="Python">
<CodeSnippet code={recallPy} section="recall-tags-strict" language="python" />
</TabItem>
</Tabs>
**AND matching** requires all specified tags to be present:
<Tabs>
<TabItem value="python" label="Python">
<CodeSnippet code={recallPy} section="recall-tags-all" language="python" />
</TabItem>
</Tabs>
### Use Cases
| Scenario | Tags | Mode | Result |
|----------|------|------|--------|
| User A's memories only | `["user:alice"]` | `any_strict` | Only memories tagged `user:alice` |
| Support + feedback | `["support", "feedback"]` | `any` | Memories with either tag + untagged |
| Multi-user room | `["user:alice", "room:general"]` | `all_strict` | Only memories with both tags |
| Global + user-specific | `["user:alice"]` | `any` | Alice's memories + shared (untagged) |
+24 -1
View File
@@ -50,10 +50,12 @@ Make sure you've completed the [Quick Start](./quickstart) to install the client
| Parameter | Type | Default | Description |
|-----------|------|---------|-------------|
| `query` | string | required | Question or prompt |
| `budget` | string | "low" | Budget level: "low", "mid", "high" |
| `budget` | string | "low" | Budget level: `low`, `mid`, `high` |
| `context` | string | None | Additional context for the query |
| `max_tokens` | int | 4096 | Maximum tokens for the response |
| `response_schema` | object | None | JSON Schema for [structured output](#structured-output) |
| `tags` | list | None | Filter memories by tags during reflection |
| `tags_match` | string | "any" | How to match tags: `any`, `all`, `any_strict`, `all_strict` |
### Response Fields
@@ -247,3 +249,24 @@ hindsight memory reflect hiring-team \
- Use `model_validate()` to parse the response back into your Pydantic model
- Keep schemas focused — extract only what you need
- Use `Optional` fields for data that may not always be available
## Filter by Tags
Like [recall](./recall#filter-by-tags), reflect supports tag filtering to scope which memories are considered during reasoning. This is essential for multi-user scenarios where reflection should only consider memories relevant to a specific user.
<Tabs>
<TabItem value="python" label="Python">
<CodeSnippet code={reflectPy} section="reflect-with-tags" language="python" />
</TabItem>
</Tabs>
The `tags_match` parameter works the same as in recall:
| Mode | Behavior |
|------|----------|
| `any` | 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 |
See [Retain API](./retain#tagging-memories) for how to tag memories and [Recall API](./recall#filter-by-tags) for more details on tag matching modes.
@@ -129,3 +129,55 @@ For large batches, use async ingestion to avoid blocking:
<CodeSnippet code={retainMjs} section="retain-async" language="javascript" />
</TabItem>
</Tabs>
## Tagging Memories
Tags enable **visibility scoping**—useful when one memory bank serves multiple users but each should only see relevant memories. For example, an agent that chats with multiple users can tag memories by user ID and filter during recall.
### Tag Individual Items
<Tabs>
<TabItem value="python" label="Python">
<CodeSnippet code={retainPy} section="retain-with-tags" language="python" />
</TabItem>
</Tabs>
### Apply Tags to All Items in a Batch
Use `document_tags` to apply the same tags to all items in a request:
<Tabs>
<TabItem value="python" label="Python">
<CodeSnippet code={retainPy} section="retain-with-document-tags" language="python" />
</TabItem>
</Tabs>
When both `document_tags` and item-level `tags` are provided, they are merged together.
### Tag Naming Conventions
Use consistent naming patterns for tags:
| Pattern | Example | Use Case |
|---------|---------|----------|
| `user:<id>` | `user:alice` | Multi-user agent filtering |
| `session:<id>` | `session:123` | Session-based scoping |
| `room:<id>` | `room:general` | Chat room isolation |
| `topic:<name>` | `topic:feedback` | Topic categorization |
### Listing Tags
Use the list tags API to discover existing tags, useful for UI autocomplete or wildcard expansion:
```python
# List all tags in a bank
tags = client.list_tags(bank_id="my-bank")
for tag in tags.items:
print(f"{tag.tag}: {tag.count} memories")
# Search with wildcards (* matches any characters)
user_tags = client.list_tags(bank_id="my-bank", q="user:*")
admin_tags = client.list_tags(bank_id="my-bank", q="*-admin")
```
See [Recall API](./recall#filter-by-tags) for filtering memories by tags during retrieval.
+46 -2
View File
@@ -139,13 +139,18 @@ export HINDSIGHT_API_REFLECT_LLM_MODEL=llama-3.3-70b-versatile
| Variable | Description | Default |
|----------|-------------|---------|
| `HINDSIGHT_API_EMBEDDINGS_PROVIDER` | Provider: `local`, `tei`, `openai`, or `cohere` | `local` |
| `HINDSIGHT_API_EMBEDDINGS_PROVIDER` | Provider: `local`, `tei`, `openai`, `cohere`, or `litellm` | `local` |
| `HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL` | Model for local provider | `BAAI/bge-small-en-v1.5` |
| `HINDSIGHT_API_EMBEDDINGS_TEI_URL` | TEI server URL | - |
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY` | OpenAI API key (falls back to `HINDSIGHT_API_LLM_API_KEY`) | - |
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL` | OpenAI embedding model | `text-embedding-3-small` |
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL` | Custom base URL for OpenAI-compatible API (e.g., Azure OpenAI) | - |
| `HINDSIGHT_API_COHERE_API_KEY` | Cohere API key (shared for embeddings and reranker) | - |
| `HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL` | Cohere embedding model | `embed-english-v3.0` |
| `HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL` | Custom base URL for Cohere-compatible API (e.g., Azure-hosted) | - |
| `HINDSIGHT_API_LITELLM_API_BASE` | LiteLLM proxy base URL (shared for embeddings and reranker) | `http://localhost:4000` |
| `HINDSIGHT_API_LITELLM_API_KEY` | LiteLLM proxy API key (optional, depends on proxy config) | - |
| `HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL` | LiteLLM embedding model (use provider prefix, e.g., `cohere/embed-english-v3.0`) | `text-embedding-3-small` |
```bash
# Local (default) - uses SentenceTransformers
@@ -157,6 +162,12 @@ export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=sk-xxxxxxxxxxxx # or reuses HINDSIGHT_API_LLM_API_KEY
export HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL=text-embedding-3-small # 1536 dimensions
# Azure OpenAI - embeddings via Azure endpoint
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=your-azure-api-key
export HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL=text-embedding-3-small
export HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL=https://your-resource.openai.azure.com/openai/deployments/your-deployment
# TEI - HuggingFace Text Embeddings Inference (recommended for production)
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=tei
export HINDSIGHT_API_EMBEDDINGS_TEI_URL=http://localhost:8080
@@ -165,6 +176,18 @@ export HINDSIGHT_API_EMBEDDINGS_TEI_URL=http://localhost:8080
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=cohere
export HINDSIGHT_API_COHERE_API_KEY=your-api-key
export HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL=embed-english-v3.0 # 1024 dimensions
# Azure-hosted Cohere - embeddings via custom endpoint
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=cohere
export HINDSIGHT_API_COHERE_API_KEY=your-azure-api-key
export HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL=embed-english-v3.0
export HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL=https://your-azure-cohere-endpoint.com
# LiteLLM proxy - unified gateway for multiple providers
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=litellm
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
export HINDSIGHT_API_LITELLM_API_KEY=your-litellm-key # optional
export HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL=text-embedding-3-small # or cohere/embed-english-v3.0
```
#### Embedding Dimensions
@@ -187,13 +210,15 @@ Supported OpenAI embedding dimensions:
| Variable | Description | Default |
|----------|-------------|---------|
| `HINDSIGHT_API_RERANKER_PROVIDER` | Provider: `local`, `tei`, or `cohere` | `local` |
| `HINDSIGHT_API_RERANKER_PROVIDER` | Provider: `local`, `tei`, `cohere`, `flashrank`, `litellm`, or `rrf` | `local` |
| `HINDSIGHT_API_RERANKER_LOCAL_MODEL` | Model for local provider | `cross-encoder/ms-marco-MiniLM-L-6-v2` |
| `HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT` | Max concurrent local reranking (prevents CPU thrashing under load) | `4` |
| `HINDSIGHT_API_RERANKER_TEI_URL` | TEI server URL | - |
| `HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE` | Batch size for TEI reranking | `128` |
| `HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT` | Max concurrent TEI reranking requests | `8` |
| `HINDSIGHT_API_RERANKER_COHERE_MODEL` | Cohere rerank model | `rerank-english-v3.0` |
| `HINDSIGHT_API_RERANKER_COHERE_BASE_URL` | Custom base URL for Cohere-compatible API (e.g., Azure-hosted) | - |
| `HINDSIGHT_API_RERANKER_LITELLM_MODEL` | LiteLLM rerank model (use provider prefix, e.g., `cohere/rerank-english-v3.0`) | `cohere/rerank-english-v3.0` |
```bash
# Local (default) - uses SentenceTransformers CrossEncoder
@@ -208,8 +233,27 @@ export HINDSIGHT_API_RERANKER_TEI_URL=http://localhost:8081
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
export HINDSIGHT_API_COHERE_API_KEY=your-api-key # shared with embeddings
export HINDSIGHT_API_RERANKER_COHERE_MODEL=rerank-english-v3.0
# Azure-hosted Cohere - reranking via custom endpoint
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
export HINDSIGHT_API_COHERE_API_KEY=your-azure-api-key
export HINDSIGHT_API_RERANKER_COHERE_MODEL=rerank-english-v3.0
export HINDSIGHT_API_RERANKER_COHERE_BASE_URL=https://your-azure-cohere-endpoint.com
# LiteLLM proxy - unified gateway for multiple reranking providers
export HINDSIGHT_API_RERANKER_PROVIDER=litellm
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
export HINDSIGHT_API_LITELLM_API_KEY=your-litellm-key # optional
export HINDSIGHT_API_RERANKER_LITELLM_MODEL=cohere/rerank-english-v3.0 # or voyage/rerank-2, together_ai/...
```
LiteLLM supports multiple reranking providers via the `/rerank` endpoint:
- Cohere (`cohere/rerank-english-v3.0`, `cohere/rerank-multilingual-v3.0`)
- Together AI (`together_ai/...`)
- Voyage AI (`voyage/rerank-2`)
- Jina AI (`jina_ai/...`)
- AWS Bedrock (`bedrock/...`)
### Authentication
By default, Hindsight runs without authentication. For production deployments, enable API key authentication using the built-in tenant extension:
+102 -12
View File
@@ -95,29 +95,71 @@ Converts text into dense vector representations for semantic similarity search.
**Default:** `BAAI/bge-small-en-v1.5` (384 dimensions, ~130MB)
**Alternatives:**
### Supported Providers
| Model | Use Case |
|-------|----------|
| `BAAI/bge-small-en-v1.5` | Default, fast, good quality |
| `sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2` | Multilingual (50+ languages) |
| Provider | Description | Best For |
|----------|-------------|----------|
| `local` | SentenceTransformers (default) | Development, low latency |
| `openai` | OpenAI embeddings API | Production, high quality |
| `cohere` | Cohere embeddings API | Production, multilingual |
| `tei` | HuggingFace Text Embeddings Inference | Production, self-hosted |
| `litellm` | LiteLLM proxy (unified gateway) | Multi-provider setups |
:::warning
All embedding models must produce **384-dimensional vectors** to match the database schema.
### Local Models
| Model | Dimensions | Use Case |
|-------|------------|----------|
| `BAAI/bge-small-en-v1.5` | 384 | Default, fast, good quality |
| `sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2` | 384 | Multilingual (50+ languages) |
### OpenAI Models
| Model | Dimensions | Use Case |
|-------|------------|----------|
| `text-embedding-3-small` | 1536 | Default OpenAI, cost-effective |
| `text-embedding-3-large` | 3072 | Higher quality, more expensive |
| `text-embedding-ada-002` | 1536 | Legacy model |
### Cohere Models
| Model | Dimensions | Use Case |
|-------|------------|----------|
| `embed-english-v3.0` | 1024 | English text |
| `embed-multilingual-v3.0` | 1024 | 100+ languages |
:::warning Embedding Dimensions
Hindsight automatically detects the embedding dimension at startup and adjusts the database schema. Once memories are stored, you cannot change dimensions without losing data.
:::
**Configuration:**
**Configuration Examples:**
```bash
# Local provider (default)
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=local
export HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL=BAAI/bge-small-en-v1.5
# TEI provider (remote)
# OpenAI
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=sk-xxxxxxxxxxxx
export HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL=text-embedding-3-small
# Cohere
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=cohere
export HINDSIGHT_API_COHERE_API_KEY=your-api-key
export HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL=embed-english-v3.0
# TEI (self-hosted)
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=tei
export HINDSIGHT_API_EMBEDDINGS_TEI_URL=http://localhost:8080
# LiteLLM proxy
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=litellm
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
export HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL=text-embedding-3-small
```
See [Configuration](./configuration#embeddings) for all options including Azure OpenAI and custom endpoints.
---
## Cross-Encoder (Reranker)
@@ -126,7 +168,18 @@ Reranks initial search results to improve precision.
**Default:** `cross-encoder/ms-marco-MiniLM-L-6-v2` (~85MB)
**Alternatives:**
### Supported Providers
| Provider | Description | Best For |
|----------|-------------|----------|
| `local` | SentenceTransformers CrossEncoder (default) | Development, low latency |
| `cohere` | Cohere rerank API | Production, high quality |
| `tei` | HuggingFace Text Embeddings Inference | Production, self-hosted |
| `flashrank` | FlashRank (lightweight, fast) | Resource-constrained environments |
| `litellm` | LiteLLM proxy (unified gateway) | Multi-provider setups |
| `rrf` | RRF-only (no neural reranking) | Testing, minimal resources |
### Local Models
| Model | Use Case |
|-------|----------|
@@ -134,14 +187,51 @@ Reranks initial search results to improve precision.
| `cross-encoder/ms-marco-MiniLM-L-12-v2` | Higher accuracy |
| `cross-encoder/mmarco-mMiniLMv2-L12-H384-v1` | Multilingual |
**Configuration:**
### Cohere Models
| Model | Use Case |
|-------|----------|
| `rerank-english-v3.0` | English text |
| `rerank-multilingual-v3.0` | 100+ languages |
### LiteLLM Supported Providers
LiteLLM supports multiple reranking providers via the `/rerank` endpoint:
| Provider | Model Example |
|----------|---------------|
| Cohere | `cohere/rerank-english-v3.0` |
| Together AI | `together_ai/...` |
| Voyage AI | `voyage/rerank-2` |
| Jina AI | `jina_ai/...` |
| AWS Bedrock | `bedrock/...` |
**Configuration Examples:**
```bash
# Local provider (default)
export HINDSIGHT_API_RERANKER_PROVIDER=local
export HINDSIGHT_API_RERANKER_LOCAL_MODEL=cross-encoder/ms-marco-MiniLM-L-6-v2
# TEI provider (remote)
# Cohere
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
export HINDSIGHT_API_COHERE_API_KEY=your-api-key
export HINDSIGHT_API_RERANKER_COHERE_MODEL=rerank-english-v3.0
# TEI (self-hosted)
export HINDSIGHT_API_RERANKER_PROVIDER=tei
export HINDSIGHT_API_RERANKER_TEI_URL=http://localhost:8081
# FlashRank (lightweight)
export HINDSIGHT_API_RERANKER_PROVIDER=flashrank
# LiteLLM proxy
export HINDSIGHT_API_RERANKER_PROVIDER=litellm
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
export HINDSIGHT_API_RERANKER_LITELLM_MODEL=cohere/rerank-english-v3.0
# RRF-only (no neural reranking)
export HINDSIGHT_API_RERANKER_PROVIDER=rrf
```
See [Configuration](./configuration#reranker) for all options including Azure-hosted endpoints and batch settings.
+1 -1
View File
@@ -183,4 +183,4 @@ Disposition creates **consistent character** across conversations while allowing
- [**Retain**](./retain) — How rich facts are stored
- [**Recall**](./retrieval) — How multi-strategy search works
- [API Reference: Reflect](./api/reflect) — Code examples and usage
- [**Reflect API**](./api/reflect) — Code examples, parameters, and tag filtering
+14 -1
View File
@@ -167,6 +167,18 @@ As facts accumulate about an entity, Hindsight synthesizes **observations** —
---
## Tagging Memories
Tags enable visibility scoping—useful when one memory bank serves multiple users but each should only see relevant memories.
- **Item tags**: Tag individual memories with specific scopes
- **Document tags**: Apply tags to all items in a batch
- **Tag filtering**: Filter during recall/reflect by tags
See [Retain API](./api/retain) for code examples and [Recall API](./api/recall) for filtering options.
---
## What You Get
After `retain()` completes:
@@ -176,6 +188,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()`.
@@ -185,4 +198,4 @@ All stored in your isolated **memory bank**, ready for `recall()` and `reflect()
- [**Recall**](./retrieval) — How multi-strategy search retrieves relevant memories
- [**Reflect**](./reflect) — How disposition influences reasoning and opinion formation
- [API Reference](./api/retain) — Code examples for retaining memories
- [**Retain API**](./api/retain) — Code examples and parameters
+4 -1
View File
@@ -133,7 +133,9 @@ Hindsight is built for AI agents, not humans. Traditional search systems return
**Parameters you control:**
- `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
- `types`: Filter by world, experience, opinion, or all
- `tags`: Filter memories by visibility tags
- `tags_match`: How to match tags (see [Recall API](./api/recall) for all options)
### Expanding Context: Chunks and Entity Observations
@@ -241,3 +243,4 @@ See [Configuration → Retrieval](./configuration#retrieval) for available algor
- [**Retain**](./retain) — How memories are stored with rich context
- [**Reflect**](./reflect) — How disposition influences reasoning
- [**Recall API**](./api/recall) — Code examples, parameters, and tag filtering
+33
View File
@@ -116,6 +116,39 @@ results = client.recall(bank_id="my-bank", query="How are Alice and Bob connecte
# [/docs:recall-budget-levels]
# [docs:recall-with-tags]
# Filter recall to only memories tagged for a specific user
response = client.recall(
bank_id="my-bank",
query="What feedback did the user give?",
tags=["user:alice"],
tags_match="any" # OR matching, includes untagged (default)
)
# [/docs:recall-with-tags]
# [docs:recall-tags-strict]
# Strict mode: only return memories that have matching tags (exclude untagged)
response = client.recall(
bank_id="my-bank",
query="What did the user say?",
tags=["user:alice"],
tags_match="any_strict" # OR matching, excludes untagged memories
)
# [/docs:recall-tags-strict]
# [docs:recall-tags-all]
# AND matching: require ALL specified tags to be present
response = client.recall(
bank_id="my-bank",
query="What bugs were reported?",
tags=["user:alice", "bug-report"],
tags_match="all_strict" # Memory must have BOTH tags
)
# [/docs:recall-tags-all]
# =============================================================================
# Cleanup (not shown in docs)
# =============================================================================
+11
View File
@@ -81,6 +81,17 @@ for fact in response.based_on or []:
# [/docs:reflect-sources]
# [docs:reflect-with-tags]
# Filter reflection to only consider memories for a specific user
response = client.reflect(
bank_id="my-bank",
query="What does this user think about our product?",
tags=["user:alice"],
tags_match="any_strict" # Only use memories tagged for this user
)
# [/docs:reflect-with-tags]
# =============================================================================
# Cleanup (not shown in docs)
# =============================================================================
+33
View File
@@ -67,6 +67,39 @@ print(result.var_async) # True
# [/docs:retain-async]
# [docs:retain-with-tags]
# Tag individual items for visibility scoping
client.retain_batch(
bank_id="my-bank",
items=[
{
"content": "User Alice said she loves the new dashboard",
"tags": ["user:alice", "feedback"]
},
{
"content": "User Bob reported a bug in the search feature",
"tags": ["user:bob", "bug-report"]
}
],
document_id="user_feedback_001"
)
# [/docs:retain-with-tags]
# [docs:retain-with-document-tags]
# Apply tags to all items in a batch
client.retain_batch(
bank_id="my-bank",
items=[
{"content": "Alice mentioned she prefers dark mode"},
{"content": "Bob asked about keyboard shortcuts"}
],
document_id="support_session_123",
document_tags=["session:123", "support"] # Applied to all items
)
# [/docs:retain-with-document-tags]
# =============================================================================
# Cleanup (not shown in docs)
# =============================================================================
+370 -1
View File
@@ -249,6 +249,72 @@
}
}
},
"/v1/default/banks/{bank_id}/memories/{memory_id}": {
"get": {
"tags": [
"Memory"
],
"summary": "Get memory unit",
"description": "Get a single memory unit by ID with all its metadata including entities and tags.",
"operationId": "get_memory",
"parameters": [
{
"name": "bank_id",
"in": "path",
"required": true,
"schema": {
"type": "string",
"title": "Bank Id"
}
},
{
"name": "memory_id",
"in": "path",
"required": true,
"schema": {
"type": "string",
"title": "Memory Id"
}
},
{
"name": "authorization",
"in": "header",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Authorization"
}
}
],
"responses": {
"200": {
"description": "Successful Response",
"content": {
"application/json": {
"schema": {}
}
}
},
"422": {
"description": "Validation Error",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
}
}
}
}
},
"/v1/default/banks/{bank_id}/memories/recall": {
"post": {
"tags": [
@@ -944,6 +1010,107 @@
}
}
},
"/v1/default/banks/{bank_id}/tags": {
"get": {
"tags": [
"Memory"
],
"summary": "List tags",
"description": "List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.",
"operationId": "list_tags",
"parameters": [
{
"name": "bank_id",
"in": "path",
"required": true,
"schema": {
"type": "string",
"title": "Bank Id"
}
},
{
"name": "q",
"in": "query",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.",
"title": "Q"
},
"description": "Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive."
},
{
"name": "limit",
"in": "query",
"required": false,
"schema": {
"type": "integer",
"description": "Maximum number of tags to return",
"default": 100,
"title": "Limit"
},
"description": "Maximum number of tags to return"
},
{
"name": "offset",
"in": "query",
"required": false,
"schema": {
"type": "integer",
"description": "Offset for pagination",
"default": 0,
"title": "Offset"
},
"description": "Offset for pagination"
},
{
"name": "authorization",
"in": "header",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Authorization"
}
}
],
"responses": {
"200": {
"description": "Successful Response",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ListTagsResponse"
}
}
}
},
"422": {
"description": "Validation Error",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
}
}
}
}
},
"/v1/default/chunks/{chunk_id}": {
"get": {
"tags": [
@@ -2219,6 +2386,14 @@
"memory_unit_count": {
"type": "integer",
"title": "Memory Unit Count"
},
"tags": {
"items": {
"type": "string"
},
"type": "array",
"title": "Tags",
"description": "Tags associated with this document"
}
},
"type": "object",
@@ -2240,6 +2415,10 @@
"id": "session_1",
"memory_unit_count": 15,
"original_text": "Full document text here...",
"tags": [
"user_a",
"session_123"
],
"updated_at": "2024-01-15T10:30:00Z"
}
},
@@ -2752,6 +2931,57 @@
"total": 150
}
},
"ListTagsResponse": {
"properties": {
"items": {
"items": {
"$ref": "#/components/schemas/TagItem"
},
"type": "array",
"title": "Items"
},
"total": {
"type": "integer",
"title": "Total"
},
"limit": {
"type": "integer",
"title": "Limit"
},
"offset": {
"type": "integer",
"title": "Offset"
}
},
"type": "object",
"required": [
"items",
"total",
"limit",
"offset"
],
"title": "ListTagsResponse",
"description": "Response model for list tags endpoint.",
"example": {
"items": [
{
"count": 42,
"tag": "user:alice"
},
{
"count": 15,
"tag": "user:bob"
},
{
"count": 8,
"tag": "session:abc123"
}
],
"limit": 100,
"offset": 0,
"total": 25
}
},
"MemoryItem": {
"properties": {
"content": {
@@ -2821,6 +3051,21 @@
],
"title": "Entities",
"description": "Optional entities to combine with auto-extracted entities."
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags",
"description": "Optional tags for visibility scoping. Memories with tags can be filtered during recall."
}
},
"type": "object",
@@ -2846,6 +3091,10 @@
"channel": "engineering",
"source": "slack"
},
"tags": [
"user_a",
"user_b"
],
"timestamp": "2024-01-15T10:30:00Z"
}
},
@@ -2999,6 +3248,33 @@
"include": {
"$ref": "#/components/schemas/IncludeOptions",
"description": "Options for including additional data (entities are included by default)"
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags",
"description": "Filter memories by tags. If not specified, all memories are returned."
},
"tags_match": {
"type": "string",
"enum": [
"any",
"all",
"any_strict",
"all_strict"
],
"title": "Tags Match",
"description": "How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
"default": "any"
}
},
"type": "object",
@@ -3017,6 +3293,10 @@
"max_tokens": 4096,
"query": "What did Alice say about machine learning?",
"query_timestamp": "2023-05-30T23:40:00",
"tags": [
"user_a"
],
"tags_match": "any",
"trace": true,
"types": [
"world",
@@ -3238,6 +3518,20 @@
}
],
"title": "Chunk Id"
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags"
}
},
"type": "object",
@@ -3262,6 +3556,10 @@
},
"occurred_end": "2024-01-15T10:30:00Z",
"occurred_start": "2024-01-15T10:30:00Z",
"tags": [
"user_a",
"user_b"
],
"text": "Alice works at Google on the AI team",
"type": "world"
}
@@ -3404,6 +3702,33 @@
],
"title": "Response Schema",
"description": "Optional JSON Schema for structured output. When provided, the response will include a 'structured_output' field with the LLM response parsed according to this schema."
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags",
"description": "Filter memories by tags during reflection. If not specified, all memories are considered."
},
"tags_match": {
"type": "string",
"enum": [
"any",
"all",
"any_strict",
"all_strict"
],
"title": "Tags Match",
"description": "How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
"default": "any"
}
},
"type": "object",
@@ -3437,7 +3762,11 @@
"key_points"
],
"type": "object"
}
},
"tags": [
"user_a"
],
"tags_match": "any"
}
},
"ReflectResponse": {
@@ -3527,6 +3856,21 @@
"title": "Async",
"description": "If true, process asynchronously in background. If false, wait for completion (default: false)",
"default": false
},
"document_tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Document Tags",
"description": "Tags applied to all items in this request. These are merged with any item-level tags."
}
},
"type": "object",
@@ -3537,6 +3881,10 @@
"description": "Request model for retain endpoint.",
"example": {
"async": false,
"document_tags": [
"user_a",
"user_b"
],
"items": [
{
"content": "Alice works at Google",
@@ -3615,6 +3963,27 @@
}
}
},
"TagItem": {
"properties": {
"tag": {
"type": "string",
"title": "Tag",
"description": "The tag value"
},
"count": {
"type": "integer",
"title": "Count",
"description": "Number of memories with this tag"
}
},
"type": "object",
"required": [
"tag",
"count"
],
"title": "TagItem",
"description": "Single tag with usage count."
},
"TokenUsage": {
"properties": {
"input_tokens": {
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "hindsight-embed"
version = "0.2.1"
version = "0.3.0"
description = "Hindsight embedded CLI - local memory operations without a server"
readme = "README.md"
requires-python = ">=3.11"
@@ -1,6 +1,6 @@
[project]
name = "hindsight-litellm"
version = "0.2.1"
version = "0.3.0"
description = "Universal LLM memory integration via LiteLLM - works with 100+ providers"
readme = "README.md"
requires-python = ">=3.10"
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "hindsight-all"
version = "0.2.1"
version = "0.3.0"
description = "Hindsight: Agent Memory That Works Like Human Memory - All-in-One Bundle"
readme = "README.md"
requires-python = ">=3.11"
@@ -54,7 +54,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(hindsight_http_requests_total)",
"expr": "sum(hindsight_http_requests_total{tenant=~\"$tenant\"})",
"refId": "A"
}
],
@@ -99,7 +99,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(rate(hindsight_http_requests_total[1m]))",
"expr": "sum(rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[1m]))",
"refId": "A"
}
],
@@ -191,7 +191,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"5xx\"}[5m])) / sum(rate(hindsight_http_requests_total[5m]))",
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"5xx\", tenant=~\"$tenant\"}[5m])) / sum(rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[5m]))",
"refId": "A"
}
],
@@ -236,7 +236,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))",
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
"refId": "A"
}
],
@@ -313,7 +313,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum by (endpoint) (rate(hindsight_http_requests_total[1m]))",
"expr": "sum by (endpoint) (rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "{{endpoint}}",
"refId": "A"
}
@@ -404,17 +404,17 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "histogram_quantile(0.50, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))",
"expr": "histogram_quantile(0.50, sum by (le) (rate(hindsight_http_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
"legendFormat": "p50",
"refId": "A"
},
{
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))",
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
"legendFormat": "p95",
"refId": "B"
},
{
"expr": "histogram_quantile(0.99, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))",
"expr": "histogram_quantile(0.99, sum by (le) (rate(hindsight_http_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
"legendFormat": "p99",
"refId": "C"
}
@@ -505,12 +505,12 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"5xx\"}[1m])) / sum(rate(hindsight_http_requests_total[1m]))",
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"5xx\", tenant=~\"$tenant\"}[1m])) / sum(rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "5xx Error Rate",
"refId": "A"
},
{
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"4xx\"}[1m])) / sum(rate(hindsight_http_requests_total[1m]))",
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"4xx\", tenant=~\"$tenant\"}[1m])) / sum(rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "4xx Error Rate",
"refId": "B"
}
@@ -1276,7 +1276,36 @@
"schemaVersion": 38,
"tags": ["hindsight", "api", "service"],
"templating": {
"list": []
"list": [
{
"allValue": ".*",
"current": {
"selected": true,
"text": "All",
"value": "$__all"
},
"datasource": {
"type": "prometheus",
"uid": "prometheus"
},
"definition": "label_values(hindsight_http_requests_total, tenant)",
"hide": 0,
"includeAll": true,
"label": "Tenant",
"multi": false,
"name": "tenant",
"options": [],
"query": {
"query": "label_values(hindsight_http_requests_total, tenant)",
"refId": "PrometheusVariableQueryEditor-VariableQuery"
},
"refresh": 2,
"regex": "",
"skipUrlSync": false,
"sort": 1,
"type": "query"
}
]
},
"time": {
"from": "now-30m",
@@ -46,7 +46,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(hindsight_llm_calls_total)",
"expr": "sum(hindsight_llm_calls_total{tenant=~\"$tenant\"})",
"refId": "A"
}
],
@@ -91,7 +91,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(hindsight_llm_tokens_input_tokens_total) + sum(hindsight_llm_tokens_output_tokens_total)",
"expr": "sum(hindsight_llm_tokens_input_tokens_total{tenant=~\"$tenant\"}) + sum(hindsight_llm_tokens_output_tokens_total{tenant=~\"$tenant\"})",
"refId": "A"
}
],
@@ -137,7 +137,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(hindsight_llm_tokens_input_tokens_total)",
"expr": "sum(hindsight_llm_tokens_input_tokens_total{tenant=~\"$tenant\"})",
"refId": "A"
}
],
@@ -183,7 +183,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(hindsight_llm_tokens_output_tokens_total)",
"expr": "sum(hindsight_llm_tokens_output_tokens_total{tenant=~\"$tenant\"})",
"refId": "A"
}
],
@@ -260,7 +260,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum by (scope) (rate(hindsight_llm_calls_total[1m]))",
"expr": "sum by (scope) (rate(hindsight_llm_calls_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "{{scope}}",
"refId": "A"
}
@@ -347,12 +347,12 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(rate(hindsight_llm_tokens_input_tokens_total[1m]))",
"expr": "sum(rate(hindsight_llm_tokens_input_tokens_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "Input",
"refId": "A"
},
{
"expr": "sum(rate(hindsight_llm_tokens_output_tokens_total[1m]))",
"expr": "sum(rate(hindsight_llm_tokens_output_tokens_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "Output",
"refId": "B"
}
@@ -430,7 +430,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "histogram_quantile(0.95, sum by (scope, le) (rate(hindsight_llm_duration_seconds_bucket[5m])))",
"expr": "histogram_quantile(0.95, sum by (scope, le) (rate(hindsight_llm_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
"legendFormat": "{{scope}}",
"refId": "A"
}
@@ -508,12 +508,12 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum by (scope) (rate(hindsight_llm_tokens_input_tokens_total[1m]))",
"expr": "sum by (scope) (rate(hindsight_llm_tokens_input_tokens_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "{{scope}} (input)",
"refId": "A"
},
{
"expr": "sum by (scope) (rate(hindsight_llm_tokens_output_tokens_total[1m]))",
"expr": "sum by (scope) (rate(hindsight_llm_tokens_output_tokens_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "{{scope}} (output)",
"refId": "B"
}
@@ -526,7 +526,36 @@
"schemaVersion": 38,
"tags": ["hindsight", "llm"],
"templating": {
"list": []
"list": [
{
"allValue": ".*",
"current": {
"selected": true,
"text": "All",
"value": "$__all"
},
"datasource": {
"type": "prometheus",
"uid": "prometheus"
},
"definition": "label_values(hindsight_llm_calls_total, tenant)",
"hide": 0,
"includeAll": true,
"label": "Tenant",
"multi": false,
"name": "tenant",
"options": [],
"query": {
"query": "label_values(hindsight_llm_calls_total, tenant)",
"refId": "PrometheusVariableQueryEditor-VariableQuery"
},
"refresh": 2,
"regex": "",
"skipUrlSync": false,
"sort": 1,
"type": "query"
}
]
},
"time": {
"from": "now-30m",
@@ -46,7 +46,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(hindsight_operation_operations_total)",
"expr": "sum(hindsight_operation_operations_total{tenant=~\"$tenant\"})",
"refId": "A"
}
],
@@ -91,7 +91,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(rate(hindsight_operation_operations_total[1m]))",
"expr": "sum(rate(hindsight_operation_operations_total{tenant=~\"$tenant\"}[1m]))",
"refId": "A"
}
],
@@ -137,7 +137,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"retain\"}[1m]))",
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"retain\", tenant=~\"$tenant\"}[1m]))",
"refId": "A"
}
],
@@ -183,7 +183,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"recall\"}[1m]))",
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"recall\", tenant=~\"$tenant\"}[1m]))",
"refId": "A"
}
],
@@ -229,7 +229,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"reflect\"}[1m]))",
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"reflect\", tenant=~\"$tenant\"}[1m]))",
"refId": "A"
}
],
@@ -319,7 +319,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum by (operation) (rate(hindsight_operation_operations_total[1m]))",
"expr": "sum by (operation) (rate(hindsight_operation_operations_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "{{operation}}",
"refId": "A"
}
@@ -410,17 +410,17 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "histogram_quantile(0.50, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\"}[5m])))",
"expr": "histogram_quantile(0.50, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\", tenant=~\"$tenant\"}[5m])))",
"legendFormat": "p50",
"refId": "A"
},
{
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\"}[5m])))",
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\", tenant=~\"$tenant\"}[5m])))",
"legendFormat": "p95",
"refId": "B"
},
{
"expr": "histogram_quantile(0.99, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\"}[5m])))",
"expr": "histogram_quantile(0.99, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\", tenant=~\"$tenant\"}[5m])))",
"legendFormat": "p99",
"refId": "C"
}
@@ -498,7 +498,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "histogram_quantile(0.95, sum by (operation, le) (rate(hindsight_operation_duration_seconds_bucket[5m])))",
"expr": "histogram_quantile(0.95, sum by (operation, le) (rate(hindsight_operation_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
"legendFormat": "{{operation}}",
"refId": "A"
}
@@ -576,7 +576,7 @@
"pluginVersion": "10.0.0",
"targets": [
{
"expr": "sum by (bank_id) (rate(hindsight_operation_operations_total[1m]))",
"expr": "sum by (bank_id) (rate(hindsight_operation_operations_total{tenant=~\"$tenant\"}[1m]))",
"legendFormat": "{{bank_id}}",
"refId": "A"
}
@@ -589,7 +589,36 @@
"schemaVersion": 38,
"tags": ["hindsight"],
"templating": {
"list": []
"list": [
{
"allValue": ".*",
"current": {
"selected": true,
"text": "All",
"value": "$__all"
},
"datasource": {
"type": "prometheus",
"uid": "prometheus"
},
"definition": "label_values(hindsight_operation_operations_total, tenant)",
"hide": 0,
"includeAll": true,
"label": "Tenant",
"multi": false,
"name": "tenant",
"options": [],
"query": {
"query": "label_values(hindsight_operation_operations_total, tenant)",
"refId": "PrometheusVariableQueryEditor-VariableQuery"
},
"refresh": 2,
"regex": "",
"skipUrlSync": false,
"sort": 1,
"type": "query"
}
]
},
"time": {
"from": "now-30m",
Generated
+5 -5
View File
@@ -1257,7 +1257,7 @@ wheels = [
[[package]]
name = "hindsight-all"
version = "0.2.1"
version = "0.3.0"
source = { editable = "hindsight" }
dependencies = [
{ name = "hindsight-api" },
@@ -1281,7 +1281,7 @@ provides-extras = ["test"]
[[package]]
name = "hindsight-api"
version = "0.2.1"
version = "0.3.0"
source = { editable = "hindsight-api" }
dependencies = [
{ name = "alembic" },
@@ -1397,7 +1397,7 @@ dev = [
[[package]]
name = "hindsight-client"
version = "0.2.1"
version = "0.3.0"
source = { editable = "hindsight-clients/python" }
dependencies = [
{ name = "aiohttp" },
@@ -1431,7 +1431,7 @@ provides-extras = ["test"]
[[package]]
name = "hindsight-dev"
version = "0.2.1"
version = "0.3.0"
source = { editable = "hindsight-dev" }
dependencies = [
{ name = "hindsight-api" },
@@ -1466,7 +1466,7 @@ dev = [
[[package]]
name = "hindsight-embed"
version = "0.2.1"
version = "0.3.0"
source = { editable = "hindsight-embed" }
dependencies = [
{ name = "httpx" },