Compare commits

...
Author SHA1 Message Date
Chris Bartholomew 3b343bc6aa Rename GCPJsonFormatter to JsonFormatter 2026-01-18 11:59:49 -05:00
Chris Bartholomew 1907a33edd Add structured JSON logging support
Add HINDSIGHT_API_LOG_FORMAT environment variable to configure log output
format. Options are "text" (default, human-readable) and "json" (structured).

JSON format outputs logs with a "severity" field that cloud logging systems
can parse for proper log level categorization. Also writes to stdout instead
of stderr so log levels are correctly interpreted.
2026-01-18 11:18:17 -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
Nicolò Boschi 7935b0accd fix: improve mpfp retrieval (#146)
* fix: improve mpfp retrieval

* fix: improve mpfp retrieval

* fix: improve embeddings service performances

* fix: improve embeddings service performances

* fix: improve embeddings service performances

* fix: improve embeddings service performances
2026-01-12 18:58:05 +01:00
Nicolò Boschi 26bf5714cd fix: entities list only show 100 entities (#142)
* fix: entities list only show 100 entities

* fix: update Rust CLI for entities pagination API changes
2026-01-12 18:50:53 +01:00
Nicolò Boschi 6232e690fc fix: improve graph retrieval on large memory banks (#141) 2026-01-09 16:43:31 +01:00
Nicolò Boschi 4135a6cee5 ci: frozen uv sync (#138)
* ci: frozen uv sync

* fix: add missing authorization parameter to get_agent_stats in CLI

The generated Rust client was updated with an authorization header
parameter for get_agent_stats, but the CLI code wasn't updated.
2026-01-09 16:43:00 +01:00
Nicolò Boschi eb2702bcba misc: performance improvements (#140)
* misc: performance improvements

* misc: performance improvements

* misc: performance improvements
2026-01-09 14:47:20 +01:00
Nicolò Boschi 0d0abaaa9f fix(typescript-client): Add error handling to all API methods (#139)
Previously, most methods in HindsightClient would silently return
undefined when API calls failed (e.g., connection refused). Only
the `recall` method had proper error checking.

This change adds a `validateResponse` helper method and applies it
consistently to all API methods:
- retain
- retainBatch
- recall
- reflect
- listMemories
- createBank
- getBankProfile

Now all methods properly throw an error with details when the API
request fails, instead of returning undefined.
2026-01-09 14:25:10 +01:00
Nicolò Boschi a6798f7e2a fix: improve tei client parameters (#137)
* fix: improve tei client parameters

* fix: improve tei client parameters

* fix: improve tei client parameters
2026-01-09 11:31:22 +01:00
Nicolò Boschi fb31a35a86 feat: retain modes (#136)
* feat: retain modes

* fix db patch
2026-01-09 11:30:36 +01:00
109 changed files with 13273 additions and 1569 deletions
+15 -15
View File
@@ -222,7 +222,7 @@ jobs:
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync --no-install-project --index-strategy unsafe-best-match
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Create .env file
run: |
@@ -352,7 +352,7 @@ jobs:
- name: Install dependencies
working-directory: ./hindsight-api
run: uv sync --extra test --no-install-project --index-strategy unsafe-best-match
run: uv sync --frozen --extra test --no-install-project --index-strategy unsafe-best-match
- name: Cache HuggingFace models
uses: actions/cache@v4
@@ -413,11 +413,11 @@ jobs:
- name: Install client test dependencies
working-directory: ./hindsight-clients/python
run: uv sync --extra test --index-strategy unsafe-best-match
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync --no-install-project --index-strategy unsafe-best-match
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Create .env file
run: |
@@ -490,7 +490,7 @@ jobs:
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync --no-install-project --index-strategy unsafe-best-match
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Install TypeScript client dependencies
working-directory: ./hindsight-clients/typescript
@@ -578,7 +578,7 @@ jobs:
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync --no-install-project --index-strategy unsafe-best-match
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Create .env file
run: |
@@ -645,11 +645,11 @@ jobs:
- name: Install API dependencies
working-directory: ./hindsight-api
run: uv sync --no-install-project --index-strategy unsafe-best-match
run: uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Install integration test dependencies
working-directory: ./hindsight-integration-tests
run: uv sync
run: uv sync --frozen
- name: Cache HuggingFace models
uses: actions/cache@v4
@@ -729,7 +729,7 @@ jobs:
- name: Install dependencies
working-directory: ./hindsight-integrations/litellm
run: uv sync --extra dev
run: uv sync --frozen --extra dev
- name: Run tests
working-directory: ./hindsight-integrations/litellm
@@ -760,7 +760,7 @@ jobs:
- name: Install dependencies
working-directory: ./hindsight-embed
run: uv sync --index-strategy unsafe-best-match
run: uv sync --frozen --index-strategy unsafe-best-match
- name: Cache HuggingFace models
uses: actions/cache@v4
@@ -820,11 +820,11 @@ jobs:
working-directory: ./hindsight-api
run: |
uv build
uv sync --no-install-project --index-strategy unsafe-best-match
uv sync --frozen --no-install-project --index-strategy unsafe-best-match
- name: Install Python client dependencies
working-directory: ./hindsight-clients/python
run: uv sync --extra test --index-strategy unsafe-best-match
run: uv sync --frozen --extra test --index-strategy unsafe-best-match
- name: Install TypeScript client
run: |
@@ -928,9 +928,9 @@ jobs:
- name: Install Python dependencies
run: |
cd hindsight-dev && uv sync --index-strategy unsafe-best-match
cd ../hindsight-api && uv sync --index-strategy unsafe-best-match
cd ../hindsight-embed && uv sync --index-strategy unsafe-best-match
cd hindsight-dev && uv sync --frozen --index-strategy unsafe-best-match
cd ../hindsight-api && uv sync --frozen --index-strategy unsafe-best-match
cd ../hindsight-embed && uv sync --frozen --index-strategy unsafe-best-match
- name: Run generate-openapi
run: ./scripts/generate-openapi.sh
+4
View File
@@ -27,6 +27,10 @@ docker-compose.override.yml
# NLTK data (will be downloaded automatically)
nltk_data/
# Monitoring stack (Prometheus/Grafana binaries and data)
.monitoring/
.pgbouncer
# Large benchmark datasets (will be downloaded automatically)
**/longmemeval_s_cleaned.json
+65
View File
@@ -108,6 +108,52 @@ PostgreSQL with pgvector. Schema managed via Alembic migrations in `hindsight-ap
Key tables: `banks`, `memory_units`, `documents`, `entities`, `entity_links`
### Adding Database Migrations
1. **Create a new migration file** in `hindsight-api/hindsight_api/alembic/versions/`:
- File name format: `<revision_id>_<description>.py` (e.g., `f1a2b3c4d5e6_add_new_index.py`)
- Use a unique hex revision ID (12 chars)
- Set `down_revision` to the previous migration's revision ID
2. **Migration template**:
```python
"""Description of the migration
Revision ID: f1a2b3c4d5e6
Revises: <previous_revision_id>
Create Date: YYYY-MM-DD
"""
from collections.abc import Sequence
from alembic import context, op
revision: str = "f1a2b3c4d5e6"
down_revision: str | Sequence[str] | None = "<previous_revision_id>"
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 (required for multi-tenant support)."""
schema = context.config.get_main_option("target_schema")
return f'"{schema}".' if schema else ""
def upgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"CREATE INDEX ... ON {schema}table_name(...)")
def downgrade() -> None:
schema = _get_schema_prefix()
op.execute(f"DROP INDEX IF EXISTS {schema}index_name")
```
3. **Run migrations locally**:
```bash
# Set database URL and run migrations
uv run hindsight-admin run-db-migration
# Run on a specific tenant schema
uv run hindsight-admin run-db-migration --schema tenant_xyz
```
## Key Conventions
### Code Quality
@@ -128,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
@@ -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,44 @@
"""add_memory_links_from_type_weight_index
Revision ID: f1a2b3c4d5e6
Revises: e0a1b2c3d4e5
Create Date: 2025-01-12
Add composite index on memory_links (from_unit_id, link_type, weight DESC)
to optimize MPFP graph traversal queries that need top-k edges per type.
"""
from collections.abc import Sequence
from alembic import context, op
# revision identifiers, used by Alembic.
revision: str = "f1a2b3c4d5e6"
down_revision: str | Sequence[str] | None = "e0a1b2c3d4e5"
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 composite index for efficient MPFP edge loading."""
schema = _get_schema_prefix()
# Create composite index for efficient top-k per (from_node, link_type) queries
# This enables LATERAL joins to use index-only scans with early termination
# Note: Not using CONCURRENTLY here as it requires running outside a transaction
# For production with large tables, consider running this manually with CONCURRENTLY
op.execute(
f"CREATE INDEX IF NOT EXISTS idx_memory_links_from_type_weight "
f"ON {schema}memory_links(from_unit_id, link_type, weight DESC)"
)
def downgrade() -> None:
"""Remove the composite index."""
schema = _get_schema_prefix()
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_links_from_type_weight")
@@ -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")
+238 -9
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):
@@ -188,12 +202,18 @@ class EntityListResponse(BaseModel):
"first_seen": "2024-01-15T10:30:00Z",
"last_seen": "2024-02-01T14:00:00Z",
}
]
],
"total": 150,
"limit": 100,
"offset": 0,
}
}
)
items: list[EntityListItem]
total: int
limit: int
offset: int
class EntityDetailResponse(BaseModel):
@@ -300,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"],
}
},
)
@@ -313,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
@@ -347,6 +372,7 @@ class RetainRequest(BaseModel):
},
],
"async": False,
"document_tags": ["user_a", "user_b"],
}
}
)
@@ -357,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):
@@ -425,6 +455,8 @@ class ReflectRequest(BaseModel):
},
"required": ["summary", "key_points"],
},
"tags": ["user_a"],
"tags_match": "any",
}
}
)
@@ -440,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):
@@ -722,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."""
@@ -735,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"],
}
}
)
@@ -746,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):
@@ -957,6 +1031,12 @@ def create_app(
await memory.initialize()
logging.info("Memory system initialized")
# Set up DB pool metrics after memory initialization
metrics_collector = get_metrics_collector()
if memory._pool is not None and hasattr(metrics_collector, "set_db_pool"):
metrics_collector.set_db_pool(memory._pool)
logging.info("DB pool metrics configured")
# Call HTTP extension startup hook
if http_extension:
await http_extension.on_startup()
@@ -993,6 +1073,30 @@ def create_app(
# This is required for mounted sub-applications where lifespan may not fire
app.state.memory = memory
# Add HTTP metrics middleware
@app.middleware("http")
async def http_metrics_middleware(request, call_next):
"""Record HTTP request metrics."""
# Normalize endpoint path to reduce cardinality
# Replace UUIDs and numeric IDs with placeholders
import re
from starlette.requests import Request
path = request.url.path
# Replace UUIDs
path = re.sub(r"/[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}", "/{id}", path)
# Replace numeric IDs
path = re.sub(r"/\d+(?=/|$)", "/{id}", path)
status_code = [500] # Default to 500, will be updated
metrics_collector = get_metrics_collector()
with metrics_collector.record_http_request(request.method, path, lambda: status_code[0]):
response = await call_next(request)
status_code[0] = response.status_code
return response
# Register all routes
_register_routes(app)
@@ -1143,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,
@@ -1160,6 +1295,9 @@ def _register_routes(app: FastAPI):
bank_id: str, request: RecallRequest, request_context: RequestContext = Depends(get_request_context)
):
"""Run a recall and return results with trace."""
import time
handler_start = time.time()
metrics = get_metrics_collector()
try:
@@ -1185,10 +1323,12 @@ def _register_routes(app: FastAPI):
include_chunks = request.include.chunks is not None
max_chunk_tokens = request.include.chunks.max_tokens if include_chunks else 8192
pre_recall = time.time() - handler_start
# Run recall with tracing (record metrics)
with metrics.record_operation(
"recall", bank_id=bank_id, source="api", budget=request.budget.value, max_tokens=request.max_tokens
):
recall_start = time.time()
core_result = await app.state.memory.recall_async(
bank_id=bank_id,
query=request.query,
@@ -1202,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)
@@ -1217,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
]
@@ -1247,9 +1390,21 @@ def _register_routes(app: FastAPI):
],
)
return RecallResponse(
response = RecallResponse(
results=recall_results, trace=core_result.trace, entities=entities_response, chunks=chunks_response
)
handler_duration = time.time() - handler_start
recall_duration = time.time() - recall_start
post_recall = handler_duration - pre_recall - recall_duration
if handler_duration > 1.0:
logging.info(
f"[RECALL HTTP] bank={bank_id} handler_total={handler_duration:.3f}s "
f"pre={pre_recall:.3f}s recall={recall_duration:.3f}s post={post_recall:.3f}s "
f"results={len(recall_results)} entities={len(entities_response) if entities_response else 0}"
)
return response
except HTTPException:
raise
except OperationValidationError as e:
@@ -1259,8 +1414,11 @@ def _register_routes(app: FastAPI):
except Exception as e:
import traceback
handler_duration = time.time() - handler_start
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/recall: {error_detail}")
logger.error(
f"[RECALL ERROR] bank={bank_id} handler_duration={handler_duration:.3f}s error={str(e)}\n{error_detail}"
)
raise HTTPException(status_code=500, detail=str(e))
@app.post(
@@ -1294,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
@@ -1486,19 +1646,27 @@ def _register_routes(app: FastAPI):
"/v1/default/banks/{bank_id}/entities",
response_model=EntityListResponse,
summary="List entities",
description="List all entities (people, organizations, etc.) known by the bank, ordered by mention count.",
description="List all entities (people, organizations, etc.) known by the bank, ordered by mention count. Supports pagination.",
operation_id="list_entities",
tags=["Entities"],
)
async def api_list_entities(
bank_id: str,
limit: int = Query(default=100, description="Maximum number of entities to return"),
offset: int = Query(default=0, description="Offset for pagination"),
request_context: RequestContext = Depends(get_request_context),
):
"""List entities for a memory bank."""
"""List entities for a memory bank with pagination."""
try:
entities = await app.state.memory.list_entities(bank_id, limit=limit, request_context=request_context)
return EntityListResponse(items=[EntityListItem(**e) for e in entities])
data = await app.state.memory.list_entities(
bank_id, limit=limit, offset=offset, request_context=request_context
)
return EntityListResponse(
items=[EntityListItem(**e) for e in data["items"]],
total=data["total"],
limit=data["limit"],
offset=data["offset"],
)
except (AuthenticationError, HTTPException):
raise
except Exception as e:
@@ -1670,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,
@@ -2032,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,
@@ -2050,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(
+158 -15
View File
@@ -4,9 +4,12 @@ Centralized configuration for Hindsight API.
All environment variables and their defaults are defined here.
"""
import json
import logging
import os
import sys
from dataclasses import dataclass
from datetime import datetime, timezone
from dotenv import find_dotenv, load_dotenv
@@ -41,20 +44,40 @@ 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"
ENV_RERANKER_LOCAL_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT"
ENV_RERANKER_TEI_URL = "HINDSIGHT_API_RERANKER_TEI_URL"
ENV_RERANKER_TEI_BATCH_SIZE = "HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE"
ENV_RERANKER_TEI_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT"
ENV_RERANKER_MAX_CANDIDATES = "HINDSIGHT_API_RERANKER_MAX_CANDIDATES"
ENV_RERANKER_FLASHRANK_MODEL = "HINDSIGHT_API_RERANKER_FLASHRANK_MODEL"
ENV_RERANKER_FLASHRANK_CACHE_DIR = "HINDSIGHT_API_RERANKER_FLASHRANK_CACHE_DIR"
ENV_HOST = "HINDSIGHT_API_HOST"
ENV_PORT = "HINDSIGHT_API_PORT"
ENV_LOG_LEVEL = "HINDSIGHT_API_LOG_LEVEL"
ENV_LOG_FORMAT = "HINDSIGHT_API_LOG_FORMAT"
ENV_WORKERS = "HINDSIGHT_API_WORKERS"
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"
@@ -66,6 +89,8 @@ ENV_OBSERVATION_TOP_ENTITIES = "HINDSIGHT_API_OBSERVATION_TOP_ENTITIES"
ENV_RETAIN_MAX_COMPLETION_TOKENS = "HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS"
ENV_RETAIN_CHUNK_SIZE = "HINDSIGHT_API_RETAIN_CHUNK_SIZE"
ENV_RETAIN_EXTRACT_CAUSAL_LINKS = "HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS"
ENV_RETAIN_EXTRACTION_MODE = "HINDSIGHT_API_RETAIN_EXTRACTION_MODE"
ENV_RETAIN_OBSERVATIONS_ASYNC = "HINDSIGHT_API_RETAIN_OBSERVATIONS_ASYNC"
# Optimization flags
ENV_SKIP_LLM_VERIFICATION = "HINDSIGHT_API_SKIP_LLM_VERIFICATION"
@@ -81,8 +106,9 @@ ENV_DB_COMMAND_TIMEOUT = "HINDSIGHT_API_DB_COMMAND_TIMEOUT"
ENV_DB_ACQUIRE_TIMEOUT = "HINDSIGHT_API_DB_ACQUIRE_TIMEOUT"
# Background task processing
ENV_TASK_BATCH_SIZE = "HINDSIGHT_API_TASK_BATCH_SIZE"
ENV_TASK_BATCH_INTERVAL = "HINDSIGHT_API_TASK_BATCH_INTERVAL"
ENV_TASK_BACKEND = "HINDSIGHT_API_TASK_BACKEND"
ENV_TASK_BACKEND_MEMORY_BATCH_SIZE = "HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_SIZE"
ENV_TASK_BACKEND_MEMORY_BATCH_INTERVAL = "HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_INTERVAL"
# Default values
DEFAULT_DATABASE_URL = "pg0"
@@ -98,15 +124,31 @@ DEFAULT_EMBEDDING_DIMENSION = 384
DEFAULT_RERANKER_PROVIDER = "local"
DEFAULT_RERANKER_LOCAL_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2"
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT = 4 # Limit concurrent CPU-bound reranking to prevent thrashing
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
DEFAULT_RERANKER_MAX_CANDIDATES = 300
DEFAULT_RERANKER_FLASHRANK_MODEL = "ms-marco-MiniLM-L-12-v2" # Best balance of speed and quality
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"
DEFAULT_LOG_FORMAT = "text" # Options: "text", "json"
DEFAULT_WORKERS = 1
DEFAULT_MCP_ENABLED = True
DEFAULT_GRAPH_RETRIEVER = "bfs" # Options: "bfs", "mpfp"
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
@@ -117,6 +159,9 @@ DEFAULT_OBSERVATION_TOP_ENTITIES = 5 # Max entities to process per retain batch
DEFAULT_RETAIN_MAX_COMPLETION_TOKENS = 64000 # Max tokens for fact extraction LLM call
DEFAULT_RETAIN_CHUNK_SIZE = 3000 # Max chars per chunk for fact extraction
DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS = True # Extract causal links between facts
DEFAULT_RETAIN_EXTRACTION_MODE = "concise" # Extraction mode: "concise" or "verbose"
RETAIN_EXTRACTION_MODES = ("concise", "verbose") # Allowed extraction modes
DEFAULT_RETAIN_OBSERVATIONS_ASYNC = False # Run observation generation async (after retain completes)
# Database migrations
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
@@ -128,8 +173,9 @@ DEFAULT_DB_COMMAND_TIMEOUT = 60 # seconds
DEFAULT_DB_ACQUIRE_TIMEOUT = 30 # seconds
# Background task processing
DEFAULT_TASK_BATCH_SIZE = 10
DEFAULT_TASK_BATCH_INTERVAL = 1.0 # seconds
DEFAULT_TASK_BACKEND = "memory" # Options: "memory", "noop"
DEFAULT_TASK_BACKEND_MEMORY_BATCH_SIZE = 10
DEFAULT_TASK_BACKEND_MEMORY_BATCH_INTERVAL = 1.0 # seconds
# Default MCP tool descriptions (can be customized via env vars)
DEFAULT_MCP_RETAIN_DESCRIPTION = """Store important information to long-term memory.
@@ -155,6 +201,48 @@ Use this tool PROACTIVELY to:
EMBEDDING_DIMENSION = DEFAULT_EMBEDDING_DIMENSION
class JsonFormatter(logging.Formatter):
"""JSON formatter for structured logging.
Outputs logs in JSON format with a 'severity' field that cloud logging
systems (GCP, AWS CloudWatch, etc.) can parse to correctly categorize log levels.
"""
SEVERITY_MAP = {
logging.DEBUG: "DEBUG",
logging.INFO: "INFO",
logging.WARNING: "WARNING",
logging.ERROR: "ERROR",
logging.CRITICAL: "CRITICAL",
}
def format(self, record: logging.LogRecord) -> str:
log_entry = {
"severity": self.SEVERITY_MAP.get(record.levelno, "DEFAULT"),
"message": record.getMessage(),
"timestamp": datetime.now(timezone.utc).isoformat(),
"logger": record.name,
}
# Add exception info if present
if record.exc_info:
log_entry["exception"] = self.formatException(record.exc_info)
return json.dumps(log_entry)
def _validate_extraction_mode(mode: str) -> str:
"""Validate and normalize extraction mode."""
mode_lower = mode.lower()
if mode_lower not in RETAIN_EXTRACTION_MODES:
logger.warning(
f"Invalid extraction mode '{mode}', must be one of {RETAIN_EXTRACTION_MODES}. "
f"Defaulting to '{DEFAULT_RETAIN_EXTRACTION_MODE}'."
)
return DEFAULT_RETAIN_EXTRACTION_MODE
return mode_lower
@dataclass
class HindsightConfig:
"""Configuration container for Hindsight API."""
@@ -185,20 +273,30 @@ 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
reranker_local_model: str
reranker_tei_url: str | None
reranker_tei_batch_size: int
reranker_tei_max_concurrent: int
reranker_max_candidates: int
reranker_cohere_base_url: str | None
# Server
host: str
port: int
log_level: str
log_format: str
mcp_enabled: bool
# Recall
graph_retriever: str
mpfp_top_k_neighbors: int
recall_max_concurrent: int
recall_connection_budget: int
# Observation thresholds
observation_min_facts: int
@@ -208,6 +306,8 @@ class HindsightConfig:
retain_max_completion_tokens: int
retain_chunk_size: int
retain_extract_causal_links: bool
retain_extraction_mode: str
retain_observations_async: bool
# Optimization flags
skip_llm_verification: bool
@@ -223,8 +323,9 @@ class HindsightConfig:
db_acquire_timeout: int
# Background task processing
task_batch_size: int
task_batch_interval: float
task_backend: str
task_backend_memory_batch_size: int
task_backend_memory_batch_interval: float
@classmethod
def from_env(cls) -> "HindsightConfig":
@@ -252,17 +353,31 @@ 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),
reranker_tei_url=os.getenv(ENV_RERANKER_TEI_URL),
reranker_tei_batch_size=int(os.getenv(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE))),
reranker_tei_max_concurrent=int(
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)),
log_level=os.getenv(ENV_LOG_LEVEL, DEFAULT_LOG_LEVEL),
log_format=os.getenv(ENV_LOG_FORMAT, DEFAULT_LOG_FORMAT).lower(),
mcp_enabled=os.getenv(ENV_MCP_ENABLED, str(DEFAULT_MCP_ENABLED)).lower() == "true",
# Recall
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",
@@ -280,6 +395,13 @@ class HindsightConfig:
ENV_RETAIN_EXTRACT_CAUSAL_LINKS, str(DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS)
).lower()
== "true",
retain_extraction_mode=_validate_extraction_mode(
os.getenv(ENV_RETAIN_EXTRACTION_MODE, DEFAULT_RETAIN_EXTRACTION_MODE)
),
retain_observations_async=os.getenv(
ENV_RETAIN_OBSERVATIONS_ASYNC, str(DEFAULT_RETAIN_OBSERVATIONS_ASYNC)
).lower()
== "true",
# Database migrations
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
# Database connection pool
@@ -288,8 +410,13 @@ class HindsightConfig:
db_command_timeout=int(os.getenv(ENV_DB_COMMAND_TIMEOUT, str(DEFAULT_DB_COMMAND_TIMEOUT))),
db_acquire_timeout=int(os.getenv(ENV_DB_ACQUIRE_TIMEOUT, str(DEFAULT_DB_ACQUIRE_TIMEOUT))),
# Background task processing
task_batch_size=int(os.getenv(ENV_TASK_BATCH_SIZE, str(DEFAULT_TASK_BATCH_SIZE))),
task_batch_interval=float(os.getenv(ENV_TASK_BATCH_INTERVAL, str(DEFAULT_TASK_BATCH_INTERVAL))),
task_backend=os.getenv(ENV_TASK_BACKEND, DEFAULT_TASK_BACKEND),
task_backend_memory_batch_size=int(
os.getenv(ENV_TASK_BACKEND_MEMORY_BATCH_SIZE, str(DEFAULT_TASK_BACKEND_MEMORY_BATCH_SIZE))
),
task_backend_memory_batch_interval=float(
os.getenv(ENV_TASK_BACKEND_MEMORY_BATCH_INTERVAL, str(DEFAULT_TASK_BACKEND_MEMORY_BATCH_INTERVAL))
),
)
def get_llm_base_url(self) -> str:
@@ -320,12 +447,28 @@ class HindsightConfig:
return log_level_map.get(self.log_level.lower(), logging.INFO)
def configure_logging(self) -> None:
"""Configure Python logging based on the log level."""
logging.basicConfig(
level=self.get_python_log_level(),
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
force=True, # Override any existing configuration
)
"""Configure Python logging based on the log level and format.
When log_format is "json", outputs structured JSON logs with a severity
field that GCP Cloud Logging can parse for proper log level categorization.
"""
root_logger = logging.getLogger()
root_logger.setLevel(self.get_python_log_level())
# Remove existing handlers
for handler in root_logger.handlers[:]:
root_logger.removeHandler(handler)
# Create handler writing to stdout (GCP treats stderr as ERROR)
handler = logging.StreamHandler(sys.stdout)
handler.setLevel(self.get_python_log_level())
if self.log_format == "json":
handler.setFormatter(JsonFormatter())
else:
handler.setFormatter(logging.Formatter("%(asctime)s - %(levelname)s - %(name)s - %(message)s"))
root_logger.addHandler(handler)
def log_config(self) -> None:
"""Log the current configuration (without sensitive values)."""
@@ -6,20 +6,38 @@ Provides an interface for reranking with different backends.
Configuration via environment variables - see hindsight_api.config for all env var names.
"""
import asyncio
import logging
import os
from abc import ABC, abstractmethod
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,
ENV_RERANKER_TEI_BATCH_SIZE,
ENV_RERANKER_TEI_MAX_CONCURRENT,
ENV_RERANKER_TEI_URL,
)
@@ -50,7 +68,7 @@ class CrossEncoderModel(ABC):
pass
@abstractmethod
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs for relevance.
@@ -73,25 +91,34 @@ class LocalSTCrossEncoder(CrossEncoderModel):
- Fast inference (~80ms for 100 pairs on CPU)
- Small model (80MB)
- Trained for passage re-ranking
Uses a dedicated thread pool to limit concurrent CPU-bound work.
"""
def __init__(self, model_name: str | None = None):
# Shared executor across all instances (one model loaded anyway)
_executor: ThreadPoolExecutor | None = None
_max_concurrent: int = 4 # Limit concurrent CPU-bound reranking calls
def __init__(self, model_name: str | None = None, max_concurrent: int = 4):
"""
Initialize local SentenceTransformers cross-encoder.
Args:
model_name: Name of the CrossEncoder model to use.
Default: cross-encoder/ms-marco-MiniLM-L-6-v2
max_concurrent: Maximum concurrent reranking calls (default: 2).
Higher values may cause CPU thrashing under load.
"""
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
self._model = None
LocalSTCrossEncoder._max_concurrent = max_concurrent
@property
def provider_name(self) -> str:
return "local"
async def initialize(self) -> None:
"""Load the cross-encoder model."""
"""Load the cross-encoder model and initialize the executor."""
if self._model is not None:
return
@@ -103,14 +130,30 @@ class LocalSTCrossEncoder(CrossEncoderModel):
"Install it with: pip install sentence-transformers"
)
# Note: We use CPU even when GPU/MPS is available because:
# 1. The reranker model (MiniLM) is tiny (~22M params)
# 2. Batch sizes are small (~100-200 pairs)
# 3. Data transfer overhead to GPU outweighs compute benefit
# 4. CPU inference is actually faster for this workload
logger.info(f"Reranker: initializing local provider with model {self.model_name}")
self._model = CrossEncoder(self.model_name)
logger.info("Reranker: local provider initialized")
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
# Initialize shared executor (limited workers naturally limits concurrency)
if LocalSTCrossEncoder._executor is None:
LocalSTCrossEncoder._executor = ThreadPoolExecutor(
max_workers=LocalSTCrossEncoder._max_concurrent,
thread_name_prefix="reranker",
)
logger.info(f"Reranker: local provider initialized (max_concurrent={LocalSTCrossEncoder._max_concurrent})")
else:
logger.info("Reranker: local provider initialized (using existing executor)")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs for relevance.
Uses a dedicated thread pool with limited workers to prevent CPU thrashing.
Args:
pairs: List of (query, document) tuples to score
@@ -119,7 +162,13 @@ class LocalSTCrossEncoder(CrossEncoderModel):
"""
if self._model is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
scores = self._model.predict(pairs, show_progress_bar=False)
# Use dedicated executor - limited workers naturally limits concurrency
loop = asyncio.get_event_loop()
scores = await loop.run_in_executor(
LocalSTCrossEncoder._executor,
lambda: self._model.predict(pairs, show_progress_bar=False),
)
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
@@ -131,13 +180,21 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
See: https://github.com/huggingface/text-embeddings-inference
Note: The TEI server must be running a cross-encoder/reranker model.
Requests are made in parallel with configurable batch size and max concurrency (backpressure).
Uses a GLOBAL semaphore to limit concurrent requests across ALL recall operations.
"""
# Global semaphore shared across all instances and calls to prevent thundering herd
_global_semaphore: asyncio.Semaphore | None = None
_global_max_concurrent: int = DEFAULT_RERANKER_TEI_MAX_CONCURRENT
def __init__(
self,
base_url: str,
timeout: float = 30.0,
batch_size: int = 32,
batch_size: int = DEFAULT_RERANKER_TEI_BATCH_SIZE,
max_concurrent: int = DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
max_retries: int = 3,
retry_delay: float = 0.5,
):
@@ -147,138 +204,187 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
Args:
base_url: Base URL of the TEI server (e.g., "http://localhost:8080")
timeout: Request timeout in seconds (default: 30.0)
batch_size: Maximum batch size for rerank requests (default: 32)
batch_size: Maximum batch size for rerank requests (default: 128)
max_concurrent: Maximum concurrent requests for backpressure (default: 8).
This is a GLOBAL limit across all parallel recall operations.
max_retries: Maximum number of retries for failed requests (default: 3)
retry_delay: Initial delay between retries in seconds, doubles each retry (default: 0.5)
"""
self.base_url = base_url.rstrip("/")
self.timeout = timeout
self.batch_size = batch_size
self.max_concurrent = max_concurrent
self.max_retries = max_retries
self.retry_delay = retry_delay
self._client: httpx.Client | None = None
self._async_client: httpx.AsyncClient | None = None
self._model_id: str | None = None
# Update global semaphore if max_concurrent changed
if (
RemoteTEICrossEncoder._global_semaphore is None
or RemoteTEICrossEncoder._global_max_concurrent != max_concurrent
):
RemoteTEICrossEncoder._global_max_concurrent = max_concurrent
RemoteTEICrossEncoder._global_semaphore = asyncio.Semaphore(max_concurrent)
@property
def provider_name(self) -> str:
return "tei"
def _request_with_retry(self, method: str, url: str, **kwargs) -> httpx.Response:
"""Make an HTTP request with automatic retries on transient errors."""
import time
async def _async_request_with_retry(
self,
client: httpx.AsyncClient,
semaphore: asyncio.Semaphore,
method: str,
url: str,
**kwargs,
) -> httpx.Response:
"""Make an async HTTP request with automatic retries on transient errors and semaphore for backpressure."""
last_error = None
delay = self.retry_delay
for attempt in range(self.max_retries + 1):
try:
if method == "GET":
response = self._client.get(url, **kwargs)
else:
response = self._client.post(url, **kwargs)
response.raise_for_status()
return response
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
last_error = e
if attempt < self.max_retries:
logger.warning(
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
)
time.sleep(delay)
delay *= 2 # Exponential backoff
except httpx.HTTPStatusError as e:
# Retry on 5xx server errors
if e.response.status_code >= 500 and attempt < self.max_retries:
async with semaphore:
for attempt in range(self.max_retries + 1):
try:
if method == "GET":
response = await client.get(url, **kwargs)
else:
response = await client.post(url, **kwargs)
response.raise_for_status()
return response
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
last_error = e
logger.warning(
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
)
time.sleep(delay)
delay *= 2
else:
raise
if attempt < self.max_retries:
logger.warning(
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. "
f"Retrying in {delay}s..."
)
await asyncio.sleep(delay)
delay *= 2 # Exponential backoff
except httpx.HTTPStatusError as e:
# Retry on 5xx server errors
if e.response.status_code >= 500 and attempt < self.max_retries:
last_error = e
logger.warning(
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. "
f"Retrying in {delay}s..."
)
await asyncio.sleep(delay)
delay *= 2
else:
raise
raise last_error
async def initialize(self) -> None:
"""Initialize the HTTP client and verify server connectivity."""
if self._client is not None:
if self._async_client is not None:
return
logger.info(f"Reranker: initializing TEI provider at {self.base_url}")
self._client = httpx.Client(timeout=self.timeout)
logger.info(
f"Reranker: initializing TEI provider at {self.base_url} "
f"(batch_size={self.batch_size}, max_concurrent={self.max_concurrent})"
)
self._async_client = httpx.AsyncClient(timeout=self.timeout)
# Verify server is reachable and get model info
# Use a temporary semaphore for initialization
init_semaphore = asyncio.Semaphore(1)
try:
response = self._request_with_retry("GET", f"{self.base_url}/info")
response = await self._async_request_with_retry(
self._async_client, init_semaphore, "GET", f"{self.base_url}/info"
)
info = response.json()
self._model_id = info.get("model_id", "unknown")
logger.info(f"Reranker: TEI provider initialized (model: {self._model_id})")
except httpx.HTTPError as e:
self._async_client = None
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
async def _rerank_query_group(
self,
client: httpx.AsyncClient,
semaphore: asyncio.Semaphore,
query: str,
texts: list[str],
) -> list[tuple[int, float]]:
"""Rerank a single query group and return list of (original_index, score) tuples."""
try:
response = await self._async_request_with_retry(
client,
semaphore,
"POST",
f"{self.base_url}/rerank",
json={
"query": query,
"texts": texts,
"return_text": False,
},
)
results = response.json()
# TEI returns results sorted by score descending, with original index
return [(result["index"], result["score"]) for result in results]
except httpx.HTTPError as e:
raise RuntimeError(f"TEI rerank request failed: {e}")
async def _predict_async(self, pairs: list[tuple[str, str]]) -> list[float]:
"""Async implementation of predict that runs requests in parallel with backpressure."""
if not pairs:
return []
# Group all pairs by query
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))
# Split each query group into batches
tasks_info: list[tuple[str, list[int], list[str]]] = [] # (query, indices, texts)
for query, indexed_texts in query_groups.items():
indices = [idx for idx, _ in indexed_texts]
texts = [text for _, text in indexed_texts]
# Split into batches
for i in range(0, len(texts), self.batch_size):
batch_indices = indices[i : i + self.batch_size]
batch_texts = texts[i : i + self.batch_size]
tasks_info.append((query, batch_indices, batch_texts))
# Run all requests in parallel with GLOBAL semaphore for backpressure
# This ensures max_concurrent is respected across ALL parallel recall operations
all_scores = [0.0] * len(pairs)
semaphore = RemoteTEICrossEncoder._global_semaphore
tasks = [
self._rerank_query_group(self._async_client, semaphore, query, texts) for query, _, texts in tasks_info
]
results = await asyncio.gather(*tasks)
# Map scores back to original positions
for (_, indices, _), result_scores in zip(tasks_info, results):
for original_idx_in_batch, score in result_scores:
global_idx = indices[original_idx_in_batch]
all_scores[global_idx] = score
return all_scores
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using the remote TEI reranker.
Requests are made in parallel with configurable backpressure.
Args:
pairs: List of (query, document) tuples to score
Returns:
List of relevance scores
"""
if self._client is None:
if self._async_client is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
if not pairs:
return []
all_scores = []
# Process in batches
for i in range(0, len(pairs), self.batch_size):
batch = pairs[i : i + self.batch_size]
# TEI rerank endpoint expects query and texts separately
# All pairs in a batch should have the same query for optimal performance
# but we handle mixed queries by making separate requests per unique query
query_groups: dict[str, list[tuple[int, str]]] = {}
for idx, (query, text) in enumerate(batch):
if query not in query_groups:
query_groups[query] = []
query_groups[query].append((idx, text))
batch_scores = [0.0] * len(batch)
for query, indexed_texts in query_groups.items():
texts = [text for _, text in indexed_texts]
indices = [idx for idx, _ in indexed_texts]
try:
response = self._request_with_retry(
"POST",
f"{self.base_url}/rerank",
json={
"query": query,
"texts": texts,
"return_text": False,
},
)
results = response.json()
# TEI returns results sorted by score descending, with original index
for result in results:
original_idx = result["index"]
score = result["score"]
# Map back to batch position
batch_scores[indices[original_idx]] = score
except httpx.HTTPError as e:
raise RuntimeError(f"TEI rerank request failed: {e}")
all_scores.extend(batch_scores)
return all_scores
return await self._predict_async(pairs)
class CohereCrossEncoder(CrossEncoderModel):
@@ -292,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,
):
"""
@@ -300,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
@@ -321,11 +430,17 @@ 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")
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using the Cohere Rerank API.
@@ -341,6 +456,12 @@ class CohereCrossEncoder(CrossEncoderModel):
if not pairs:
return []
# Run sync Cohere API calls in thread pool
loop = asyncio.get_event_loop()
return await loop.run_in_executor(None, self._predict_sync, pairs)
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
"""Synchronous predict implementation for Cohere API."""
# Group pairs by query for efficient batching
# Cohere rerank expects one query with multiple documents
query_groups: dict[str, list[tuple[int, str]]] = {}
@@ -371,6 +492,280 @@ class CohereCrossEncoder(CrossEncoderModel):
return all_scores
class RRFPassthroughCrossEncoder(CrossEncoderModel):
"""
Passthrough cross-encoder that preserves RRF scores without neural reranking.
This is useful for:
- Testing retrieval quality without reranking overhead
- Deployments where reranking latency is unacceptable
- Debugging to isolate retrieval vs reranking issues
"""
def __init__(self):
"""Initialize RRF passthrough cross-encoder."""
pass
@property
def provider_name(self) -> str:
return "rrf"
async def initialize(self) -> None:
"""No initialization needed."""
logger.info("Reranker: RRF passthrough provider initialized (neural reranking disabled)")
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Return neutral scores - actual ranking uses RRF scores from retrieval.
Args:
pairs: List of (query, document) tuples (ignored)
Returns:
List of 0.5 scores (neutral, lets RRF scores dominate)
"""
# Return neutral scores so RRF ranking is preserved
return [0.5] * len(pairs)
class FlashRankCrossEncoder(CrossEncoderModel):
"""
FlashRank cross-encoder implementation.
FlashRank is an ultra-lite reranking library that runs on CPU without
requiring PyTorch or Transformers. It's ideal for serverless deployments
with minimal cold-start overhead.
Available models:
- ms-marco-TinyBERT-L-2-v2: Fastest, ~4MB
- ms-marco-MiniLM-L-12-v2: Best quality, ~34MB (default)
- rank-T5-flan: Best zero-shot, ~110MB
- ms-marco-MultiBERT-L-12: Multi-lingual, ~150MB
"""
# Shared executor for CPU-bound reranking
_executor: ThreadPoolExecutor | None = None
_max_concurrent: int = 4
def __init__(
self,
model_name: str | None = None,
cache_dir: str | None = None,
max_length: int = 512,
max_concurrent: int = 4,
):
"""
Initialize FlashRank cross-encoder.
Args:
model_name: FlashRank model name. Default: ms-marco-MiniLM-L-12-v2
cache_dir: Directory to cache downloaded models. Default: system cache
max_length: Maximum sequence length for reranking. Default: 512
max_concurrent: Maximum concurrent reranking calls. Default: 4
"""
self.model_name = model_name or DEFAULT_RERANKER_FLASHRANK_MODEL
self.cache_dir = cache_dir or DEFAULT_RERANKER_FLASHRANK_CACHE_DIR
self.max_length = max_length
self._ranker = None
FlashRankCrossEncoder._max_concurrent = max_concurrent
@property
def provider_name(self) -> str:
return "flashrank"
async def initialize(self) -> None:
"""Load the FlashRank model."""
if self._ranker is not None:
return
try:
from flashrank import Ranker # type: ignore[import-untyped]
except ImportError:
raise ImportError("flashrank is required for FlashRankCrossEncoder. Install it with: pip install flashrank")
logger.info(f"Reranker: initializing FlashRank provider with model {self.model_name}")
# Initialize ranker with optional cache directory
ranker_kwargs = {"model_name": self.model_name, "max_length": self.max_length}
if self.cache_dir:
ranker_kwargs["cache_dir"] = self.cache_dir
self._ranker = Ranker(**ranker_kwargs)
# Initialize shared executor
if FlashRankCrossEncoder._executor is None:
FlashRankCrossEncoder._executor = ThreadPoolExecutor(
max_workers=FlashRankCrossEncoder._max_concurrent,
thread_name_prefix="flashrank",
)
logger.info(
f"Reranker: FlashRank provider initialized (max_concurrent={FlashRankCrossEncoder._max_concurrent})"
)
else:
logger.info("Reranker: FlashRank provider initialized (using existing executor)")
def _predict_sync(self, pairs: list[tuple[str, str]]) -> list[float]:
"""Synchronous predict - processes each query group."""
from flashrank import RerankRequest # type: ignore[import-untyped]
if not pairs:
return []
# Group pairs by query
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():
# Build passages list for FlashRank
passages = [{"id": i, "text": text} for i, (_, text) in enumerate(indexed_texts)]
global_indices = [idx for idx, _ in indexed_texts]
# Create rerank request
request = RerankRequest(query=query, passages=passages)
results = self._ranker.rerank(request)
# Map scores back to original positions
for result in results:
local_idx = result["id"]
score = result["score"]
global_idx = global_indices[local_idx]
all_scores[global_idx] = score
return all_scores
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using FlashRank.
Args:
pairs: List of (query, document) tuples to score
Returns:
List of relevance scores (higher = more relevant)
"""
if self._ranker is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
# Run in thread pool to avoid blocking event loop
loop = asyncio.get_event_loop()
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.
@@ -386,16 +781,35 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
url = os.environ.get(ENV_RERANKER_TEI_URL)
if not url:
raise ValueError(f"{ENV_RERANKER_TEI_URL} is required when {ENV_RERANKER_PROVIDER} is 'tei'")
return RemoteTEICrossEncoder(base_url=url)
batch_size = int(os.environ.get(ENV_RERANKER_TEI_BATCH_SIZE, str(DEFAULT_RERANKER_TEI_BATCH_SIZE)))
max_concurrent = int(os.environ.get(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT)))
return RemoteTEICrossEncoder(base_url=url, batch_size=batch_size, max_concurrent=max_concurrent)
elif provider == "local":
model = os.environ.get(ENV_RERANKER_LOCAL_MODEL)
model_name = model or DEFAULT_RERANKER_LOCAL_MODEL
return LocalSTCrossEncoder(model_name=model_name)
max_concurrent = int(
os.environ.get(ENV_RERANKER_LOCAL_MAX_CONCURRENT, str(DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT))
)
return LocalSTCrossEncoder(model_name=model_name, max_concurrent=max_concurrent)
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_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'")
raise ValueError(
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)
@@ -83,11 +83,22 @@ async def acquire_with_retry(pool: asyncpg.Pool, max_retries: int = DEFAULT_MAX_
Yields:
An asyncpg connection
"""
import time
start = time.time()
async def acquire():
return await pool.acquire()
conn = await retry_with_backoff(acquire, max_retries=max_retries)
acquire_time = time.time() - start
# Log slow connection acquisitions (indicates pool contention)
if acquire_time > 0.05: # 50ms threshold
pool_size = pool.get_size()
pool_free = pool.get_idle_size()
logger.warning(f"[DB POOL] Slow acquire: {acquire_time:.3f}s | size={pool_size}, idle={pool_free}")
try:
yield conn
finally:
@@ -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'"
)
@@ -406,18 +406,20 @@ class MemoryEngineInterface(ABC):
bank_id: str,
*,
limit: int = 100,
offset: int = 0,
request_context: "RequestContext",
) -> list[dict[str, Any]]:
) -> dict[str, Any]:
"""
List entities for a bank.
List entities for a bank with pagination.
Args:
bank_id: The memory bank ID.
limit: Maximum results.
offset: Offset for pagination.
request_context: Request context for authentication.
Returns:
List of entity dicts.
Dict with items, total, limit, offset.
"""
...
@@ -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,7 +151,8 @@ 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 .task_backend import AsyncIOQueueBackend, TaskBackend
from .search.tags import TagsMatch
from .task_backend import AsyncIOQueueBackend, NoopTaskBackend, TaskBackend
class Budget(str, Enum):
@@ -257,8 +259,8 @@ class MemoryEngine(MemoryEngineInterface):
db_command_timeout: PostgreSQL command timeout in seconds. Defaults to HINDSIGHT_API_DB_COMMAND_TIMEOUT.
db_acquire_timeout: Connection acquisition timeout in seconds. Defaults to HINDSIGHT_API_DB_ACQUIRE_TIMEOUT.
task_backend: Custom task backend. If not provided, uses AsyncIOQueueBackend.
task_batch_size: Background task batch size. Defaults to HINDSIGHT_API_TASK_BATCH_SIZE.
task_batch_interval: Background task batch interval in seconds. Defaults to HINDSIGHT_API_TASK_BATCH_INTERVAL.
task_batch_size: Background task batch size. Defaults to HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_SIZE.
task_batch_interval: Background task batch interval in seconds. Defaults to HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_INTERVAL.
run_migrations: Whether to run database migrations during initialize(). Default: True
operation_validator: Optional extension to validate operations before execution.
If provided, retain/recall/reflect operations will be validated.
@@ -396,17 +398,21 @@ class MemoryEngine(MemoryEngineInterface):
self._cross_encoder_reranker = CrossEncoderReranker(cross_encoder=cross_encoder)
# Initialize task backend
_task_batch_size = task_batch_size if task_batch_size is not None else config.task_batch_size
_task_batch_interval = task_batch_interval if task_batch_interval is not None else config.task_batch_interval
self._task_backend = task_backend or AsyncIOQueueBackend(
batch_size=_task_batch_size, batch_interval=_task_batch_interval
)
if task_backend:
self._task_backend = task_backend
elif config.task_backend == "noop":
self._task_backend = NoopTaskBackend()
else:
# Default to memory (AsyncIOQueueBackend)
_task_batch_size = task_batch_size if task_batch_size is not None else config.task_backend_memory_batch_size
_task_batch_interval = (
task_batch_interval if task_batch_interval is not None else config.task_backend_memory_batch_interval
)
self._task_backend = AsyncIOQueueBackend(batch_size=_task_batch_size, batch_interval=_task_batch_interval)
# Backpressure mechanism: limit concurrent searches to prevent overwhelming the database
# Limit concurrent searches to prevent connection pool exhaustion
# Each search can use 2-4 connections, so with 10 concurrent searches
# we use ~20-40 connections max, staying well within pool limits
self._search_semaphore = asyncio.Semaphore(10)
# Configurable via HINDSIGHT_API_RECALL_MAX_CONCURRENT (default: 50)
self._search_semaphore = asyncio.Semaphore(get_config().recall_max_concurrent)
# Backpressure for put operations: limit concurrent puts to prevent database contention
# Each put_batch holds a connection for the entire transaction, so we limit to 5
@@ -1054,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,
):
"""
@@ -1186,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
@@ -1204,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
@@ -1238,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.
@@ -1254,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)
@@ -1278,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(
@@ -1336,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).
@@ -1361,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:
@@ -1412,7 +1428,9 @@ class MemoryEngine(MemoryEngineInterface):
# Backpressure: limit concurrent recalls to prevent overwhelming the database
result = None
error_msg = None
semaphore_wait_start = time.time()
async with self._search_semaphore:
semaphore_wait = time.time() - semaphore_wait_start
# Retry loop for connection errors
max_retries = 3
for attempt in range(max_retries + 1):
@@ -1430,6 +1448,9 @@ class MemoryEngine(MemoryEngineInterface):
include_chunks,
max_chunk_tokens,
request_context,
semaphore_wait=semaphore_wait,
tags=tags,
tags_match=tags_match,
)
break # Success - exit retry loop
except Exception as e:
@@ -1547,6 +1568,9 @@ class MemoryEngine(MemoryEngineInterface):
include_chunks: bool = False,
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.
@@ -1576,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()
@@ -1586,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:
@@ -1601,37 +1628,67 @@ 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 retrieve_parallel
from .search.retrieval import (
get_default_graph_retriever,
retrieve_all_fact_types_parallel,
)
# Track each retrieval start time
retrieval_start = time.time()
# Run retrieval for each fact type in parallel
retrieval_tasks = [
retrieve_parallel(
pool, query, query_embedding_str, bank_id, ft, thinking_budget, question_date, self.query_analyzer
# 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,
fact_type, # Pass all fact types at once
thinking_budget,
question_date,
self.query_analyzer,
tags=tags,
tags_match=tags_match,
)
for ft in fact_type
]
all_retrievals = await asyncio.gather(*retrieval_tasks)
parallel_duration = time.time() - parallel_start
# Combine all results from all fact types and aggregate timings
semantic_results = []
bm25_results = []
graph_results = []
temporal_results = []
aggregated_timings = {"semantic": 0.0, "bm25": 0.0, "graph": 0.0, "temporal": 0.0}
aggregated_timings = {
"semantic": 0.0,
"bm25": 0.0,
"graph": 0.0,
"temporal": 0.0,
"temporal_extraction": 0.0,
}
all_mpfp_timings = []
detected_temporal_constraint = None
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)
@@ -1645,6 +1702,8 @@ class MemoryEngine(MemoryEngineInterface):
# Capture temporal constraint (same across all fact types)
if retrieval_result.temporal_constraint:
detected_temporal_constraint = retrieval_result.temporal_constraint
# Collect MPFP timings
all_mpfp_timings.extend(retrieval_result.mpfp_timings)
# If no temporal results from any fact type, set to None
if not temporal_results:
@@ -1663,12 +1722,12 @@ class MemoryEngine(MemoryEngineInterface):
retrieval_duration = time.time() - retrieval_start
step_duration = time.time() - step_start
total_retrievals = len(fact_type) * (4 if temporal_results else 3)
# Format per-method timings
# Format per-method timings (these are the actual parallel retrieval times)
timing_parts = [
f"semantic={len(semantic_results)}({aggregated_timings['semantic']:.3f}s)",
f"bm25={len(bm25_results)}({aggregated_timings['bm25']:.3f}s)",
f"graph={len(graph_results)}({aggregated_timings['graph']:.3f}s)",
f"temporal_extraction={aggregated_timings['temporal_extraction']:.3f}s",
]
temporal_info = ""
if detected_temporal_constraint:
@@ -1677,9 +1736,41 @@ class MemoryEngine(MemoryEngineInterface):
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')}"
log_buffer.append(
f" [2] {total_retrievals}-way retrieval ({len(fact_type)} fact_types): {', '.join(timing_parts)} in {step_duration:.3f}s{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
if all_mpfp_timings:
retriever_name = get_default_graph_retriever().name.upper()
mpfp_total = all_mpfp_timings[0] # Take first fact type's timing as representative
mpfp_parts = [
f"db_queries={mpfp_total.db_queries}",
f"edge_load={mpfp_total.edge_load_time:.3f}s",
f"edges={mpfp_total.edge_count}",
f"patterns={mpfp_total.pattern_count}",
]
if mpfp_total.seeds_time > 0.01:
mpfp_parts.append(f"seeds={mpfp_total.seeds_time:.3f}s")
if mpfp_total.fusion > 0.001:
mpfp_parts.append(f"fusion={mpfp_total.fusion:.3f}s")
if mpfp_total.fetch > 0.001:
mpfp_parts.append(f"fetch={mpfp_total.fetch:.3f}s")
log_buffer.append(f" [{retriever_name}] {', '.join(mpfp_parts)}")
# Log detailed hop timing for debugging slow queries
if mpfp_total.hop_details:
for hd in mpfp_total.hop_details:
log_buffer.append(
f" hop{hd['hop']}: exec={hd.get('exec_time', 0) * 1000:.0f}ms, "
f"uncached={hd.get('uncached_after_filter', 0)}, "
f"load={hd.get('load_time', 0) * 1000:.0f}ms, "
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
@@ -1687,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(
@@ -1720,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,
)
@@ -1777,11 +1878,24 @@ class MemoryEngine(MemoryEngineInterface):
# Ensure reranker is initialized (for lazy initialization mode)
await reranker_instance.ensure_initialized()
# Pre-filter candidates to reduce reranking cost (RRF already provides good ranking)
# This is especially important for remote rerankers with network latency
reranker_max_candidates = get_config().reranker_max_candidates
pre_filtered_count = 0
if len(merged_candidates) > reranker_max_candidates:
# Sort by RRF score and take top candidates
merged_candidates.sort(key=lambda mc: mc.rrf_score, reverse=True)
pre_filtered_count = len(merged_candidates) - reranker_max_candidates
merged_candidates = merged_candidates[:reranker_max_candidates]
# Rerank using cross-encoder
scored_results = reranker_instance.rerank(query, merged_candidates)
scored_results = await reranker_instance.rerank(query, merged_candidates)
step_duration = time.time() - step_start
log_buffer.append(f" [4] Reranking: {len(scored_results)} candidates scored in {step_duration:.3f}s")
pre_filter_note = f" (pre-filtered {pre_filtered_count})" if pre_filtered_count > 0 else ""
log_buffer.append(
f" [4] Reranking: {len(scored_results)} candidates scored in {step_duration:.3f}s{pre_filter_note}"
)
# Step 4.5: Combine cross-encoder score with retrieval signals
# This preserves retrieval work (RRF, temporal, recency) instead of pure cross-encoder ranking
@@ -1831,9 +1945,6 @@ class MemoryEngine(MemoryEngineInterface):
# Re-sort by combined score
scored_results.sort(key=lambda x: x.weight, reverse=True)
log_buffer.append(
" [4.6] Combined scoring: cross_encoder(0.6) + rrf(0.2) + temporal(0.1) + recency(0.1)"
)
# Add reranked results to tracer AFTER combined scoring (so normalized values are included)
if tracer:
@@ -1852,7 +1963,6 @@ class MemoryEngine(MemoryEngineInterface):
# Step 5: Truncate to thinking_budget * 2 for token filtering
rerank_limit = thinking_budget * 2
top_scored = scored_results[:rerank_limit]
log_buffer.append(f" [5] Truncated to top {len(top_scored)} results")
# Step 6: Token budget filtering
step_start = time.time()
@@ -1867,7 +1977,7 @@ class MemoryEngine(MemoryEngineInterface):
step_duration = time.time() - step_start
log_buffer.append(
f" [6] Token filtering: {len(top_scored)} results, {total_tokens}/{max_tokens} tokens in {step_duration:.3f}s"
f" [5] Token filtering: {len(top_scored)} results, {total_tokens}/{max_tokens} tokens in {step_duration:.3f}s"
)
if tracer:
@@ -1901,7 +2011,6 @@ class MemoryEngine(MemoryEngineInterface):
visited_ids = list(set([sr.id for sr in scored_results[:50]])) # Top 50
if visited_ids:
await self._task_backend.submit_task({"type": "access_count_update", "node_ids": visited_ids})
log_buffer.append(f" [7] Queued access count updates for {len(visited_ids)} nodes")
# Log fact_type distribution in results
fact_type_counts = {}
@@ -1934,6 +2043,7 @@ class MemoryEngine(MemoryEngineInterface):
top_results_dicts.append(result_dict)
# Get entities for each fact if include_entities is requested
step_start = time.time()
fact_entity_map = {} # unit_id -> list of (entity_id, entity_name)
if include_entities and top_scored:
unit_ids = [uuid.UUID(sr.id) for sr in top_scored]
@@ -1955,6 +2065,7 @@ class MemoryEngine(MemoryEngineInterface):
fact_entity_map[unit_id].append(
{"entity_id": str(row["entity_id"]), "canonical_name": row["canonical_name"]}
)
entity_map_duration = time.time() - step_start
# Convert results to MemoryFact objects
memory_facts = []
@@ -1977,10 +2088,12 @@ 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"),
)
)
# Fetch entity observations if requested
step_start = time.time()
entities_dict = None
total_entity_tokens = 0
total_chunk_tokens = 0
@@ -2001,7 +2114,13 @@ class MemoryEngine(MemoryEngineInterface):
entities_ordered.append((entity_id, entity_name))
seen_entity_ids.add(entity_id)
# Fetch observations for each entity (respect token budget, in order)
# Fetch all observations in a single batched query
entity_ids = [eid for eid, _ in entities_ordered]
all_observations = await self.get_entity_observations_batch(
bank_id, entity_ids, limit_per_entity=5, request_context=request_context
)
# Build entities_dict respecting token budget, in relevance order
entities_dict = {}
encoding = _get_tiktoken_encoding()
@@ -2009,9 +2128,7 @@ class MemoryEngine(MemoryEngineInterface):
if total_entity_tokens >= max_entity_tokens:
break
observations = await self.get_entity_observations(
bank_id, entity_id, limit=5, request_context=request_context
)
observations = all_observations.get(entity_id, [])
# Calculate tokens for this entity's observations
entity_tokens = 0
@@ -2029,8 +2146,10 @@ class MemoryEngine(MemoryEngineInterface):
entity_id=entity_id, canonical_name=entity_name, observations=included_observations
)
total_entity_tokens += entity_tokens
entity_obs_duration = time.time() - step_start
# Fetch chunks if requested
step_start = time.time()
chunks_dict = None
if include_chunks and top_scored:
from .response_models import ChunkInfo
@@ -2090,6 +2209,12 @@ class MemoryEngine(MemoryEngineInterface):
chunk_text=chunk_text, chunk_index=row["chunk_index"], truncated=False
)
total_chunk_tokens += chunk_tokens
chunks_duration = time.time() - step_start
# Log entity/chunk fetch timing (only if any enrichment was requested)
log_buffer.append(
f" [6] Response enrichment: entity_map={entity_map_duration:.3f}s, entity_obs={entity_obs_duration:.3f}s, chunks={chunks_duration:.3f}s"
)
# Finalize trace if enabled
trace_dict = None
@@ -2101,8 +2226,15 @@ class MemoryEngine(MemoryEngineInterface):
total_time = time.time() - recall_start
num_chunks = len(chunks_dict) if chunks_dict else 0
num_entities = len(entities_dict) if entities_dict else 0
# Include wait times in log if significant
wait_parts = []
if semaphore_wait > 0.01:
wait_parts.append(f"sem={semaphore_wait:.3f}s")
if max_conn_wait > 0.01:
wait_parts.append(f"conn={max_conn_wait:.3f}s")
wait_info = f" | waits: {', '.join(wait_parts)}" if wait_parts else ""
log_buffer.append(
f"[RECALL {recall_id}] Complete: {len(top_scored)} facts ({total_tokens} tok), {num_chunks} chunks ({total_chunk_tokens} tok), {num_entities} entities ({total_entity_tokens} tok) | {fact_type_summary} | {total_time:.3f}s"
f"[RECALL {recall_id}] Complete: {len(top_scored)} facts ({total_tokens} tok), {num_chunks} chunks ({total_chunk_tokens} tok), {num_entities} entities ({total_entity_tokens} tok) | {fact_type_summary} | {total_time:.3f}s{wait_info}"
)
logger.info("\n" + "\n".join(log_buffer))
@@ -2172,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,
@@ -2193,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(
@@ -2298,9 +2431,10 @@ class MemoryEngine(MemoryEngineInterface):
await self._authenticate_tenant(request_context)
pool = await self._get_pool()
async with acquire_with_retry(pool) as conn:
# Ensure connection is not in read-only mode (can happen with connection poolers)
await conn.execute("SET SESSION CHARACTERISTICS AS TRANSACTION READ WRITE")
async with conn.transaction():
# Ensure transaction is not in read-only mode (can happen with connection poolers)
# Using SET LOCAL so it only affects this transaction, not the session
await conn.execute("SET LOCAL transaction_read_only TO off")
try:
if fact_type:
# Delete only memories of a specific fact type
@@ -2680,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,
@@ -3203,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.
@@ -3270,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
@@ -3485,37 +3685,110 @@ Guidelines:
observations.append(EntityObservation(text=row["text"], mentioned_at=mentioned_at))
return observations
async def get_entity_observations_batch(
self,
bank_id: str,
entity_ids: list[str],
*,
limit_per_entity: int = 5,
request_context: "RequestContext",
) -> dict[str, list[Any]]:
"""
Get observations for multiple entities in a single query.
Args:
bank_id: bank IDentifier
entity_ids: List of entity UUIDs to get observations for
limit_per_entity: Maximum observations per entity
request_context: Request context for authentication.
Returns:
Dict mapping entity_id -> list of EntityObservation objects
"""
if not entity_ids:
return {}
await self._authenticate_tenant(request_context)
pool = await self._get_pool()
async with acquire_with_retry(pool) as conn:
# Use window function to limit observations per entity
rows = await conn.fetch(
f"""
WITH ranked AS (
SELECT
ue.entity_id,
mu.text,
mu.mentioned_at,
ROW_NUMBER() OVER (PARTITION BY ue.entity_id ORDER BY mu.mentioned_at DESC) as rn
FROM {fq_table("memory_units")} mu
JOIN {fq_table("unit_entities")} ue ON mu.id = ue.unit_id
WHERE mu.bank_id = $1
AND mu.fact_type = 'observation'
AND ue.entity_id = ANY($2::uuid[])
)
SELECT entity_id, text, mentioned_at
FROM ranked
WHERE rn <= $3
ORDER BY entity_id, rn
""",
bank_id,
[uuid.UUID(eid) for eid in entity_ids],
limit_per_entity,
)
result: dict[str, list[Any]] = {eid: [] for eid in entity_ids}
for row in rows:
entity_id = str(row["entity_id"])
mentioned_at = row["mentioned_at"].isoformat() if row["mentioned_at"] else None
result[entity_id].append(EntityObservation(text=row["text"], mentioned_at=mentioned_at))
return result
async def list_entities(
self,
bank_id: str,
*,
limit: int = 100,
offset: int = 0,
request_context: "RequestContext",
) -> list[dict[str, Any]]:
) -> dict[str, Any]:
"""
List all entities for a bank.
List all entities for a bank with pagination.
Args:
bank_id: bank IDentifier
limit: Maximum number of entities to return
offset: Offset for pagination
request_context: Request context for authentication.
Returns:
List of entity dicts with id, canonical_name, mention_count, first_seen, last_seen
Dict with items, total, limit, offset
"""
await self._authenticate_tenant(request_context)
pool = await self._get_pool()
async with acquire_with_retry(pool) as conn:
# Get total count
total_row = await conn.fetchrow(
f"""
SELECT COUNT(*) as total
FROM {fq_table("entities")}
WHERE bank_id = $1
""",
bank_id,
)
total = total_row["total"] if total_row else 0
# Get paginated entities
rows = await conn.fetch(
f"""
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
LIMIT $2
ORDER BY mention_count DESC, last_seen DESC, id ASC
LIMIT $2 OFFSET $3
""",
bank_id,
limit,
offset,
)
entities = []
@@ -3542,7 +3815,91 @@ Guidelines:
"metadata": metadata,
}
)
return entities
return {
"items": entities,
"total": total,
"limit": limit,
"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,
@@ -4188,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)
@@ -4211,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}")
@@ -84,7 +84,7 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
Performance:
- ~10-50ms per query
- No model loading required
- No model loading required (lazy import on first use)
"""
def __init__(self):
@@ -112,8 +112,6 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
Returns:
QueryAnalysis with temporal_constraint if found
"""
self.load()
if reference_date is None:
reference_date = datetime.now()
@@ -123,6 +121,9 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
if period_result is not None:
return QueryAnalysis(temporal_constraint=period_result)
# Lazy load dateparser (only imports on first call, then cached)
self.load()
# Use dateparser's search_dates to find temporal expressions
settings = {
"RELATIVE_BASE": reference_date,
@@ -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):
@@ -210,6 +210,98 @@ class FactExtractionResponse(BaseModel):
facts: list[ExtractedFact] = Field(description="List of extracted factual statements")
class ExtractedFactVerbose(BaseModel):
"""A single extracted fact with verbose field descriptions for detailed extraction."""
model_config = ConfigDict(
json_schema_mode="validation",
json_schema_extra={"required": ["what", "when", "where", "who", "why", "fact_type"]},
)
what: str = Field(
description="WHAT happened - COMPLETE, DETAILED description with ALL specifics. "
"NEVER summarize or omit details. Include: exact actions, objects, quantities, specifics. "
"BE VERBOSE - capture every detail that was mentioned. "
"Example: 'Emily got married to Sarah at a rooftop garden ceremony with 50 guests attending and a live jazz band playing' "
"NOT: 'A wedding happened' or 'Emily got married'"
)
when: str = Field(
description="WHEN it happened - ALWAYS include temporal information if mentioned. "
"Include: specific dates, times, durations, relative time references. "
"Examples: 'on June 15th, 2024 at 3pm', 'last weekend', 'for the past 3 years', 'every morning at 6am'. "
"Write 'N/A' ONLY if absolutely no temporal context exists. Prefer converting to absolute dates when possible."
)
where: str = Field(
description="WHERE it happened or is about - SPECIFIC locations, places, areas, regions if applicable. "
"Include: cities, neighborhoods, venues, buildings, countries, specific addresses when mentioned. "
"Examples: 'downtown San Francisco at a rooftop garden venue', 'at the user's home in Brooklyn', 'online via Zoom', 'Paris, France'. "
"Write 'N/A' ONLY if absolutely no location context exists or if the fact is completely location-agnostic."
)
who: str = Field(
description="WHO is involved - ALL people/entities with FULL context and relationships. "
"Include: names, roles, relationships to user, background details. "
"Resolve coreferences (if 'my roommate' is later named 'Emily', write 'Emily, the user's college roommate'). "
"BE DETAILED about relationships and roles. "
"Example: 'Emily (user's college roommate from Stanford, now works at Google), Sarah (Emily's partner of 5 years, software engineer)' "
"NOT: 'my friend' or 'Emily and Sarah'"
)
why: str = Field(
description="WHY it matters - ALL emotional, contextual, and motivational details. "
"Include EVERYTHING: feelings, preferences, motivations, observations, context, background, significance. "
"BE VERBOSE - capture all the nuance and meaning. "
"FOR ASSISTANT FACTS: MUST include what the user asked/requested that led to this interaction! "
"Example (world): 'The user felt thrilled and inspired, has always dreamed of an outdoor ceremony, mentioned wanting a similar garden venue, was particularly moved by the intimate atmosphere and personal vows' "
"Example (assistant): 'User asked how to fix slow API performance with 1000+ concurrent users, expected 70-80% reduction in database load' "
"NOT: 'User liked it' or 'To help user'"
)
fact_kind: str = Field(
default="conversation",
description="'event' = specific datable occurrence (set occurred dates), 'conversation' = general info (no occurred dates)",
)
occurred_start: str | None = Field(
default=None,
description="WHEN the event happened (ISO timestamp). Only for fact_kind='event'. Leave null for conversations.",
)
occurred_end: str | None = Field(
default=None,
description="WHEN the event ended (ISO timestamp). Only for events with duration. Leave null for conversations.",
)
fact_type: Literal["world", "assistant"] = Field(
description="'world' = about the user/others (background, experiences). 'assistant' = experience with the assistant."
)
entities: list[Entity] | None = Field(
default=None,
description="Named entities, objects, AND abstract concepts from the fact. Include: people names, organizations, places, significant objects (e.g., 'coffee maker', 'car'), AND abstract concepts/themes (e.g., 'friendship', 'career growth', 'loss', 'celebration'). Extract anything that could help link related facts together.",
)
causal_relations: list[FactCausalRelation] | None = Field(
default=None,
description="Causal links to PREVIOUS facts only. target_index MUST be less than this fact's position. "
"Example: fact #3 can only reference facts 0, 1, or 2. Max 2 relations per fact.",
)
@field_validator("entities", mode="before")
@classmethod
def ensure_entities_list(cls, v):
if v is None:
return []
return v
class FactExtractionResponseVerbose(BaseModel):
"""Response for verbose fact extraction."""
facts: list[ExtractedFactVerbose] = Field(description="List of extracted factual statements")
class ExtractedFactNoCausal(BaseModel):
"""A single extracted fact WITHOUT causal relations (for when causal extraction is disabled)."""
@@ -342,35 +434,12 @@ def _chunk_conversation(turns: list[dict], max_chars: int) -> list[str]:
return chunks if chunks else [json.dumps(turns, ensure_ascii=False)]
async def _extract_facts_from_chunk(
chunk: str,
chunk_index: int,
total_chunks: int,
event_date: datetime,
context: str,
llm_config: "LLMConfig",
agent_name: str = None,
extract_opinions: bool = False,
) -> tuple[list[dict[str, str]], TokenUsage]:
"""
Extract facts from a single chunk (internal helper for parallel processing).
# =============================================================================
# FACT EXTRACTION PROMPTS
# =============================================================================
Note: event_date parameter is kept for backward compatibility but not used in prompt.
The LLM extracts temporal information from the context string instead.
"""
memory_bank_context = f"\n- Your name: {agent_name}" if agent_name and extract_opinions else ""
# Determine which fact types to extract based on the flag
# Note: We use "assistant" in the prompt but convert to "bank" for storage
if extract_opinions:
# Opinion extraction uses a separate prompt (not this one)
fact_types_instruction = "Extract ONLY 'opinion' type facts (formed opinions, beliefs, and perspectives). DO NOT extract 'world' or 'assistant' facts."
else:
fact_types_instruction = (
"Extract ONLY 'world' and 'assistant' type facts. DO NOT extract opinions - those are extracted separately."
)
prompt = f"""Extract SIGNIFICANT facts from text. Be SELECTIVE - only extract facts worth remembering long-term.
# Concise extraction prompt (default) - selective, high-quality facts
CONCISE_FACT_EXTRACTION_PROMPT = """Extract SIGNIFICANT facts from text. Be SELECTIVE - only extract facts worth remembering long-term.
LANGUAGE RULE (CRITICAL): Output facts in the EXACT SAME language as the input text. If input is Japanese, output Japanese. If input is Chinese, output Chinese. NEVER translate to English. Preserve original language completely.
@@ -470,8 +539,123 @@ QUALITY OVER QUANTITY
Ask: "Would this be useful to recall in 6 months?" If no, skip it."""
# Causal relationships section - only included if enabled in config
causal_relationships_section = """
# Verbose extraction prompt - detailed, comprehensive facts (legacy mode)
VERBOSE_FACT_EXTRACTION_PROMPT = """Extract facts from text into structured format with FIVE required dimensions - BE EXTREMELY DETAILED.
LANGUAGE REQUIREMENT: Detect the language of the input text. All extracted facts, entity names, descriptions,
and other output MUST be in the SAME language as the input. Do not translate to English if the input is in another language.
{fact_types_instruction}
══════════════════════════════════════════════════════════════════════════
FACT FORMAT - ALL FIVE DIMENSIONS REQUIRED - MAXIMUM VERBOSITY
══════════════════════════════════════════════════════════════════════════
For EACH fact, CAPTURE ALL DETAILS - NEVER SUMMARIZE OR OMIT:
1. **what**: WHAT happened - COMPLETE description with ALL specifics (objects, actions, quantities, details)
2. **when**: WHEN it happened - ALWAYS include temporal info with DAY OF WEEK (e.g., "Monday, June 10, 2024")
- Always include the day name: Monday, Tuesday, Wednesday, Thursday, Friday, Saturday, Sunday
- Format: "day_name, month day, year" (e.g., "Saturday, June 9, 2024")
3. **where**: WHERE it happened or is about - SPECIFIC locations, places, areas, regions (if applicable)
4. **who**: WHO is involved - ALL people/entities with FULL relationships and background
5. **why**: WHY it matters - ALL emotions, preferences, motivations, significance, nuance
- For assistant facts: MUST include what the user asked/requested that triggered this!
Plus: fact_type, fact_kind, entities, occurred_start/end (for structured dates), where (structured location)
VERBOSITY REQUIREMENT: Include EVERY detail mentioned. More detail is ALWAYS better than less.
══════════════════════════════════════════════════════════════════════════
COREFERENCE RESOLUTION (CRITICAL)
══════════════════════════════════════════════════════════════════════════
When text uses BOTH a generic relation AND a name for the same person → LINK THEM!
Example input: "I went to my college roommate's wedding last June. Emily finally married Sarah after 5 years together."
CORRECT output:
- what: "Emily got married to Sarah at a rooftop garden ceremony"
- when: "Saturday, June 8, 2024, after dating for 5 years"
- where: "downtown San Francisco, at a rooftop garden venue"
- who: "Emily (user's college roommate), Sarah (Emily's partner of 5 years)"
- why: "User found it romantic and beautiful, dreams of similar outdoor ceremony"
- where (structured): "San Francisco"
WRONG output:
- what: "User's roommate got married" ← LOSES THE NAME!
- who: "the roommate" ← WRONG - use the actual name!
- where: (missing) ← WRONG - include the location!
══════════════════════════════════════════════════════════════════════════
FACT_KIND CLASSIFICATION (CRITICAL FOR TEMPORAL HANDLING)
══════════════════════════════════════════════════════════════════════════
⚠️ MUST set fact_kind correctly - this determines whether occurred_start/end are set!
fact_kind="event" - USE FOR:
- Actions that happened at a specific time: "went to", "attended", "visited", "bought", "made"
- Past events: "yesterday I...", "last week...", "in March 2020..."
- Future plans with dates: "will go to", "scheduled for"
- Examples: "I went to a pottery workshop" → event
"Alice visited Paris in February" → event
"I bought a new car yesterday" → event
"The user graduated from MIT in March 2020" → event
fact_kind="conversation" - USE FOR:
- Ongoing states: "works as", "lives in", "is married to"
- Preferences: "loves", "prefers", "enjoys"
- Traits/abilities: "speaks fluent French", "knows Python"
- Examples: "I love Italian food" → conversation
"Alice works at Google" → conversation
"I prefer outdoor dining" → conversation
══════════════════════════════════════════════════════════════════════════
TEMPORAL HANDLING (CRITICAL - USE EVENT DATE AS REFERENCE)
══════════════════════════════════════════════════════════════════════════
⚠️ IMPORTANT: Use the "Event Date" provided in the input as your reference point!
All relative dates ("yesterday", "last week", "recently") must be resolved relative to the Event Date, NOT today's date.
For EVENTS (fact_kind="event") - MUST SET BOTH occurred_start AND occurred_end:
- Convert relative dates → absolute using Event Date as reference
- If Event Date is "Saturday, March 15, 2020", then "yesterday" = Friday, March 14, 2020
- Dates mentioned in text (e.g., "in March 2020") should use THAT year, not current year
- Always include the day name (Monday, Tuesday, etc.) in the 'when' field
- Set occurred_start AND occurred_end to WHEN IT HAPPENED (not when mentioned)
- For single-day/point events: set occurred_end = occurred_start (same timestamp)
For CONVERSATIONS (fact_kind="conversation"):
- General info, preferences, ongoing states → NO occurred dates
- Examples: "loves coffee", "works as engineer"
══════════════════════════════════════════════════════════════════════════
FACT TYPE
══════════════════════════════════════════════════════════════════════════
- **world**: User's life, other people, events (would exist without this conversation)
- **assistant**: Interactions with assistant (requests, recommendations, help)
⚠️ CRITICAL for assistant facts: ALWAYS capture the user's request/question in the fact!
Include: what the user asked, what problem they wanted solved, what context they provided
══════════════════════════════════════════════════════════════════════════
ENTITIES - EXTRACT EVERYTHING
══════════════════════════════════════════════════════════════════════════
Extract ALL of the following from the fact:
- People names (Emily, Alice, Dr. Smith)
- Organizations (Google, MIT, local coffee shop)
- Places (San Francisco, Brooklyn, Paris)
- Significant objects mentioned (coffee maker, new car, wedding dress)
- Abstract concepts/themes (friendship, career growth, loss, celebration)
ALWAYS include "user" when fact is about the user.
Extract anything that could help link related facts together."""
# Causal relationships section - appended when causal extraction is enabled
CAUSAL_RELATIONSHIPS_SECTION = """
══════════════════════════════════════════════════════════════════════════
CAUSAL RELATIONSHIPS
@@ -485,14 +669,57 @@ Example: "Lost job → couldn't pay rent → moved apartment"
- Fact 1: Couldn't pay rent, causal_relations: [{target_index: 0, relation_type: "caused_by"}]
- Fact 2: Moved apartment, causal_relations: [{target_index: 1, relation_type: "caused_by"}]"""
# Check config for causal link extraction
async def _extract_facts_from_chunk(
chunk: str,
chunk_index: int,
total_chunks: int,
event_date: datetime,
context: str,
llm_config: "LLMConfig",
agent_name: str = None,
extract_opinions: bool = False,
) -> tuple[list[dict[str, str]], TokenUsage]:
"""
Extract facts from a single chunk (internal helper for parallel processing).
Note: event_date parameter is kept for backward compatibility but not used in prompt.
The LLM extracts temporal information from the context string instead.
"""
memory_bank_context = f"\n- Your name: {agent_name}" if agent_name and extract_opinions else ""
# Determine which fact types to extract based on the flag
# Note: We use "assistant" in the prompt but convert to "bank" for storage
if extract_opinions:
# Opinion extraction uses a separate prompt (not this one)
fact_types_instruction = "Extract ONLY 'opinion' type facts (formed opinions, beliefs, and perspectives). DO NOT extract 'world' or 'assistant' facts."
else:
fact_types_instruction = (
"Extract ONLY 'world' and 'assistant' type facts. DO NOT extract opinions - those are extracted separately."
)
# Check config for extraction mode and causal link extraction
config = get_config()
extraction_mode = config.retain_extraction_mode
extract_causal_links = config.retain_extract_causal_links
# Select base prompt based on extraction mode
if extraction_mode == "verbose":
base_prompt = VERBOSE_FACT_EXTRACTION_PROMPT
else:
base_prompt = CONCISE_FACT_EXTRACTION_PROMPT
# Format the prompt with fact types instruction
prompt = base_prompt.format(fact_types_instruction=fact_types_instruction)
# Build the full prompt with or without causal relationships section
# Select appropriate response schema based on extraction mode and causal links
if extract_causal_links:
prompt = prompt + causal_relationships_section
response_schema = FactExtractionResponse
prompt = prompt + CAUSAL_RELATIONSHIPS_SECTION
if extraction_mode == "verbose":
response_schema = FactExtractionResponseVerbose
else:
response_schema = FactExtractionResponse
else:
response_schema = FactExtractionResponseNoCausal
@@ -898,7 +1125,7 @@ async def extract_facts_from_text(
# Log chunk count before starting LLM requests
total_chars = sum(len(c) for c in chunks)
if len(chunks) > 1:
logger.info(
logger.debug(
f"[FACT_EXTRACTION] Text chunked into {len(chunks)} chunks ({total_chars:,} chars total, "
f"chunk_size={config.retain_chunk_size:,}) - starting parallel LLM extraction"
)
@@ -1041,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 [],
)
@@ -9,6 +9,7 @@ import time
import uuid
from datetime import UTC, datetime
from ...config import get_config
from ..db_utils import acquire_with_retry
from . import bank_utils
@@ -48,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.
@@ -66,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)
@@ -87,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)
@@ -130,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
@@ -158,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
@@ -224,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
@@ -268,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)
@@ -395,16 +408,26 @@ async def retain_batch(
causal_link_count = await link_creation.create_causal_links_batch(conn, unit_ids, non_duplicate_facts)
log_buffer.append(f"[10] Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s")
# Regenerate observations INSIDE transaction for atomicity
await observation_regeneration.regenerate_observations_batch(
conn, embeddings_model, llm_config, bank_id, entity_links, log_buffer
)
# Regenerate observations - sync (in transaction) or async (background task)
config = get_config()
if config.retain_observations_async:
# Queue for async processing after transaction commits
entity_ids_for_async = list(set(link.entity_id for link in entity_links)) if entity_links else []
log_buffer.append(
f"[11] Observations: queued {len(entity_ids_for_async)} entities for async processing"
)
else:
# Run synchronously inside transaction for atomicity
await observation_regeneration.regenerate_observations_batch(
conn, embeddings_model, llm_config, bank_id, entity_links, log_buffer
)
entity_ids_for_async = []
# Map results back to original content items
result_unit_ids = _map_results_to_contents(contents, extracted_facts, is_duplicate_flags, unit_ids)
# Trigger background tasks AFTER transaction commits (opinion reinforcement only)
await _trigger_background_tasks(task_backend, bank_id, unit_ids, non_duplicate_facts)
# Trigger background tasks AFTER transaction commits
await _trigger_background_tasks(task_backend, bank_id, unit_ids, non_duplicate_facts, entity_ids_for_async)
# Log final summary
total_time = time.time() - start_time
@@ -454,8 +477,9 @@ async def _trigger_background_tasks(
bank_id: str,
unit_ids: list[str],
facts: list[ProcessedFact],
entity_ids_for_observations: list[str] | None = None,
) -> None:
"""Trigger opinion reinforcement as background task (after transaction commits)."""
"""Trigger background tasks after transaction commits."""
# Trigger opinion reinforcement if there are entities
fact_entities = [[e.name for e in fact.entities] for fact in facts]
if any(fact_entities):
@@ -468,3 +492,13 @@ async def _trigger_background_tasks(
"unit_entities": fact_entities,
}
)
# Trigger observation regeneration if async mode is enabled
if entity_ids_for_observations:
await task_backend.submit_task(
{
"type": "regenerate_observations",
"bank_id": bank_id,
"entity_ids": entity_ids_for_observations,
}
)
@@ -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,7 +11,8 @@ from abc import ABC, abstractmethod
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .types import RetrievalResult
from .tags import TagsMatch, filter_results_by_tags
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -42,7 +43,10 @@ class GraphRetriever(ABC):
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
) -> list[RetrievalResult]:
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.
@@ -55,9 +59,11 @@ class GraphRetriever(ABC):
query_text: Original query text (optional, for some strategies)
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:
List of RetrievalResult objects with activation scores set
Tuple of (List of RetrievalResult with activation scores, optional timing info)
"""
pass
@@ -111,7 +117,10 @@ class BFSGraphRetriever(GraphRetriever):
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
) -> list[RetrievalResult]:
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.
@@ -122,11 +131,14 @@ class BFSGraphRetriever(GraphRetriever):
4. Return visited nodes up to budget
Note: BFS finds its own entry points via embedding search.
The semantic_seeds and temporal_seeds parameters are accepted
The semantic_seeds, temporal_seeds, and adjacency parameters are accepted
for interface compatibility but not used.
"""
async with acquire_with_retry(pool) as conn:
return 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(
self,
@@ -135,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 = []
@@ -192,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
@@ -232,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
@@ -0,0 +1,256 @@
"""
Link Expansion graph retrieval.
A simple, fast graph retrieval that expands from seeds via:
1. Entity links: Find facts sharing entities with seeds (filtered by entity frequency)
2. Causal links: Find facts causally linked to seeds (top-k by weight)
Characteristics:
- 2-3 DB queries (seed finding + parallel entity/causal expansion)
- Sublinear: only touches connected facts via indexes
- No iteration, no propagation, no normalization
- Target: <100ms
"""
import logging
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__)
async def _find_semantic_seeds(
conn,
query_embedding_str: str,
bank_id: str,
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, 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
""",
*params,
)
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
class LinkExpansionRetriever(GraphRetriever):
"""
Graph retrieval via direct link expansion from seeds.
Expands through entity co-occurrence and causal links in a single query.
Fast and simple alternative to MPFP.
"""
def __init__(
self,
max_entity_frequency: int = 500,
causal_weight_threshold: float = 0.3,
causal_limit_per_seed: int = 10,
):
"""
Initialize link expansion retriever.
Args:
max_entity_frequency: Skip entities appearing in more than this many facts
causal_weight_threshold: Minimum weight for causal links
causal_limit_per_seed: Max causal links to follow per seed
"""
self.max_entity_frequency = max_entity_frequency
self.causal_weight_threshold = causal_weight_threshold
self.causal_limit_per_seed = causal_limit_per_seed
@property
def name(self) -> str:
return "link_expansion"
async def retrieve(
self,
pool,
query_embedding_str: str,
bank_id: str,
fact_type: str,
budget: int,
query_text: str | None = None,
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.
Args:
pool: Database connection pool
query_embedding_str: Query embedding (unused, kept for interface)
bank_id: Memory bank ID
fact_type: Fact type to filter
budget: Maximum results to return
query_text: Original query text (unused)
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)
"""
start_time = time.time()
timings = MPFPTimings(fact_type=fact_type)
# Use single connection for all queries to reduce pool pressure
# (queries are fast ~50ms each, connection acquisition is the bottleneck)
async with acquire_with_retry(pool) as conn:
# Find seeds if not provided
if semantic_seeds:
all_seeds = list(semantic_seeds)
else:
seeds_start = time.time()
all_seeds = await _find_semantic_seeds(
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})
timings.pattern_count = len(seed_ids)
# Run entity and causal expansion sequentially on same connection
query_start = time.time()
entity_rows = 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,
COUNT(*)::float AS score
FROM {fq_table("unit_entities")} seed_ue
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
JOIN {fq_table("unit_entities")} other_ue ON seed_ue.entity_id = other_ue.entity_id
JOIN {fq_table("memory_units")} mu ON other_ue.unit_id = mu.id
WHERE seed_ue.unit_id = ANY($1::uuid[])
AND e.mention_count < $2
AND mu.id != ALL($1::uuid[])
AND mu.fact_type = $3
GROUP BY mu.id
ORDER BY score DESC
LIMIT $4
""",
seed_ids,
self.max_entity_frequency,
fact_type,
budget,
)
causal_rows = await conn.fetch(
f"""
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.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
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.link_type IN ('causes', 'caused_by', 'enables', 'prevents')
AND ml.weight >= $2
AND mu.fact_type = $3
ORDER BY mu.id, ml.weight DESC
LIMIT $4
""",
seed_ids,
self.causal_weight_threshold,
fact_type,
budget,
)
timings.edge_load_time = time.time() - query_start
timings.db_queries = 2
timings.edge_count = len(entity_rows) + len(causal_rows)
# Merge results, taking max score per fact
score_map: dict[str, float] = {}
row_map: dict[str, dict] = {}
for row in entity_rows:
fact_id = str(row["id"])
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
row_map[fact_id] = dict(row)
for row in causal_rows:
fact_id = str(row["id"])
score_map[fact_id] = max(score_map.get(fact_id, 0), row["score"])
if fact_id not in row_map:
row_map[fact_id] = dict(row)
# Sort by score and limit
sorted_ids = sorted(score_map.keys(), key=lambda x: score_map[x], reverse=True)[:budget]
rows = [row_map[fact_id] for fact_id in sorted_ids]
# Convert to results
results = []
for row in rows:
result = RetrievalResult.from_db_row(dict(row))
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
logger.debug(
f"LinkExpansion: {len(results)} results from {len(seed_ids)} seeds "
f"in {timings.traverse * 1000:.1f}ms (query: {timings.edge_load_time * 1000:.1f}ms)"
)
return results, timings
@@ -9,6 +9,7 @@ propagation from Approximate PPR.
Key properties:
- Sublinear in graph size (threshold pruning bounds active nodes)
- Lazy edge loading: only loads edges for frontier nodes, not entire graph
- Predefined patterns capture different retrieval intents
- All patterns run in parallel, results fused via RRF
- No LLM in the loop during traversal
@@ -22,7 +23,8 @@ from dataclasses import dataclass, field
from ..db_utils import acquire_with_retry
from ..memory_engine import fq_table
from .graph_retrieval import GraphRetriever
from .types import RetrievalResult
from .tags import TagsMatch
from .types import MPFPTimings, RetrievalResult
logger = logging.getLogger(__name__)
@@ -41,11 +43,27 @@ class EdgeTarget:
@dataclass
class TypedAdjacency:
"""Adjacency lists split by edge type."""
class EdgeCache:
"""
Cache for lazily-loaded edges.
# edge_type -> from_node_id -> list of (to_node_id, weight)
Grows per-hop as edges are loaded for frontier nodes.
Shared across patterns to avoid redundant loads.
Loads ALL edge types at once to minimize DB queries.
Thread-safe via asyncio lock to prevent redundant concurrent loads.
"""
# edge_type -> from_node_id -> list of EdgeTarget
graphs: dict[str, dict[str, list[EdgeTarget]]] = field(default_factory=dict)
# Track which nodes have been fully loaded (all edge types)
_fully_loaded: set[str] = field(default_factory=set)
# Timing stats
db_queries: int = 0
edge_load_time: float = 0.0
# Detailed hop timing for debugging
hop_details: list[dict] = field(default_factory=list)
# Lock to prevent redundant concurrent loads
_lock: asyncio.Lock = field(default_factory=asyncio.Lock)
def get_neighbors(self, edge_type: str, node_id: str) -> list[EdgeTarget]:
"""Get neighbors for a node via a specific edge type."""
@@ -63,6 +81,31 @@ class TypedAdjacency:
return [EdgeTarget(node_id=n.node_id, weight=n.weight / total) for n in neighbors]
def is_fully_loaded(self, node_id: str) -> bool:
"""Check if all edges for this node have been loaded."""
return node_id in self._fully_loaded
def get_uncached(self, node_ids: list[str]) -> list[str]:
"""Get node IDs that haven't been fully loaded yet."""
return [n for n in node_ids if not self.is_fully_loaded(n)]
def add_all_edges(self, edges_by_type: dict[str, dict[str, list[EdgeTarget]]], all_queried: list[str]):
"""
Add loaded edges to the cache (all edge types at once).
Args:
edges_by_type: Dict mapping edge_type -> from_node_id -> list of EdgeTarget
all_queried: All node IDs that were queried (marks them as fully loaded)
"""
for edge_type, edges in edges_by_type.items():
if edge_type not in self.graphs:
self.graphs[edge_type] = {}
for node_id, neighbors in edges.items():
self.graphs[edge_type][node_id] = neighbors
# Mark all queried nodes as fully loaded (even if they have no edges)
self._fully_loaded.update(all_queried)
@dataclass
class PatternResult:
@@ -109,66 +152,249 @@ class SeedNode:
# -----------------------------------------------------------------------------
# Core Algorithm
# Lazy Edge Loading
# -----------------------------------------------------------------------------
def mpfp_traverse(
seeds: list[SeedNode],
pattern: list[str],
adjacency: TypedAdjacency,
config: MPFPConfig,
) -> PatternResult:
async def load_all_edges_for_frontier(
pool,
node_ids: list[str],
top_k_per_type: int = 20,
) -> dict[str, dict[str, list[EdgeTarget]]]:
"""
Forward Push traversal following a meta-path pattern.
Load top-k edges per (node, edge_type) for frontier nodes.
Uses a LATERAL join to efficiently fetch only the top-k edges per type,
avoiding loading hundreds of entity edges when only 20 are needed.
Requires composite index: (from_unit_id, link_type, weight DESC)
Args:
seeds: Entry point nodes with initial scores
pattern: Sequence of edge types to follow
adjacency: Typed adjacency structure
config: Algorithm parameters
pool: Database connection pool
node_ids: Frontier node IDs to load edges for
top_k_per_type: Max edges to load per (node, link_type) pair
Returns:
PatternResult with accumulated scores per node
Dict mapping edge_type -> from_node_id -> list of EdgeTarget
"""
if not node_ids:
return {}
async with acquire_with_retry(pool) as conn:
# Use LATERAL join to get top-k per (from_node, link_type)
# This leverages the composite index for efficient early termination
rows = await conn.fetch(
f"""
WITH frontier(node_id) AS (SELECT unnest($1::uuid[]))
SELECT f.node_id as from_unit_id, lt.link_type, edges.to_unit_id, edges.weight
FROM frontier f
CROSS JOIN (VALUES ('semantic'), ('temporal'), ('entity'), ('causes'), ('caused_by')) AS lt(link_type)
CROSS JOIN LATERAL (
SELECT ml.to_unit_id, ml.weight
FROM {fq_table("memory_links")} ml
WHERE ml.from_unit_id = f.node_id
AND ml.link_type = lt.link_type
AND ml.weight >= 0.1
ORDER BY ml.weight DESC
LIMIT $2
) edges
""",
node_ids,
top_k_per_type,
)
# Group by edge_type -> from_node -> neighbors
result: dict[str, dict[str, list[EdgeTarget]]] = defaultdict(lambda: defaultdict(list))
for row in rows:
edge_type = row["link_type"]
from_id = str(row["from_unit_id"])
to_id = str(row["to_unit_id"])
weight = row["weight"]
result[edge_type][from_id].append(EdgeTarget(node_id=to_id, weight=weight))
# Convert nested defaultdicts to regular dicts
return {edge_type: dict(edges) for edge_type, edges in result.items()}
# -----------------------------------------------------------------------------
# Core Algorithm (Async with Lazy Loading)
# -----------------------------------------------------------------------------
@dataclass
class PatternState:
"""State for a pattern traversal between hops."""
pattern: list[str]
hop_index: int
scores: dict[str, float]
frontier: dict[str, float]
def _init_pattern_state(seeds: list[SeedNode], pattern: list[str]) -> PatternState:
"""Initialize pattern state from seeds."""
if not seeds:
return PatternState(pattern=pattern, hop_index=0, scores={}, frontier={})
total_seed_score = sum(s.score for s in seeds)
if total_seed_score == 0:
total_seed_score = len(seeds)
frontier = {s.node_id: s.score / total_seed_score for s in seeds}
return PatternState(pattern=pattern, hop_index=0, scores={}, frontier=frontier)
def _execute_hop(state: PatternState, cache: EdgeCache, config: MPFPConfig) -> set[str]:
"""
Execute ONE hop of traversal, return frontier nodes for next hop.
This is a pure function that uses cached edges (no DB access).
Returns set of uncached nodes needed for next hop.
"""
if state.hop_index >= len(state.pattern):
return set()
edge_type = state.pattern[state.hop_index]
# Collect active nodes above threshold
active_nodes = [node_id for node_id, mass in state.frontier.items() if mass >= config.threshold]
if not active_nodes:
state.frontier = {}
return set()
# Propagate mass using cached edges
next_frontier: dict[str, float] = {}
uncached_for_next: set[str] = set()
for node_id, mass in state.frontier.items():
if mass < config.threshold:
continue
# Keep α portion for this node
state.scores[node_id] = state.scores.get(node_id, 0) + config.alpha * mass
# Push (1-α) to neighbors
push_mass = (1 - config.alpha) * mass
neighbors = cache.get_normalized_neighbors(edge_type, node_id, config.top_k_neighbors)
for neighbor in neighbors:
next_frontier[neighbor.node_id] = next_frontier.get(neighbor.node_id, 0) + push_mass * neighbor.weight
# Track if we'll need edges for this node in the next hop
if not cache.is_fully_loaded(neighbor.node_id):
uncached_for_next.add(neighbor.node_id)
state.frontier = next_frontier
state.hop_index += 1
return uncached_for_next
def _finalize_pattern(state: PatternState, config: MPFPConfig) -> PatternResult:
"""Finalize pattern by adding remaining frontier mass to scores."""
for node_id, mass in state.frontier.items():
if mass >= config.threshold:
state.scores[node_id] = state.scores.get(node_id, 0) + mass
return PatternResult(pattern=state.pattern, scores=state.scores)
async def mpfp_traverse_hop_synchronized(
pool,
pattern_jobs: list[tuple[list[SeedNode], list[str]]],
config: MPFPConfig,
cache: EdgeCache,
) -> list[PatternResult]:
"""
Execute ALL patterns with hop-synchronized edge loading.
Instead of running each pattern independently (causing multiple DB queries),
this function:
1. Runs hop 1 for ALL patterns (using pre-warmed seed edges)
2. Collects ALL unique hop-2 frontier nodes across patterns
3. Pre-warms hop-2 edges in ONE query
4. Runs hop 2 for ALL patterns
This reduces DB queries from O(patterns * hops) to O(hops).
Args:
pool: Database connection pool
pattern_jobs: List of (seeds, pattern) tuples
config: Algorithm parameters
cache: Shared edge cache (should be pre-warmed with seed edges)
Returns:
List of PatternResult for each pattern
"""
import time
# Initialize all pattern states
states = [_init_pattern_state(seeds, pattern) for seeds, pattern in pattern_jobs]
# Determine max hops (all patterns should be same length, but be safe)
max_hops = max((len(p) for _, p in pattern_jobs), default=0)
# Detailed timing for debugging
hop_times: list[dict] = []
# Execute hop-by-hop across ALL patterns
for hop in range(max_hops):
hop_start = time.time()
hop_timing = {"hop": hop, "patterns_executed": 0, "uncached_count": 0, "load_time": 0.0}
# Execute this hop for all patterns, collect uncached nodes for next hop
all_uncached: set[str] = set()
exec_start = time.time()
for state in states:
if state.hop_index < len(state.pattern):
uncached = _execute_hop(state, cache, config)
all_uncached.update(uncached)
hop_timing["patterns_executed"] += 1
hop_timing["exec_time"] = time.time() - exec_start
# Pre-warm edges for ALL uncached nodes before next hop
hop_timing["uncached_count"] = len(all_uncached)
if all_uncached:
uncached_list = list(all_uncached - cache._fully_loaded)
hop_timing["uncached_after_filter"] = len(uncached_list)
if uncached_list:
load_start = time.time()
edges_by_type = await load_all_edges_for_frontier(pool, uncached_list, config.top_k_neighbors)
hop_timing["load_time"] = time.time() - load_start
cache.edge_load_time += hop_timing["load_time"]
cache.db_queries += 1
cache.add_all_edges(edges_by_type, uncached_list)
hop_timing["edges_loaded"] = sum(
len(neighbors) for edges in edges_by_type.values() for neighbors in edges.values()
)
hop_timing["total_time"] = time.time() - hop_start
hop_times.append(hop_timing)
# Store hop timing details in cache for logging
cache.hop_details = hop_times
# Finalize all patterns
return [_finalize_pattern(state, config) for state in states]
async def mpfp_traverse_async(
pool,
seeds: list[SeedNode],
pattern: list[str],
config: MPFPConfig,
cache: EdgeCache,
) -> PatternResult:
"""
Async Forward Push traversal with lazy edge loading.
NOTE: For better performance with multiple patterns, use mpfp_traverse_hop_synchronized().
This function is kept for single-pattern use cases.
"""
if not seeds:
return PatternResult(pattern=pattern, scores={})
scores: dict[str, float] = {}
# Initialize frontier with seed masses (normalized)
total_seed_score = sum(s.score for s in seeds)
if total_seed_score == 0:
total_seed_score = len(seeds) # fallback to uniform
frontier: dict[str, float] = {s.node_id: s.score / total_seed_score for s in seeds}
# Follow pattern hop by hop
for edge_type in pattern:
next_frontier: dict[str, float] = {}
for node_id, mass in frontier.items():
if mass < config.threshold:
continue
# Keep α portion for this node
scores[node_id] = scores.get(node_id, 0) + config.alpha * mass
# Push (1-α) to neighbors
push_mass = (1 - config.alpha) * mass
neighbors = adjacency.get_normalized_neighbors(edge_type, node_id, config.top_k_neighbors)
for neighbor in neighbors:
next_frontier[neighbor.node_id] = next_frontier.get(neighbor.node_id, 0) + push_mass * neighbor.weight
frontier = next_frontier
# Final frontier nodes get their remaining mass
for node_id, mass in frontier.items():
if mass >= config.threshold:
scores[node_id] = scores.get(node_id, 0) + mass
return PatternResult(pattern=pattern, scores=scores)
results = await mpfp_traverse_hop_synchronized(pool, [(seeds, pattern)], config, cache)
return results[0] if results else PatternResult(pattern=pattern, scores={})
def rrf_fusion(
@@ -210,38 +436,6 @@ def rrf_fusion(
# -----------------------------------------------------------------------------
async def load_typed_adjacency(pool, bank_id: str) -> TypedAdjacency:
"""
Load all edges for a bank, split by edge type.
Single query, then organize in-memory for fast traversal.
"""
async with acquire_with_retry(pool) as conn:
rows = await conn.fetch(
f"""
SELECT ml.from_unit_id, ml.to_unit_id, ml.link_type, ml.weight
FROM {fq_table("memory_links")} ml
JOIN {fq_table("memory_units")} mu ON ml.from_unit_id = mu.id
WHERE mu.bank_id = $1
AND ml.weight >= 0.1
ORDER BY ml.from_unit_id, ml.weight DESC
""",
bank_id,
)
graphs: dict[str, dict[str, list[EdgeTarget]]] = defaultdict(lambda: defaultdict(list))
for row in rows:
from_id = str(row["from_unit_id"])
to_id = str(row["to_unit_id"])
link_type = row["link_type"]
weight = row["weight"]
graphs[link_type][from_id].append(EdgeTarget(node_id=to_id, weight=weight))
return TypedAdjacency(graphs=dict(graphs))
async def fetch_memory_units_by_ids(
pool,
node_ids: list[str],
@@ -255,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
@@ -274,10 +468,10 @@ async def fetch_memory_units_by_ids(
class MPFPGraphRetriever(GraphRetriever):
"""
Graph retrieval using Meta-Path Forward Push.
Graph retrieval using Meta-Path Forward Push with lazy edge loading.
Runs predefined patterns in parallel from semantic and temporal seeds,
then fuses results via RRF.
loading edges on-demand per hop instead of loading entire graph upfront.
"""
def __init__(self, config: MPFPConfig | None = None):
@@ -287,8 +481,13 @@ class MPFPGraphRetriever(GraphRetriever):
Args:
config: Algorithm configuration (uses defaults if None)
"""
self.config = config or MPFPConfig()
self._adjacency_cache: dict[str, TypedAdjacency] = {}
if config is None:
# Read top_k_neighbors from global config
from ...config import get_config
global_config = get_config()
config = MPFPConfig(top_k_neighbors=global_config.mpfp_top_k_neighbors)
self.config = config
@property
def name(self) -> str:
@@ -304,9 +503,12 @@ class MPFPGraphRetriever(GraphRetriever):
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
) -> list[RetrievalResult]:
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.
Retrieve facts using MPFP algorithm with lazy edge loading.
Args:
pool: Database connection pool
@@ -317,12 +519,15 @@ class MPFPGraphRetriever(GraphRetriever):
query_text: Original query text (optional)
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:
List of RetrievalResult with activation scores
Tuple of (List of RetrievalResult with activation scores, MPFPTimings)
"""
# Load typed adjacency (could cache per bank_id with TTL)
adjacency = await load_typed_adjacency(pool, bank_id)
import time
timings = MPFPTimings(fact_type=fact_type)
# Convert seeds to SeedNode format
semantic_seed_nodes = self._convert_seeds(semantic_seeds, "similarity")
@@ -330,54 +535,88 @@ class MPFPGraphRetriever(GraphRetriever):
# If no semantic seeds provided, fall back to finding our own
if not semantic_seed_nodes:
semantic_seed_nodes = await self._find_semantic_seeds(pool, query_embedding_str, bank_id, fact_type)
seeds_start = time.time()
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})"
)
# Run all patterns in parallel
tasks = []
# Collect all pattern jobs
pattern_jobs = []
# Patterns from semantic seeds
for pattern in self.config.patterns_semantic:
if semantic_seed_nodes:
tasks.append(
asyncio.to_thread(
mpfp_traverse,
semantic_seed_nodes,
pattern,
adjacency,
self.config,
)
)
pattern_jobs.append((semantic_seed_nodes, pattern))
# Patterns from temporal seeds
for pattern in self.config.patterns_temporal:
if temporal_seed_nodes:
tasks.append(
asyncio.to_thread(
mpfp_traverse,
temporal_seed_nodes,
pattern,
adjacency,
self.config,
)
)
pattern_jobs.append((temporal_seed_nodes, pattern))
if not tasks:
return []
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
# Gather pattern results
pattern_results = await asyncio.gather(*tasks)
timings.pattern_count = len(pattern_jobs)
# Shared edge cache across all patterns
cache = EdgeCache()
# Pre-warm cache with ALL seed node edges BEFORE running patterns
# This prevents redundant DB queries at hop 1
all_seed_ids = list({s.node_id for seeds, _ in pattern_jobs for s in seeds})
if all_seed_ids:
import time as time_module
prewarm_start = time_module.time()
edges_by_type = await load_all_edges_for_frontier(pool, all_seed_ids, self.config.top_k_neighbors)
cache.edge_load_time += time_module.time() - prewarm_start
cache.db_queries += 1
cache.add_all_edges(edges_by_type, all_seed_ids)
# Run all patterns with HOP-SYNCHRONIZED edge loading
# This batches hop-2 edge loads across ALL patterns into ONE query
# Reduces DB queries from O(patterns * hops) to O(hops)
step_start = time.time()
pattern_results = await mpfp_traverse_hop_synchronized(pool, pattern_jobs, self.config, cache)
timings.traverse = time.time() - step_start
# Record edge loading stats from cache
timings.edge_count = sum(len(neighbors) for g in cache.graphs.values() for neighbors in g.values())
timings.db_queries = cache.db_queries
timings.edge_load_time = cache.edge_load_time
timings.hop_details = cache.hop_details
# Fuse results
step_start = time.time()
fused = rrf_fusion(pattern_results, top_k=budget)
timings.fusion = time.time() - step_start
if not fused:
return []
logger.debug(f"[MPFP] No fused results after RRF fusion (pattern_count={len(pattern_results)})")
return [], timings
# Get top result IDs (don't exclude seeds - they may be highly relevant)
# Get top result IDs
result_ids = [node_id for node_id, score in fused][:budget]
# Fetch full details
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
score_map = {node_id: score for node_id, score in fused}
@@ -387,7 +626,7 @@ class MPFPGraphRetriever(GraphRetriever):
# Sort by activation
results.sort(key=lambda r: r.activation or 0, reverse=True)
return results
return results, timings
def _convert_seeds(
self,
@@ -415,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"""
@@ -426,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]
@@ -44,7 +44,7 @@ class CrossEncoderReranker:
await cross_encoder.initialize()
self._initialized = True
def rerank(self, query: str, candidates: list[MergedCandidate]) -> list[ScoredResult]:
async def rerank(self, query: str, candidates: list[MergedCandidate]) -> list[ScoredResult]:
"""
Rerank candidates using cross-encoder scores.
@@ -85,7 +85,7 @@ class CrossEncoderReranker:
pairs.append([query, doc_text])
# Get cross-encoder scores
scores = self.cross_encoder.predict(pairs)
scores = await self.cross_encoder.predict(pairs)
# Normalize scores using sigmoid to [0, 1] range
# Cross-encoder returns logits which can be negative
File diff suppressed because it is too large Load Diff
@@ -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
@@ -10,6 +10,24 @@ from datetime import datetime
from typing import Any
@dataclass
class MPFPTimings:
"""Timing breakdown for a single MPFP retrieval call."""
fact_type: str
edge_count: int = 0 # Total edges loaded
db_queries: int = 0 # Number of DB queries for edge loading
edge_load_time: float = 0.0 # Time spent loading edges from DB
traverse: float = 0.0 # Total traversal time (includes edge loading)
pattern_count: int = 0 # Number of patterns executed
fusion: float = 0.0 # Time for RRF fusion
fetch: float = 0.0 # Time to fetch memory unit details
seeds_time: float = 0.0 # Time to find semantic seeds (if fallback used)
result_count: int = 0 # Number of results returned
# Detailed per-hop timing: list of {hop, exec_time, uncached, load_time, edges_loaded, total_time}
hop_details: list[dict] = field(default_factory=list)
@dataclass
class RetrievalResult:
"""
@@ -30,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
@@ -54,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"),
@@ -138,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,
}
@@ -121,6 +121,29 @@ class SyncTaskBackend(TaskBackend):
logger.debug("SyncTaskBackend shutdown")
class NoopTaskBackend(TaskBackend):
"""
No-op task backend that discards all tasks.
This is useful for tests where background task execution is not needed
and would only slow down the test suite.
"""
async def initialize(self):
"""No-op."""
self._initialized = True
logger.debug("NoopTaskBackend initialized")
async def submit_task(self, task_dict: dict[str, Any]):
"""Discard the task (do nothing)."""
pass
async def shutdown(self):
"""No-op."""
self._initialized = False
logger.debug("NoopTaskBackend shutdown")
class AsyncIOQueueBackend(TaskBackend):
"""
Task backend implementation using asyncio queues.
+36 -5
View File
@@ -23,7 +23,7 @@ import uvicorn
from . import MemoryEngine
from .api import create_app
from .banner import print_banner
from .config import HindsightConfig, get_config
from .config import DEFAULT_WORKERS, ENV_WORKERS, HindsightConfig, get_config
from .daemon import (
DEFAULT_DAEMON_PORT,
DEFAULT_IDLE_TIMEOUT,
@@ -95,7 +95,12 @@ def main():
# Development options
parser.add_argument("--reload", action="store_true", help="Enable auto-reload on code changes (development only)")
parser.add_argument("--workers", type=int, default=1, help="Number of worker processes (default: 1)")
parser.add_argument(
"--workers",
type=int,
default=int(os.getenv(ENV_WORKERS, str(DEFAULT_WORKERS))),
help=f"Number of worker processes (env: {ENV_WORKERS}, default: {DEFAULT_WORKERS})",
)
# Access log options
parser.add_argument("--access-log", action="store_true", help="Enable access log")
@@ -182,19 +187,31 @@ 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,
log_format=config.log_format,
mcp_enabled=config.mcp_enabled,
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,
retain_chunk_size=config.retain_chunk_size,
retain_extract_causal_links=config.retain_extract_causal_links,
retain_extraction_mode=config.retain_extraction_mode,
retain_observations_async=config.retain_observations_async,
skip_llm_verification=config.skip_llm_verification,
lazy_reranker=config.lazy_reranker,
run_migrations_on_startup=config.run_migrations_on_startup,
@@ -202,8 +219,9 @@ def main():
db_pool_max_size=config.db_pool_max_size,
db_command_timeout=config.db_command_timeout,
db_acquire_timeout=config.db_acquire_timeout,
task_batch_size=config.task_batch_size,
task_batch_interval=config.task_batch_interval,
task_backend=config.task_backend,
task_backend_memory_batch_size=config.task_backend_memory_batch_size,
task_backend_memory_batch_interval=config.task_backend_memory_batch_interval,
)
config.configure_logging()
if not args.daemon:
@@ -260,14 +278,27 @@ def main():
app = idle_middleware
# Prepare uvicorn config
# When using workers or reload, we must use import string so each worker can import the app
use_import_string = args.workers > 1 or args.reload
# Check for uvloop availability
try:
import uvloop # noqa: F401
loop_impl = "uvloop"
print("uvloop available, will use for event loop")
except ImportError:
loop_impl = "asyncio"
print("uvloop not installed, using default asyncio event loop")
uvicorn_config = {
"app": app,
"app": "hindsight_api.server:app" if use_import_string else app,
"host": args.host,
"port": args.port,
"log_level": args.log_level,
"access_log": args.access_log,
"proxy_headers": args.proxy_headers,
"ws": "wsproto", # Use wsproto instead of websockets to avoid deprecation warnings
"loop": loop_impl, # Explicitly set event loop implementation
}
# Add optional parameters if provided
+262 -1
View File
@@ -6,11 +6,18 @@ This module provides metrics for:
- Token usage (input/output) per operation
- Per-bank granularity via labels
- LLM call latency and token usage with scope dimension
- HTTP request metrics (latency, count by endpoint/method/status)
- Process metrics (CPU, memory, file descriptors, threads)
- Database connection pool metrics
"""
import logging
import os
import resource
import threading
import time
from contextlib import contextmanager
from typing import TYPE_CHECKING, Callable
from opentelemetry import metrics
from opentelemetry.exporter.prometheus import PrometheusMetricReader
@@ -18,6 +25,18 @@ from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.view import ExplicitBucketHistogramAggregation, View
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)
@@ -25,6 +44,9 @@ DURATION_BUCKETS = (0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 2
# LLM duration buckets (finer granularity for faster LLM calls)
LLM_DURATION_BUCKETS = (0.1, 0.25, 0.5, 1.0, 2.0, 3.0, 5.0, 10.0, 15.0, 30.0, 60.0, 120.0)
# HTTP request duration buckets (millisecond-level for fast endpoints)
HTTP_DURATION_BUCKETS = (0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0)
def get_token_bucket(token_count: int) -> str:
"""
@@ -107,9 +129,17 @@ def initialize_metrics(service_name: str = "hindsight-api", service_version: str
aggregation=ExplicitBucketHistogramAggregation(boundaries=LLM_DURATION_BUCKETS),
)
# Create view with custom bucket boundaries for HTTP request duration histogram
http_duration_view = View(
instrument_name="hindsight.http.duration",
aggregation=ExplicitBucketHistogramAggregation(boundaries=HTTP_DURATION_BUCKETS),
)
# Create meter provider with Prometheus exporter and custom views
provider = MeterProvider(
resource=resource, metric_readers=[prometheus_reader], views=[duration_view, llm_duration_view]
resource=resource,
metric_readers=[prometheus_reader],
views=[duration_view, llm_duration_view, http_duration_view],
)
# Set the global meter provider
@@ -167,6 +197,15 @@ class MetricsCollectorBase:
"""
raise NotImplementedError
@contextmanager
def record_http_request(self, method: str, endpoint: str, status_code_getter: Callable[[], int]):
"""Context manager to record HTTP request metrics."""
raise NotImplementedError
def set_db_pool(self, pool: "asyncpg.Pool"):
"""Set the database pool for metrics collection."""
pass
class NoOpMetricsCollector(MetricsCollectorBase):
"""No-op metrics collector that does nothing. Used when metrics are disabled."""
@@ -196,6 +235,11 @@ class NoOpMetricsCollector(MetricsCollectorBase):
"""No-op LLM call recording."""
pass
@contextmanager
def record_http_request(self, method: str, endpoint: str, status_code_getter: Callable[[], int]):
"""No-op HTTP request recording."""
yield
class MetricsCollector(MetricsCollectorBase):
"""
@@ -238,6 +282,27 @@ class MetricsCollector(MetricsCollectorBase):
name="hindsight.llm.calls.total", description="Total number of LLM API calls", unit="calls"
)
# HTTP request metrics
self.http_request_duration = self.meter.create_histogram(
name="hindsight.http.duration", description="Duration of HTTP requests in seconds", unit="s"
)
self.http_requests_total = self.meter.create_counter(
name="hindsight.http.requests.total", description="Total number of HTTP requests", unit="requests"
)
self.http_requests_in_progress = self.meter.create_up_down_counter(
name="hindsight.http.requests.in_progress",
description="Number of HTTP requests in progress",
unit="requests",
)
# Process metrics (observable gauges - collected on scrape)
self._setup_process_metrics()
# DB pool metrics holder (set via set_db_pool)
self._db_pool: "asyncpg.Pool | None" = None
@contextmanager
def record_operation(
self,
@@ -267,6 +332,7 @@ class MetricsCollector(MetricsCollectorBase):
"operation": operation,
"bank_id": bank_id,
"source": source,
"tenant": _get_tenant(),
}
if budget:
attributes["budget"] = budget
@@ -317,6 +383,7 @@ class MetricsCollector(MetricsCollectorBase):
"model": model,
"scope": scope,
"success": str(success).lower(),
"tenant": _get_tenant(),
}
# Record duration
@@ -340,6 +407,200 @@ class MetricsCollector(MetricsCollectorBase):
}
self.llm_tokens_output.add(output_tokens, output_attributes)
@contextmanager
def record_http_request(self, method: str, endpoint: str, status_code_getter: Callable[[], int]):
"""
Context manager to record HTTP request metrics.
Usage:
status_code = [200] # Use list for mutability
with metrics.record_http_request("GET", "/api/banks", lambda: status_code[0]):
# ... handle request
status_code[0] = response.status_code
Args:
method: HTTP method (GET, POST, etc.)
endpoint: Request endpoint path
status_code_getter: Callable that returns the status code after request completes
"""
start_time = time.time()
base_attributes = {"method": method, "endpoint": endpoint}
# Track in-progress
self.http_requests_in_progress.add(1, base_attributes)
try:
yield
finally:
duration = time.time() - start_time
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
self.http_request_duration.record(duration, attributes)
self.http_requests_total.add(1, attributes)
# Decrement in-progress
self.http_requests_in_progress.add(-1, base_attributes)
def _setup_process_metrics(self):
"""Set up observable gauges for process metrics."""
def get_cpu_times(_options):
"""Get process CPU times."""
try:
rusage = resource.getrusage(resource.RUSAGE_SELF)
yield metrics.Observation(rusage.ru_utime, {"type": "user"})
yield metrics.Observation(rusage.ru_stime, {"type": "system"})
except Exception:
pass
def get_memory_usage(_options):
"""Get process memory usage in bytes."""
try:
rusage = resource.getrusage(resource.RUSAGE_SELF)
# ru_maxrss is in kilobytes on Linux, bytes on macOS
max_rss = rusage.ru_maxrss
if os.uname().sysname == "Linux":
max_rss *= 1024 # Convert KB to bytes
yield metrics.Observation(max_rss, {"type": "rss_max"})
except Exception:
pass
def get_open_file_descriptors(_options):
"""Get number of open file descriptors."""
try:
# Try to count open FDs by checking /proc on Linux
if os.path.exists("/proc/self/fd"):
count = len(os.listdir("/proc/self/fd"))
yield metrics.Observation(count)
else:
# Fallback: use resource limits
soft, hard = resource.getrlimit(resource.RLIMIT_NOFILE)
yield metrics.Observation(soft, {"limit": "soft"})
except Exception:
pass
def get_thread_count(_options):
"""Get number of active threads."""
try:
yield metrics.Observation(threading.active_count())
except Exception:
pass
# Create observable gauges
self.meter.create_observable_gauge(
name="hindsight.process.cpu.seconds",
callbacks=[get_cpu_times],
description="Process CPU time in seconds",
unit="s",
)
self.meter.create_observable_gauge(
name="hindsight.process.memory.bytes",
callbacks=[get_memory_usage],
description="Process memory usage in bytes",
unit="By",
)
self.meter.create_observable_gauge(
name="hindsight.process.open_fds",
callbacks=[get_open_file_descriptors],
description="Number of open file descriptors",
unit="{fds}",
)
self.meter.create_observable_gauge(
name="hindsight.process.threads",
callbacks=[get_thread_count],
description="Number of active threads",
unit="{threads}",
)
def set_db_pool(self, pool: "asyncpg.Pool"):
"""
Set the database pool for metrics collection.
Args:
pool: asyncpg connection pool instance
"""
self._db_pool = pool
self._setup_db_pool_metrics()
def _setup_db_pool_metrics(self):
"""Set up observable gauges for database pool metrics."""
def get_pool_size(_options):
"""Get current pool size."""
if self._db_pool is not None:
try:
yield metrics.Observation(self._db_pool.get_size())
except Exception:
pass
def get_pool_free_size(_options):
"""Get number of free connections in pool."""
if self._db_pool is not None:
try:
yield metrics.Observation(self._db_pool.get_idle_size())
except Exception:
pass
def get_pool_min_size(_options):
"""Get pool minimum size."""
if self._db_pool is not None:
try:
yield metrics.Observation(self._db_pool.get_min_size())
except Exception:
pass
def get_pool_max_size(_options):
"""Get pool maximum size."""
if self._db_pool is not None:
try:
yield metrics.Observation(self._db_pool.get_max_size())
except Exception:
pass
# Create observable gauges for pool metrics
self.meter.create_observable_gauge(
name="hindsight.db.pool.size",
callbacks=[get_pool_size],
description="Current number of connections in the pool",
unit="{connections}",
)
self.meter.create_observable_gauge(
name="hindsight.db.pool.idle",
callbacks=[get_pool_free_size],
description="Number of idle connections in the pool",
unit="{connections}",
)
self.meter.create_observable_gauge(
name="hindsight.db.pool.min",
callbacks=[get_pool_min_size],
description="Minimum pool size",
unit="{connections}",
)
self.meter.create_observable_gauge(
name="hindsight.db.pool.max",
callbacks=[get_pool_max_size],
description="Maximum pool size",
unit="{connections}",
)
# Global metrics collector instance (defaults to no-op)
_metrics_collector: MetricsCollectorBase = NoOpMetricsCollector()
+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}'")
+39 -2
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,12 +32,42 @@ 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
_memory = MemoryEngine()
# Note: run_migrations=True by default, but migrations are idempotent so safe with workers
_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(memory=_memory, http_api_enabled=True, mcp_api_enabled=config.mcp_enabled, mcp_mount_path="/mcp")
app = create_app(
memory=_memory,
http_api_enabled=True,
mcp_api_enabled=config.mcp_enabled,
mcp_mount_path="/mcp",
initialize_memory=True,
)
if __name__ == "__main__":
+3 -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"
@@ -37,10 +37,12 @@ dependencies = [
"anthropic>=0.40.0",
"typer>=0.9.0",
"cohere>=5.0.0",
"flashrank>=0.2.0",
# Local ML models for embeddings/reranking - can be excluded in Docker with INCLUDE_LOCAL_MODELS=false
"sentence-transformers>=3.0.0,<3.3.0",
"transformers>=4.30.0,<4.46.0",
"torch>=2.0.0",
"uvloop>=0.22.1",
]
[project.optional-dependencies]
@@ -514,14 +514,15 @@ class TestCohereCrossEncoder:
"""Test that Cohere cross-encoder initializes correctly."""
assert cohere_cross_encoder.provider_name == "cohere"
def test_cohere_cross_encoder_predict(self, cohere_cross_encoder):
@pytest.mark.asyncio
async def test_cohere_cross_encoder_predict(self, cohere_cross_encoder):
"""Test that Cohere cross-encoder can score pairs."""
pairs = [
("What is the capital of France?", "Paris is the capital of France."),
("What is the capital of France?", "The Eiffel Tower is in Paris."),
("What is the capital of France?", "Python is a programming language."),
]
scores = cohere_cross_encoder.predict(pairs)
scores = await cohere_cross_encoder.predict(pairs)
assert len(scores) == 3
assert all(isinstance(s, float) for s in scores)
@@ -250,11 +250,34 @@ async def test_full_api_workflow(api_client, test_bank_id):
# 8. Test Entity Endpoints
# ================================================================
# List entities
# List entities with pagination
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/entities")
assert response.status_code == 200
entities_data = response.json()
assert "items" in entities_data
assert "total" in entities_data
assert "limit" in entities_data
assert "offset" in entities_data
assert entities_data["offset"] == 0
assert entities_data["limit"] == 100 # default limit
# Test pagination with custom limit and offset
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/entities?limit=5&offset=0")
assert response.status_code == 200
paginated_data = response.json()
assert paginated_data["limit"] == 5
assert paginated_data["offset"] == 0
assert len(paginated_data["items"]) <= 5
# Test offset
if entities_data["total"] > 1:
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/entities?limit=1&offset=1")
assert response.status_code == 200
offset_data = response.json()
assert offset_data["offset"] == 1
# With offset=1, we should get different entity than first one (if there are multiple)
if len(offset_data["items"]) > 0 and len(entities_data["items"]) > 1:
assert offset_data["items"][0]["id"] != entities_data["items"][0]["id"]
# Get specific entity if any exist
if len(entities_data['items']) > 0:
+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()
+8 -8
View File
@@ -64,12 +64,12 @@ class TestMetricsCollector:
def mock_meter(self):
"""Create a mock meter for testing."""
meter = MagicMock()
# Create separate mocks for each histogram (operation_duration, llm_duration)
histogram_mocks = [MagicMock(), MagicMock()]
# Create separate mocks for each histogram (operation_duration, llm_duration, http_request_duration)
histogram_mocks = [MagicMock(), MagicMock(), MagicMock()]
meter.create_histogram.side_effect = histogram_mocks
# Create separate mocks for each counter
# (operation_total, llm_tokens_input, llm_tokens_output, llm_calls_total)
counter_mocks = [MagicMock() for _ in range(4)]
# (operation_total, llm_tokens_input, llm_tokens_output, llm_calls_total, http_requests_total)
counter_mocks = [MagicMock() for _ in range(5)]
meter.create_counter.side_effect = counter_mocks
return meter
@@ -257,12 +257,12 @@ class TestLLMMetrics:
def mock_meter(self):
"""Create a mock meter for testing."""
meter = MagicMock()
# Create separate mocks for each histogram (operation_duration, llm_duration)
histogram_mocks = [MagicMock(), MagicMock()]
# Create separate mocks for each histogram (operation_duration, llm_duration, http_request_duration)
histogram_mocks = [MagicMock(), MagicMock(), MagicMock()]
meter.create_histogram.side_effect = histogram_mocks
# Create separate mocks for each counter
# (operation_total, llm_tokens_input, llm_tokens_output, llm_calls_total)
counter_mocks = [MagicMock() for _ in range(4)]
# (operation_total, llm_tokens_input, llm_tokens_output, llm_calls_total, http_requests_total)
counter_mocks = [MagicMock() for _ in range(5)]
meter.create_counter.side_effect = counter_mocks
return meter
+819
View File
@@ -0,0 +1,819 @@
"""
Tests for MPFP (Meta-Path Forward Push) graph retrieval.
Tests cover:
1. EdgeCache - lazy caching behavior
2. mpfp_traverse_async - core traversal algorithm
3. load_edges_for_frontier - lazy edge loading
4. rrf_fusion - result fusion
5. MPFPGraphRetriever - full integration
"""
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from datetime import datetime, timezone
from hindsight_api.engine.search.mpfp_retrieval import (
EdgeCache,
EdgeTarget,
MPFPConfig,
MPFPGraphRetriever,
PatternResult,
SeedNode,
load_all_edges_for_frontier,
mpfp_traverse_async,
rrf_fusion,
)
from hindsight_api.engine.search.types import RetrievalResult
class TestEdgeCache:
"""Tests for the EdgeCache lazy loading cache."""
def test_empty_cache_returns_empty_neighbors(self):
"""Empty cache should return empty list for any node."""
cache = EdgeCache()
neighbors = cache.get_neighbors("semantic", "node-1")
assert neighbors == []
def test_is_fully_loaded_false_for_uncached(self):
"""is_fully_loaded should return False for nodes not yet loaded."""
cache = EdgeCache()
assert cache.is_fully_loaded("node-1") is False
def test_add_all_edges_marks_as_fully_loaded(self):
"""Adding edges should mark nodes as fully loaded."""
cache = EdgeCache()
edges_by_type = {
"semantic": {"node-1": [EdgeTarget("node-2", 0.8), EdgeTarget("node-3", 0.6)]},
}
cache.add_all_edges(edges_by_type, ["node-1", "node-4"]) # node-4 has no edges
assert cache.is_fully_loaded("node-1") is True
assert cache.is_fully_loaded("node-4") is True # Marked even with no edges
assert cache.is_fully_loaded("node-2") is False # Target, not source
def test_get_neighbors_returns_added_edges(self):
"""get_neighbors should return edges after add_all_edges."""
cache = EdgeCache()
edges_by_type = {
"semantic": {"node-1": [EdgeTarget("node-2", 0.8), EdgeTarget("node-3", 0.6)]},
}
cache.add_all_edges(edges_by_type, ["node-1"])
neighbors = cache.get_neighbors("semantic", "node-1")
assert len(neighbors) == 2
assert neighbors[0].node_id == "node-2"
assert neighbors[0].weight == 0.8
def test_get_uncached_filters_loaded_nodes(self):
"""get_uncached should only return nodes not yet fully loaded."""
cache = EdgeCache()
# Load some nodes (all edge types)
cache.add_all_edges({"semantic": {"node-1": []}}, ["node-1", "node-2"])
# Check uncached
uncached = cache.get_uncached(["node-1", "node-2", "node-3", "node-4"])
assert set(uncached) == {"node-3", "node-4"}
def test_get_normalized_neighbors_normalizes_weights(self):
"""get_normalized_neighbors should normalize weights to sum to 1."""
cache = EdgeCache()
edges_by_type = {
"semantic": {
"node-1": [
EdgeTarget("node-2", 0.8),
EdgeTarget("node-3", 0.4),
EdgeTarget("node-4", 0.2),
],
},
}
cache.add_all_edges(edges_by_type, ["node-1"])
# Get top 2, normalized
neighbors = cache.get_normalized_neighbors("semantic", "node-1", top_k=2)
assert len(neighbors) == 2
# Weights should sum to 1
total = sum(n.weight for n in neighbors)
assert abs(total - 1.0) < 0.001
# node-2 should have higher normalized weight than node-3
assert neighbors[0].node_id == "node-2"
assert neighbors[1].node_id == "node-3"
# Original: 0.8 and 0.4, so normalized: 0.8/1.2 and 0.4/1.2
assert abs(neighbors[0].weight - 0.8 / 1.2) < 0.001
assert abs(neighbors[1].weight - 0.4 / 1.2) < 0.001
def test_different_edge_types_are_separate(self):
"""Different edge types should be stored separately."""
cache = EdgeCache()
edges_by_type = {
"semantic": {"node-1": [EdgeTarget("node-2", 0.8)]},
"temporal": {"node-1": [EdgeTarget("node-3", 0.5)]},
}
cache.add_all_edges(edges_by_type, ["node-1"])
semantic_neighbors = cache.get_neighbors("semantic", "node-1")
temporal_neighbors = cache.get_neighbors("temporal", "node-1")
assert len(semantic_neighbors) == 1
assert semantic_neighbors[0].node_id == "node-2"
assert len(temporal_neighbors) == 1
assert temporal_neighbors[0].node_id == "node-3"
class TestRRFFusion:
"""Tests for RRF (Reciprocal Rank Fusion)."""
def test_empty_results(self):
"""Empty results should return empty fusion."""
fused = rrf_fusion([])
assert fused == []
def test_single_pattern_ranking(self):
"""Single pattern should preserve ranking order."""
result = PatternResult(
pattern=["semantic"],
scores={"node-1": 0.9, "node-2": 0.7, "node-3": 0.5},
)
fused = rrf_fusion([result], top_k=3)
assert len(fused) == 3
# node-1 should be first (highest score)
assert fused[0][0] == "node-1"
assert fused[1][0] == "node-2"
assert fused[2][0] == "node-3"
def test_multiple_patterns_boost_common_nodes(self):
"""Nodes appearing in multiple patterns should get boosted."""
result1 = PatternResult(
pattern=["semantic", "semantic"],
scores={"node-1": 0.9, "node-2": 0.7},
)
result2 = PatternResult(
pattern=["entity", "temporal"],
scores={"node-1": 0.8, "node-3": 0.6}, # node-1 in both
)
fused = rrf_fusion([result1, result2], top_k=3)
# node-1 should be first (appears in both patterns)
assert fused[0][0] == "node-1"
# Its score should be higher than others
assert fused[0][1] > fused[1][1]
def test_top_k_limits_results(self):
"""top_k should limit the number of results."""
result = PatternResult(
pattern=["semantic"],
scores={f"node-{i}": 1.0 / (i + 1) for i in range(10)},
)
fused = rrf_fusion([result], top_k=3)
assert len(fused) == 3
def test_empty_pattern_scores_ignored(self):
"""Patterns with empty scores should be ignored."""
result1 = PatternResult(pattern=["semantic"], scores={})
result2 = PatternResult(
pattern=["entity"],
scores={"node-1": 0.5},
)
fused = rrf_fusion([result1, result2], top_k=3)
assert len(fused) == 1
assert fused[0][0] == "node-1"
class TestMPFPTraverseAsync:
"""Tests for the async MPFP traversal algorithm."""
@pytest.mark.asyncio
async def test_empty_seeds_returns_empty(self):
"""Empty seeds should return empty result."""
cache = EdgeCache()
config = MPFPConfig()
result = await mpfp_traverse_async(
pool=None, # Not used when no seeds
seeds=[],
pattern=["semantic"],
config=config,
cache=cache,
)
assert result.scores == {}
@pytest.mark.asyncio
async def test_single_hop_no_edges(self):
"""Single hop with no edges should deposit mass at seeds."""
cache = EdgeCache()
config = MPFPConfig(alpha=0.15, threshold=1e-6)
# Pre-populate cache with empty edges for seed (marks as fully loaded)
cache.add_all_edges({}, ["seed-1"])
seeds = [SeedNode("seed-1", 1.0)]
with patch(
"hindsight_api.engine.search.mpfp_retrieval.load_all_edges_for_frontier",
new_callable=AsyncMock,
return_value={},
):
result = await mpfp_traverse_async(
pool=MagicMock(),
seeds=seeds,
pattern=["semantic"],
config=config,
cache=cache,
)
# Seed should have alpha portion of its mass
assert "seed-1" in result.scores
assert result.scores["seed-1"] == pytest.approx(config.alpha, rel=0.01)
@pytest.mark.asyncio
async def test_single_hop_with_edges(self):
"""Single hop should spread mass to neighbors."""
cache = EdgeCache()
config = MPFPConfig(alpha=0.15, threshold=1e-6, top_k_neighbors=10)
seeds = [SeedNode("seed-1", 1.0)]
# Pre-populate cache with seed edges (mimics pre-warming in retrieve())
cache.add_all_edges(
{
"semantic": {
"seed-1": [
EdgeTarget("neighbor-1", 0.8),
EdgeTarget("neighbor-2", 0.4),
]
}
},
["seed-1"],
)
# Mock for loading neighbor edges (after hop 0)
async def mock_load_all_edges(pool, node_ids, top_k=20):
return {}
with patch(
"hindsight_api.engine.search.mpfp_retrieval.load_all_edges_for_frontier",
side_effect=mock_load_all_edges,
):
result = await mpfp_traverse_async(
pool=MagicMock(),
seeds=seeds,
pattern=["semantic"],
config=config,
cache=cache,
)
# Seed keeps alpha portion
assert "seed-1" in result.scores
assert result.scores["seed-1"] == pytest.approx(config.alpha, rel=0.01)
# Neighbors get remaining mass (normalized)
assert "neighbor-1" in result.scores
assert "neighbor-2" in result.scores
# neighbor-1 should get more (higher weight)
assert result.scores["neighbor-1"] > result.scores["neighbor-2"]
@pytest.mark.asyncio
async def test_two_hops(self):
"""Two-hop pattern should traverse through neighbors."""
cache = EdgeCache()
config = MPFPConfig(alpha=0.15, threshold=1e-6, top_k_neighbors=10)
seeds = [SeedNode("seed-1", 1.0)]
# Pre-populate cache with seed edges (mimics pre-warming in retrieve())
cache.add_all_edges(
{"semantic": {"seed-1": [EdgeTarget("hop1-node", 1.0)]}},
["seed-1"],
)
# Mock edge loading for hop 1 nodes
async def mock_load_all_edges(pool, node_ids, top_k=20):
edges: dict[str, dict[str, list[EdgeTarget]]] = {"semantic": {}}
if "hop1-node" in node_ids:
edges["semantic"]["hop1-node"] = [EdgeTarget("hop2-node", 1.0)]
return edges
with patch(
"hindsight_api.engine.search.mpfp_retrieval.load_all_edges_for_frontier",
side_effect=mock_load_all_edges,
):
result = await mpfp_traverse_async(
pool=MagicMock(),
seeds=seeds,
pattern=["semantic", "semantic"], # Two hops
config=config,
cache=cache,
)
# Should have scores for all three nodes
assert "seed-1" in result.scores
assert "hop1-node" in result.scores
assert "hop2-node" in result.scores
@pytest.mark.asyncio
async def test_cache_reuse(self):
"""Cache should prevent redundant edge loading for already-cached nodes."""
cache = EdgeCache()
config = MPFPConfig(alpha=0.15, threshold=1e-6)
# Pre-load cache (marks seed-1 AND neighbor-1 as fully loaded)
# neighbor-1 is also cached because after hop 0, the frontier contains neighbor-1
# and the algorithm tries to pre-warm edges for the next hop
cache.add_all_edges(
{"semantic": {"seed-1": [EdgeTarget("neighbor-1", 1.0)], "neighbor-1": []}},
["seed-1", "neighbor-1"],
)
seeds = [SeedNode("seed-1", 1.0)]
load_mock = AsyncMock(return_value={})
with patch(
"hindsight_api.engine.search.mpfp_retrieval.load_all_edges_for_frontier",
load_mock,
):
await mpfp_traverse_async(
pool=MagicMock(),
seeds=seeds,
pattern=["semantic"],
config=config,
cache=cache,
)
# Should not call load_all_edges_for_frontier since all nodes are already cached
load_mock.assert_not_called()
class TestMPFPGraphRetriever:
"""Tests for the MPFPGraphRetriever class."""
def test_name_is_mpfp(self):
"""Retriever name should be 'mpfp'."""
retriever = MPFPGraphRetriever()
assert retriever.name == "mpfp"
def test_default_config(self):
"""Default config should have expected patterns."""
# Use explicit config to avoid global config dependency
config = MPFPConfig()
retriever = MPFPGraphRetriever(config=config)
assert len(retriever.config.patterns_semantic) > 0
assert len(retriever.config.patterns_temporal) > 0
assert retriever.config.alpha == 0.15
assert retriever.config.top_k_neighbors == 20
def test_custom_config(self):
"""Custom config should be used."""
config = MPFPConfig(alpha=0.3, top_k_neighbors=10)
retriever = MPFPGraphRetriever(config=config)
assert retriever.config.alpha == 0.3
assert retriever.config.top_k_neighbors == 10
def test_convert_seeds_from_retrieval_results(self):
"""_convert_seeds should extract scores from RetrievalResult."""
retriever = MPFPGraphRetriever()
results = [
RetrievalResult(id="id-1", text="text1", fact_type="world", similarity=0.9),
RetrievalResult(id="id-2", text="text2", fact_type="world", similarity=0.7),
]
seeds = retriever._convert_seeds(results, "similarity")
assert len(seeds) == 2
assert seeds[0].node_id == "id-1"
assert seeds[0].score == 0.9
assert seeds[1].node_id == "id-2"
assert seeds[1].score == 0.7
def test_convert_seeds_empty(self):
"""_convert_seeds should handle empty/None input."""
retriever = MPFPGraphRetriever()
assert retriever._convert_seeds(None, "similarity") == []
assert retriever._convert_seeds([], "similarity") == []
@pytest.mark.asyncio
async def test_retrieve_no_seeds_returns_empty(self):
"""Retrieve with no seeds should return empty results."""
# Use explicit config to avoid global config dependency
config = MPFPConfig()
retriever = MPFPGraphRetriever(config=config)
# Mock _find_semantic_seeds to return empty
with patch.object(retriever, "_find_semantic_seeds", new_callable=AsyncMock, return_value=[]):
results, timings = await retriever.retrieve(
pool=MagicMock(),
query_embedding_str="[0.1, 0.2]",
bank_id="test",
fact_type="world",
budget=10,
)
assert results == []
assert timings is not None
assert timings.pattern_count == 0
@pytest.mark.asyncio
async def test_retrieve_with_semantic_seeds(self):
"""Retrieve with semantic seeds should run patterns and return results."""
# Use explicit config to avoid global config dependency
config = MPFPConfig()
retriever = MPFPGraphRetriever(config=config)
semantic_seeds = [
RetrievalResult(id="seed-1", text="seed text", fact_type="world", similarity=0.9),
]
# Mock the internal functions
# mpfp_traverse_hop_synchronized returns a list of PatternResult (one per pattern)
async def mock_traverse(*args, **kwargs):
return [PatternResult(pattern=["semantic"], scores={"seed-1": 0.5, "result-1": 0.3})]
async def mock_fetch(pool, node_ids, fact_type):
return [
RetrievalResult(id="seed-1", text="seed text", fact_type="world"),
RetrievalResult(id="result-1", text="result text", fact_type="world"),
]
with (
patch(
"hindsight_api.engine.search.mpfp_retrieval.mpfp_traverse_hop_synchronized",
side_effect=mock_traverse,
),
patch(
"hindsight_api.engine.search.mpfp_retrieval.fetch_memory_units_by_ids",
side_effect=mock_fetch,
),
patch(
"hindsight_api.engine.search.mpfp_retrieval.load_all_edges_for_frontier",
new_callable=AsyncMock,
return_value={},
),
):
results, timings = await retriever.retrieve(
pool=MagicMock(),
query_embedding_str="[0.1, 0.2]",
bank_id="test",
fact_type="world",
budget=10,
semantic_seeds=semantic_seeds,
)
assert len(results) == 2
assert timings is not None
assert timings.pattern_count > 0
@pytest.mark.asyncio
async def test_mpfp_integration(memory, request_context):
"""Integration test: MPFP retrieval with real database."""
bank_id = f"test_mpfp_{datetime.now(timezone.utc).timestamp()}"
try:
# Store memories with entity relationships
await memory.retain_async(
bank_id=bank_id,
content="Alice works at TechCorp as a software engineer",
context="employee info",
request_context=request_context,
)
await memory.retain_async(
bank_id=bank_id,
content="TechCorp is located in San Francisco",
context="company info",
request_context=request_context,
)
await memory.retain_async(
bank_id=bank_id,
content="Bob is Alice's manager at TechCorp",
context="employee info",
request_context=request_context,
)
await memory.retain_async(
bank_id=bank_id,
content="San Francisco has many tech companies",
context="city info",
request_context=request_context,
)
# Query should find related facts via graph traversal
from hindsight_api.engine.memory_engine import Budget
result = await memory.recall_async(
bank_id=bank_id,
query="Tell me about Alice",
fact_type=["world"],
budget=Budget.MID,
max_tokens=2048,
request_context=request_context,
)
# Should return results
assert result.results is not None
assert len(result.results) > 0
# Should find Alice-related facts
fact_texts = [f.text for f in result.results]
alice_facts = [t for t in fact_texts if "Alice" in t or "TechCorp" in t]
assert len(alice_facts) > 0, f"Should find Alice-related facts, got: {fact_texts}"
print(f"\n✓ MPFP integration test passed! Found {len(result.results)} facts")
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_mpfp_lazy_loading_efficiency(memory, request_context):
"""Test that MPFP loads edges lazily, not upfront."""
bank_id = f"test_mpfp_lazy_{datetime.now(timezone.utc).timestamp()}"
try:
# Store many memories to create a larger graph
for i in range(20):
await memory.retain_async(
bank_id=bank_id,
content=f"Fact number {i} about topic {i % 5}",
context=f"context {i}",
request_context=request_context,
)
from hindsight_api.engine.memory_engine import Budget
# Query - MPFP should only load edges for relevant frontier nodes
result = await memory.recall_async(
bank_id=bank_id,
query="topic 0",
fact_type=["world"],
budget=Budget.LOW,
max_tokens=1024,
enable_trace=True,
request_context=request_context,
)
assert result.results is not None
# Check trace for timing info
if result.trace:
print(f"\n✓ MPFP lazy loading test passed!")
print(f" - Facts returned: {len(result.results)}")
finally:
await memory.delete_bank(bank_id, request_context=request_context)
# ============================================================================
# MPFP Performance Benchmark Tests
# ============================================================================
# These tests require an external database with a large memory bank to be useful.
# Set EXTERNAL_DATABASE_URL and BENCHMARK_BANK_ID environment variables to run.
# Example:
# EXTERNAL_DATABASE_URL=postgresql://user:pass@host:port/db \
# BENCHMARK_BANK_ID=load-test \
# pytest tests/test_mpfp_retrieval.py::test_mpfp_edge_loading_performance -v -s
import os
import asyncpg
EXTERNAL_DATABASE_URL = os.environ.get("EXTERNAL_DATABASE_URL")
BENCHMARK_BANK_ID = os.environ.get("BENCHMARK_BANK_ID", "load-test")
requires_external_db = pytest.mark.skipif(
EXTERNAL_DATABASE_URL is None,
reason="EXTERNAL_DATABASE_URL not set - skipping external DB benchmark",
)
@requires_external_db
@pytest.mark.asyncio
async def test_mpfp_edge_loading_performance():
"""
Benchmark MPFP edge loading performance.
This test measures the performance of the LATERAL query optimization
for loading edges in the MPFP graph traversal algorithm.
Set EXTERNAL_DATABASE_URL to point to a database with existing data.
Set BENCHMARK_BANK_ID to specify which bank to query (default: load-test).
Example usage:
EXTERNAL_DATABASE_URL=postgresql://hindsight:hindsight@localhost:5435/hindsight \
BENCHMARK_BANK_ID=load-test \
pytest tests/test_mpfp_retrieval.py::test_mpfp_edge_loading_performance -v -s
"""
import time
# Connect to external database
pool = await asyncpg.create_pool(EXTERNAL_DATABASE_URL, min_size=2, max_size=10)
try:
# Get some sample node IDs from the database
async with pool.acquire() as conn:
# First check how many links exist
stats = await conn.fetchrow("""
SELECT
count(*) as total_links,
count(DISTINCT from_unit_id) as unique_sources
FROM memory_links
""")
print(f"\n📊 Database Stats:")
print(f" Total links: {stats['total_links']:,}")
print(f" Unique sources: {stats['unique_sources']:,}")
# Get edge distribution by type
type_stats = await conn.fetch("""
SELECT link_type, count(*) as cnt,
round(avg(weight)::numeric, 3) as avg_weight
FROM memory_links
GROUP BY link_type
ORDER BY cnt DESC
""")
print(f"\n Edge distribution:")
for row in type_stats:
print(f" - {row['link_type']}: {row['cnt']:,} (avg_weight={row['avg_weight']})")
# Get sample frontier nodes (from memory_units in the benchmark bank)
# bank_id is the text primary key in banks table
frontier_rows = await conn.fetch("""
SELECT id FROM memory_units
WHERE bank_id = $1
LIMIT 100
""", BENCHMARK_BANK_ID)
if not frontier_rows:
pytest.skip(f"No memory units found for bank '{BENCHMARK_BANK_ID}'")
frontier_node_ids = [str(row['id']) for row in frontier_rows]
print(f"\n🎯 Testing with {len(frontier_node_ids)} frontier nodes from bank '{BENCHMARK_BANK_ID}'")
# Test 1: Original query approach (all edges, no per-type limit)
async with pool.acquire() as conn:
start = time.time()
original_rows = await conn.fetch("""
SELECT ml.from_unit_id, ml.to_unit_id, ml.link_type, ml.weight
FROM memory_links ml
WHERE ml.from_unit_id = ANY($1::uuid[])
AND ml.weight >= 0.1
ORDER BY ml.from_unit_id, ml.link_type, ml.weight DESC
""", frontier_node_ids)
original_time = time.time() - start
original_count = len(original_rows)
# Test 2: New LATERAL query approach (top-k per type)
async with pool.acquire() as conn:
start = time.time()
lateral_rows = await conn.fetch("""
WITH frontier(node_id) AS (SELECT unnest($1::uuid[]))
SELECT f.node_id as from_unit_id, lt.link_type, edges.to_unit_id, edges.weight
FROM frontier f
CROSS JOIN (VALUES ('semantic'), ('temporal'), ('entity'), ('causes'), ('caused_by')) AS lt(link_type)
CROSS JOIN LATERAL (
SELECT ml.to_unit_id, ml.weight
FROM memory_links ml
WHERE ml.from_unit_id = f.node_id
AND ml.link_type = lt.link_type
AND ml.weight >= 0.1
ORDER BY ml.weight DESC
LIMIT 20
) edges
""", frontier_node_ids)
lateral_time = time.time() - start
lateral_count = len(lateral_rows)
# Print results
print(f"\n⏱️ Performance Comparison ({len(frontier_node_ids)} nodes):")
print(f"\n Original (all edges):")
print(f" - Time: {original_time * 1000:.2f}ms")
print(f" - Rows: {original_count:,}")
print(f" - Rows/node: {original_count / len(frontier_node_ids):.1f}")
print(f"\n LATERAL (top-20 per type):")
print(f" - Time: {lateral_time * 1000:.2f}ms")
print(f" - Rows: {lateral_count:,}")
print(f" - Rows/node: {lateral_count / len(frontier_node_ids):.1f}")
speedup = original_time / lateral_time if lateral_time > 0 else float('inf')
reduction = (1 - lateral_count / original_count) * 100 if original_count > 0 else 0
print(f"\n 📈 Improvement:")
print(f" - Speedup: {speedup:.2f}x faster")
print(f" - Data reduction: {reduction:.1f}% fewer rows")
# Assert improvement (should be at least some improvement for large datasets)
if original_count > 1000:
# For large datasets, expect significant improvement
assert speedup >= 1.5, f"Expected at least 1.5x speedup, got {speedup:.2f}x"
assert reduction >= 30, f"Expected at least 30% data reduction, got {reduction:.1f}%"
print(f"\n✅ Performance test PASSED!")
else:
print(f"\n⚠️ Dataset too small ({original_count} rows) for meaningful performance comparison")
finally:
await pool.close()
@requires_external_db
@pytest.mark.asyncio
async def test_mpfp_full_retrieval_performance():
"""
Benchmark full MPFP retrieval including traversal and reranking.
This test measures end-to-end MPFP retrieval performance.
"""
import time
pool = await asyncpg.create_pool(EXTERNAL_DATABASE_URL, min_size=2, max_size=10)
try:
# Get a sample query embedding from an existing memory unit
async with pool.acquire() as conn:
# Check if bank exists
bank_exists = await conn.fetchval("""
SELECT 1 FROM banks WHERE bank_id = $1
""", BENCHMARK_BANK_ID)
if not bank_exists:
pytest.skip(f"Bank '{BENCHMARK_BANK_ID}' not found")
sample = await conn.fetchrow("""
SELECT embedding::text as embedding_str
FROM memory_units
WHERE bank_id = $1
AND embedding IS NOT NULL
LIMIT 1
""", BENCHMARK_BANK_ID)
if not sample:
pytest.skip("No memory units with embeddings found")
query_embedding_str = sample['embedding_str']
# Run MPFP retrieval
retriever = MPFPGraphRetriever()
print(f"\n🔍 Running MPFP retrieval benchmark on bank '{BENCHMARK_BANK_ID}'...")
# Warm-up run
await retriever.retrieve(
pool=pool,
query_embedding_str=query_embedding_str,
bank_id=BENCHMARK_BANK_ID,
fact_type="world",
budget=100,
query_text="test query",
)
# Timed runs
timings_list = []
for i in range(3):
start = time.time()
results, timings = await retriever.retrieve(
pool=pool,
query_embedding_str=query_embedding_str,
bank_id=BENCHMARK_BANK_ID,
fact_type="opinion",
budget=100,
query_text="What did I say about training models?",
)
elapsed = time.time() - start
timings_list.append((elapsed, timings, len(results)))
# Print results
print(f"\n⏱️ MPFP Retrieval Results (3 runs):")
for i, (elapsed, timings, count) in enumerate(timings_list):
print(f"\n Run {i + 1}:")
print(f" - Total: {elapsed * 1000:.2f}ms")
print(f" - Results: {count}")
if timings:
print(f" - Seeds: {timings.seeds_time * 1000:.2f}ms")
print(f" - Patterns: {timings.pattern_count}")
print(f" - Traverse: {timings.traverse * 1000:.2f}ms")
print(f" - Edge load: {timings.edge_load_time * 1000:.2f}ms")
print(f" - Edges: {timings.edge_count:,}")
print(f" - DB queries: {timings.db_queries}")
print(f" - Fusion: {timings.fusion * 1000:.2f}ms")
print(f" - Fetch: {timings.fetch * 1000:.2f}ms")
avg_time = sum(t[0] for t in timings_list) / len(timings_list)
print(f"\n 📊 Average: {avg_time * 1000:.2f}ms")
print(f"\n✅ MPFP retrieval benchmark complete!")
finally:
await pool.close()
+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")
@@ -0,0 +1,786 @@
"""
Tests for RemoteTEICrossEncoder (TEI reranker client).
Tests cover:
- Initialization and server connectivity
- Basic predict functionality
- Batch splitting
- Parallel request handling
- Backpressure/semaphore behavior
- Retry logic on transient errors
- Multiple queries handling
"""
import asyncio
import time
from unittest.mock import MagicMock, patch
import httpx
import pytest
from hindsight_api.engine.cross_encoder import RemoteTEICrossEncoder
class TestRemoteTEICrossEncoderInitialization:
"""Tests for TEI cross-encoder initialization."""
@pytest.mark.asyncio
async def test_initialize_success(self):
"""Test successful initialization with valid TEI server."""
async def mock_handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/info":
return httpx.Response(
200,
json={"model_id": "BAAI/bge-reranker-base", "version": "1.0"},
)
return httpx.Response(404)
transport = httpx.MockTransport(mock_handler)
with patch.object(httpx, "AsyncClient", return_value=httpx.AsyncClient(transport=transport)):
encoder = RemoteTEICrossEncoder(base_url="http://localhost:8080")
await encoder.initialize()
assert encoder._model_id == "BAAI/bge-reranker-base"
assert encoder._async_client is not None
@pytest.mark.asyncio
async def test_initialize_server_unreachable(self):
"""Test initialization fails when server is unreachable."""
async def mock_handler(request: httpx.Request) -> httpx.Response:
raise httpx.ConnectError("Connection refused")
transport = httpx.MockTransport(mock_handler)
with patch.object(httpx, "AsyncClient", return_value=httpx.AsyncClient(transport=transport)):
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
max_retries=1,
retry_delay=0.01,
)
with pytest.raises(RuntimeError, match="Failed to connect to TEI server"):
await encoder.initialize()
@pytest.mark.asyncio
async def test_initialize_idempotent(self):
"""Test that initialize() is idempotent."""
call_count = 0
async def mock_handler(request: httpx.Request) -> httpx.Response:
nonlocal call_count
if request.url.path == "/info":
call_count += 1
return httpx.Response(200, json={"model_id": "test-model"})
return httpx.Response(404)
transport = httpx.MockTransport(mock_handler)
with patch.object(httpx, "AsyncClient", return_value=httpx.AsyncClient(transport=transport)):
encoder = RemoteTEICrossEncoder(base_url="http://localhost:8080")
await encoder.initialize()
await encoder.initialize()
await encoder.initialize()
assert call_count == 1
def create_mock_async_client(handler):
"""Create a mock AsyncClient that uses the given handler for requests."""
class MockAsyncClient:
def __init__(self, **kwargs):
self.timeout = kwargs.get("timeout", 30.0)
async def __aenter__(self):
return self
async def __aexit__(self, *args):
pass
async def post(self, url, **kwargs):
return await handler("POST", url, **kwargs)
async def get(self, url, **kwargs):
return await handler("GET", url, **kwargs)
return MockAsyncClient()
class TestRemoteTEICrossEncoderPredict:
"""Tests for TEI cross-encoder predict functionality."""
@pytest.mark.asyncio
async def test_predict_not_initialized(self):
"""Test predict raises error when not initialized."""
encoder = RemoteTEICrossEncoder(base_url="http://localhost:8080")
with pytest.raises(RuntimeError, match="Reranker not initialized"):
await encoder.predict([("query", "doc")])
@pytest.mark.asyncio
async def test_predict_empty_pairs(self):
"""Test predict returns empty list for empty input."""
encoder = RemoteTEICrossEncoder(base_url="http://localhost:8080")
encoder._async_client = httpx.AsyncClient()
encoder._model_id = "test-model"
result = await encoder.predict([])
assert result == []
@pytest.mark.asyncio
async def test_predict_single_query(self):
"""Test predict with single query and multiple documents."""
rerank_calls = []
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
body = kwargs.get("json", {})
rerank_calls.append(body)
texts = body["texts"]
# Return scores in descending order with original indices
results = [{"index": i, "score": 1.0 - (i * 0.1)} for i in range(len(texts))]
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
response.json = MagicMock(return_value=results)
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
encoder = RemoteTEICrossEncoder(base_url="http://localhost:8080")
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
pairs = [
("What is Python?", "Python is a programming language."),
("What is Python?", "Python is a snake."),
("What is Python?", "Java is also a language."),
]
scores = await encoder.predict(pairs)
assert len(scores) == 3
assert len(rerank_calls) == 1
assert rerank_calls[0]["query"] == "What is Python?"
assert len(rerank_calls[0]["texts"]) == 3
# Scores should be mapped back correctly
assert scores[0] == 1.0
assert scores[1] == 0.9
assert scores[2] == pytest.approx(0.8, rel=0.01)
@pytest.mark.asyncio
async def test_predict_multiple_queries(self):
"""Test predict with multiple different queries."""
rerank_calls = []
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
body = kwargs.get("json", {})
rerank_calls.append(body)
texts = body["texts"]
results = [{"index": i, "score": 0.5 + (i * 0.1)} for i in range(len(texts))]
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
response.json = MagicMock(return_value=results)
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
encoder = RemoteTEICrossEncoder(base_url="http://localhost:8080")
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
pairs = [
("Query A", "Doc A1"),
("Query B", "Doc B1"),
("Query A", "Doc A2"),
("Query B", "Doc B2"),
]
scores = await encoder.predict(pairs)
assert len(scores) == 4
# Two queries = two rerank calls (run in parallel)
assert len(rerank_calls) == 2
class TestRemoteTEICrossEncoderBatching:
"""Tests for batch splitting behavior."""
@pytest.mark.asyncio
async def test_batch_splitting(self):
"""Test that large inputs are split into batches."""
rerank_calls = []
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
body = kwargs.get("json", {})
rerank_calls.append(body)
texts = body["texts"]
results = [{"index": i, "score": 0.5} for i in range(len(texts))]
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
response.json = MagicMock(return_value=results)
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
batch_size=3, # Small batch for testing
)
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
# 7 documents with same query, batch_size=3 -> 3 batches (3+3+1)
pairs = [("Query", f"Doc {i}") for i in range(7)]
scores = await encoder.predict(pairs)
assert len(scores) == 7
assert len(rerank_calls) == 3
# Check batch sizes
batch_sizes = sorted([len(call["texts"]) for call in rerank_calls])
assert batch_sizes == [1, 3, 3]
@pytest.mark.asyncio
async def test_score_mapping_across_batches(self):
"""Test that scores are correctly mapped back across batches."""
call_counter = [0]
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
body = kwargs.get("json", {})
batch_num = call_counter[0]
call_counter[0] += 1
texts = body["texts"]
# Each batch returns different scores to verify mapping
base_score = batch_num * 10
results = [{"index": i, "score": float(base_score + i)} for i in range(len(texts))]
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
response.json = MagicMock(return_value=results)
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
batch_size=3,
)
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
pairs = [("Query", f"Doc {i}") for i in range(7)]
scores = await encoder.predict(pairs)
assert len(scores) == 7
# All scores should be present (exact values depend on batch ordering)
assert all(isinstance(s, (int, float)) for s in scores)
class TestRemoteTEICrossEncoderParallelism:
"""Tests for parallel request handling and backpressure."""
@pytest.mark.asyncio
async def test_parallel_requests(self):
"""Test that requests are made in parallel."""
concurrent_count = [0]
max_concurrent_observed = [0]
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
concurrent_count[0] += 1
max_concurrent_observed[0] = max(max_concurrent_observed[0], concurrent_count[0])
await asyncio.sleep(0.03) # Simulate latency
concurrent_count[0] -= 1
body = kwargs.get("json", {})
texts = body["texts"]
results = [{"index": i, "score": 0.5} for i in range(len(texts))]
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
response.json = MagicMock(return_value=results)
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
batch_size=2,
max_concurrent=10, # High limit to allow parallelism
)
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
# 6 docs = 3 batches, should run in parallel
pairs = [("Query", f"Doc {i}") for i in range(6)]
start = time.time()
scores = await encoder.predict(pairs)
elapsed = time.time() - start
assert len(scores) == 6
# If parallel, 3 batches with 30ms each should take ~30ms, not 90ms
assert elapsed < 0.08, f"Requests should run in parallel, took {elapsed}s"
assert max_concurrent_observed[0] > 1, "Multiple requests should run concurrently"
@pytest.mark.asyncio
async def test_backpressure_semaphore(self):
"""Test that semaphore limits concurrent requests."""
concurrent_count = [0]
max_concurrent_observed = [0]
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
concurrent_count[0] += 1
max_concurrent_observed[0] = max(max_concurrent_observed[0], concurrent_count[0])
await asyncio.sleep(0.01) # Simulate latency
concurrent_count[0] -= 1
body = kwargs.get("json", {})
texts = body["texts"]
results = [{"index": i, "score": 0.5} for i in range(len(texts))]
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
response.json = MagicMock(return_value=results)
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
max_concurrent_limit = 2
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
batch_size=1, # 1 doc per batch to maximize requests
max_concurrent=max_concurrent_limit,
)
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
# 10 docs = 10 batches, but only 2 should run at a time
pairs = [("Query", f"Doc {i}") for i in range(10)]
scores = await encoder.predict(pairs)
assert len(scores) == 10
assert max_concurrent_observed[0] <= max_concurrent_limit, (
f"Semaphore should limit to {max_concurrent_limit}, observed {max_concurrent_observed[0]}"
)
class TestRemoteTEICrossEncoderRetry:
"""Tests for retry logic on transient errors."""
@pytest.mark.asyncio
async def test_retry_on_connect_error(self):
"""Test that connect errors trigger retries."""
attempt_count = [0]
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
attempt_count[0] += 1
if attempt_count[0] < 3:
raise httpx.ConnectError("Connection refused")
body = kwargs.get("json", {})
texts = body["texts"]
results = [{"index": i, "score": 0.5} for i in range(len(texts))]
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
response.json = MagicMock(return_value=results)
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
max_retries=3,
retry_delay=0.01,
)
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
pairs = [("Query", "Doc 1")]
scores = await encoder.predict(pairs)
assert len(scores) == 1
assert attempt_count[0] == 3 # 2 failures + 1 success
@pytest.mark.asyncio
async def test_retry_on_server_error(self):
"""Test that 5xx errors trigger retries."""
attempt_count = [0]
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
attempt_count[0] += 1
if attempt_count[0] < 2:
response = MagicMock()
response.status_code = 503
def raise_for_status():
raise httpx.HTTPStatusError(
"Service unavailable",
request=MagicMock(),
response=response,
)
response.raise_for_status = raise_for_status
return response
body = kwargs.get("json", {})
texts = body["texts"]
results = [{"index": i, "score": 0.5} for i in range(len(texts))]
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
response.json = MagicMock(return_value=results)
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
max_retries=3,
retry_delay=0.01,
)
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
pairs = [("Query", "Doc 1")]
scores = await encoder.predict(pairs)
assert len(scores) == 1
assert attempt_count[0] == 2
@pytest.mark.asyncio
async def test_no_retry_on_client_error(self):
"""Test that 4xx errors do not trigger retries."""
attempt_count = [0]
async def mock_handler(method, url, **kwargs):
if "/rerank" in url:
attempt_count[0] += 1
response = MagicMock()
response.status_code = 400
def raise_for_status():
raise httpx.HTTPStatusError(
"Bad request",
request=MagicMock(),
response=response,
)
response.raise_for_status = raise_for_status
return response
raise httpx.HTTPStatusError("Not found", request=MagicMock(), response=MagicMock())
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
max_retries=3,
retry_delay=0.01,
)
encoder._async_client = create_mock_async_client(mock_handler)
encoder._model_id = "test-model"
pairs = [("Query", "Doc 1")]
with pytest.raises(RuntimeError, match="TEI rerank request failed"):
await encoder.predict(pairs)
assert attempt_count[0] == 1 # No retries for 4xx
class TestRemoteTEICrossEncoderConfig:
"""Tests for configuration from environment variables."""
def test_default_values(self):
"""Test default configuration values."""
encoder = RemoteTEICrossEncoder(base_url="http://localhost:8080")
assert encoder.batch_size == 128
assert encoder.max_concurrent == 8
assert encoder.timeout == 30.0
assert encoder.max_retries == 3
def test_custom_values(self):
"""Test custom configuration values."""
encoder = RemoteTEICrossEncoder(
base_url="http://localhost:8080",
batch_size=64,
max_concurrent=4,
timeout=60.0,
max_retries=5,
retry_delay=1.0,
)
assert encoder.batch_size == 64
assert encoder.max_concurrent == 4
assert encoder.timeout == 60.0
assert encoder.max_retries == 5
assert encoder.retry_delay == 1.0
def test_create_from_env(self):
"""Test creating encoder from environment variables."""
import os
from hindsight_api.engine.cross_encoder import create_cross_encoder_from_env
with patch.dict(
os.environ,
{
"HINDSIGHT_API_RERANKER_PROVIDER": "tei",
"HINDSIGHT_API_RERANKER_TEI_URL": "http://test:9000",
"HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE": "256",
"HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT": "16",
},
):
encoder = create_cross_encoder_from_env()
assert isinstance(encoder, RemoteTEICrossEncoder)
assert encoder.base_url == "http://test:9000"
assert encoder.batch_size == 256
assert encoder.max_concurrent == 16
# ============================================================================
# TEI Reranker Performance Benchmark Tests
# ============================================================================
# These tests require a running TEI server to measure actual performance.
# Set TEI_RERANKER_URL environment variable to run.
# Example:
# TEI_RERANKER_URL=http://localhost:8000 \
# pytest tests/test_tei_cross_encoder.py::test_tei_reranker_performance -v -s -n0
import os
TEI_RERANKER_URL = os.environ.get("TEI_RERANKER_URL")
requires_tei_server = pytest.mark.skipif(
TEI_RERANKER_URL is None,
reason="TEI_RERANKER_URL not set - skipping TEI performance benchmark",
)
@requires_tei_server
@pytest.mark.asyncio
async def test_tei_reranker_performance():
"""
Benchmark TEI reranker performance with different configurations.
This test measures latency for different batch sizes and concurrency levels
to find the optimal configuration for your TEI server.
Example usage:
TEI_RERANKER_URL=http://localhost:8000 \
pytest tests/test_tei_cross_encoder.py::test_tei_reranker_performance -v -s -n0
"""
import httpx
# Get server info
async with httpx.AsyncClient() as client:
response = await client.get(f"{TEI_RERANKER_URL}/info")
info = response.json()
print(f"\n📊 TEI Server Info:")
print(f" URL: {TEI_RERANKER_URL}")
print(f" Model: {info.get('model_id', 'unknown')}")
if "reranker_model" in info:
print(f" Reranker Model: {info['reranker_model']}")
# Generate test data (800 pairs to simulate real workload)
num_pairs = 800
query = "What did I say about training machine learning models and artificial intelligence?"
test_pairs = [
(query, f"Document {i} about machine learning, neural networks, and AI training techniques.")
for i in range(num_pairs)
]
# Test configurations: (batch_size, max_concurrent)
configs = [
(128, 8), # Default
(256, 4), # Larger batches, fewer concurrent
(256, 8), # Larger batches, same concurrent
(512, 2), # Very large batches, few concurrent
(512, 4), # Very large batches, moderate concurrent
(64, 16), # Smaller batches, more concurrent
(800, 1), # Single batch (all at once)
]
results = []
print(f"\n⏱️ Benchmarking {num_pairs} pairs with different configurations:\n")
for batch_size, max_concurrent in configs:
encoder = RemoteTEICrossEncoder(
base_url=TEI_RERANKER_URL,
batch_size=batch_size,
max_concurrent=max_concurrent,
timeout=60.0,
)
await encoder.initialize()
# Warm-up run
await encoder.predict(test_pairs[:100])
# Timed runs (3 iterations)
times = []
for _ in range(3):
start = time.time()
scores = await encoder.predict(test_pairs)
elapsed = time.time() - start
times.append(elapsed)
assert len(scores) == num_pairs
avg_time = sum(times) / len(times)
min_time = min(times)
results.append({
"batch_size": batch_size,
"max_concurrent": max_concurrent,
"avg_ms": avg_time * 1000,
"min_ms": min_time * 1000,
"num_batches": (num_pairs + batch_size - 1) // batch_size,
})
print(f" batch_size={batch_size:4d}, max_concurrent={max_concurrent:2d}: "
f"avg={avg_time * 1000:6.1f}ms, min={min_time * 1000:6.1f}ms "
f"({results[-1]['num_batches']} batches)")
# Find best configuration
best = min(results, key=lambda x: x["avg_ms"])
print(f"\n🏆 Best Configuration:")
print(f" batch_size={best['batch_size']}, max_concurrent={best['max_concurrent']}")
print(f" Average: {best['avg_ms']:.1f}ms, Min: {best['min_ms']:.1f}ms")
# Performance target check
target_ms = 100
if best["avg_ms"] <= target_ms:
print(f"\n✅ Target met! Average {best['avg_ms']:.1f}ms <= {target_ms}ms")
else:
print(f"\n⚠️ Target NOT met. Average {best['avg_ms']:.1f}ms > {target_ms}ms")
print(f" Consider: larger batch size, GPU optimization, or faster network")
@requires_tei_server
@pytest.mark.asyncio
async def test_tei_reranker_concurrent_requests():
"""
Test TEI reranker performance under concurrent request load.
This simulates multiple parallel recall requests hitting the reranker
at the same time.
"""
# Smaller batches to simulate typical recall workload
num_pairs_per_request = 200
num_concurrent_requests = 4
query = "Tell me about machine learning and AI training"
test_pairs = [
(query, f"Document {i} about ML and training.")
for i in range(num_pairs_per_request)
]
# Test configurations
configs = [
(128, 8), # Default
(256, 4), # Larger batches
(512, 2), # Very large batches
(200, 1), # Single batch per request
]
print(f"\n⏱️ Concurrent Load Test: {num_concurrent_requests} parallel requests, "
f"{num_pairs_per_request} pairs each:\n")
for batch_size, max_concurrent in configs:
encoder = RemoteTEICrossEncoder(
base_url=TEI_RERANKER_URL,
batch_size=batch_size,
max_concurrent=max_concurrent,
timeout=60.0,
)
await encoder.initialize()
# Warm-up
await encoder.predict(test_pairs[:50])
async def run_single_request():
start = time.time()
scores = await encoder.predict(test_pairs)
return time.time() - start, len(scores)
# Run concurrent requests
times = []
for _ in range(3): # 3 iterations
start = time.time()
results = await asyncio.gather(*[run_single_request() for _ in range(num_concurrent_requests)])
total_time = time.time() - start
individual_times = [r[0] for r in results]
times.append({
"total": total_time,
"max_individual": max(individual_times),
"avg_individual": sum(individual_times) / len(individual_times),
})
avg_total = sum(t["total"] for t in times) / len(times)
avg_max_individual = sum(t["max_individual"] for t in times) / len(times)
print(f" batch_size={batch_size:4d}, max_concurrent={max_concurrent:2d}: "
f"total={avg_total * 1000:6.1f}ms, slowest_req={avg_max_individual * 1000:6.1f}ms")
@requires_tei_server
@pytest.mark.asyncio
async def test_tei_reranker_latency_breakdown():
"""
Measure latency breakdown for TEI reranker requests.
This helps identify where time is spent: network vs processing.
"""
import httpx
print(f"\n⏱️ Latency Breakdown Test:\n")
# Test single document latency (network overhead)
async with httpx.AsyncClient(timeout=30.0) as client:
times = []
for _ in range(10):
start = time.time()
await client.post(
f"{TEI_RERANKER_URL}/rerank",
json={
"query": "test query",
"texts": ["test document"],
"return_text": False,
},
)
times.append((time.time() - start) * 1000)
avg_single = sum(times) / len(times)
print(f" Single doc latency (raw HTTP): {avg_single:.2f}ms")
# Test batch latencies
batch_sizes = [10, 50, 100, 200, 500]
for batch_size in batch_sizes:
texts = [f"Document {i} about machine learning" for i in range(batch_size)]
async with httpx.AsyncClient(timeout=30.0) as client:
times = []
for _ in range(5):
start = time.time()
await client.post(
f"{TEI_RERANKER_URL}/rerank",
json={
"query": "What about machine learning?",
"texts": texts,
"return_text": False,
},
)
times.append((time.time() - start) * 1000)
avg = sum(times) / len(times)
per_doc = avg / batch_size
print(f" Batch size {batch_size:4d}: {avg:6.1f}ms total, {per_doc:.2f}ms/doc")
print(f"\n 💡 Insight: Higher per-doc time at small batches = network overhead dominant")
print(f" 💡 Insight: Lower per-doc time at large batches = GPU efficiently utilized")
+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"
+3 -3
View File
@@ -103,7 +103,7 @@ impl ApiClient {
pub fn get_stats(&self, agent_id: &str, _verbose: bool) -> Result<AgentStats> {
self.runtime.block_on(async {
let response = self.client.get_agent_stats(agent_id).await?;
let response = self.client.get_agent_stats(agent_id, None).await?;
let value = response.into_inner();
// Convert to JSON Value first, then parse into our type
let json_value = serde_json::to_value(&value)?;
@@ -241,9 +241,9 @@ impl ApiClient {
})
}
pub fn list_entities(&self, bank_id: &str, limit: Option<i64>, _verbose: bool) -> Result<types::EntityListResponse> {
pub fn list_entities(&self, bank_id: &str, limit: Option<i64>, offset: Option<i64>, _verbose: bool) -> Result<types::EntityListResponse> {
self.runtime.block_on(async {
let response = self.client.list_entities(bank_id, limit, None).await?;
let response = self.client.list_entities(bank_id, limit, offset, None).await?;
Ok(response.into_inner())
})
}
+1 -1
View File
@@ -16,7 +16,7 @@ pub fn list(
None
};
let response = client.list_entities(bank_id, Some(limit), verbose)?;
let response = client.list_entities(bank_id, Some(limit), None, verbose)?;
if let Some(mut sp) = spinner {
sp.finish();
+6 -2
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},
@@ -283,7 +283,7 @@ impl App {
}
fn load_entities(&mut self, bank_id: &str) -> Result<()> {
let response = self.client.list_entities(bank_id, Some(100), false)?;
let response = self.client.list_entities(bank_id, Some(100), None, false)?;
self.entities = response.items;
if !self.entities.is_empty() && self.entities_state.selected().is_none() {
@@ -341,6 +341,8 @@ impl App {
trace: false,
query_timestamp: None,
include: None,
tags: None,
tags_match: TagsMatch::Any,
};
let result = client.recall(&bank_id, &request, false)
@@ -357,6 +359,8 @@ impl App {
max_tokens: 4096,
include: None,
response_schema: None,
tags: None,
tags_match: TagsMatch::Any,
};
let result = client.reflect(&bank_id, &request, false)
+9 -1
View File
@@ -9,7 +9,7 @@ use crate::output::{self, OutputFormat};
use crate::ui;
// Import types from generated client
use hindsight_client::types::{Budget, ChunkIncludeOptions, IncludeOptions};
use hindsight_client::types::{Budget, ChunkIncludeOptions, IncludeOptions, TagsMatch};
use serde_json;
// Helper function to parse budget string to Budget enum
@@ -60,6 +60,8 @@ pub fn recall(
trace,
query_timestamp: None,
include,
tags: None,
tags_match: TagsMatch::Any,
};
let response = client.recall(agent_id, &request, verbose);
@@ -116,6 +118,8 @@ pub fn reflect(
max_tokens: max_tokens.unwrap_or(4096),
include: None,
response_schema,
tags: None,
tags_match: TagsMatch::Any,
};
let response = client.reflect(agent_id, &request, verbose);
@@ -162,11 +166,13 @@ pub fn retain(
timestamp: None,
document_id: Some(doc_id.clone()),
entities: None,
tags: None,
};
let request = RetainRequest {
items: vec![item],
async_: r#async,
document_tags: None,
};
let response = client.retain(agent_id, &request, r#async, verbose);
@@ -272,6 +278,7 @@ pub fn retain_files(
timestamp: None,
document_id: Some(doc_id),
entities: None,
tags: None,
});
pb.inc(1);
@@ -288,6 +295,7 @@ pub fn retain_files(
let request = RetainRequest {
items,
async_: r#async,
document_tags: None,
};
let response = client.retain(agent_id, &request, r#async, verbose);
@@ -39,6 +39,7 @@ hindsight_client_api/models/http_validation_error.py
hindsight_client_api/models/include_options.py
hindsight_client_api/models/list_documents_response.py
hindsight_client_api/models/list_memory_units_response.py
hindsight_client_api/models/list_tags_response.py
hindsight_client_api/models/memory_item.py
hindsight_client_api/models/operation_response.py
hindsight_client_api/models/operations_list_response.py
@@ -51,6 +52,7 @@ hindsight_client_api/models/reflect_request.py
hindsight_client_api/models/reflect_response.py
hindsight_client_api/models/retain_request.py
hindsight_client_api/models/retain_response.py
hindsight_client_api/models/tag_item.py
hindsight_client_api/models/token_usage.py
hindsight_client_api/models/update_disposition_request.py
hindsight_client_api/models/validation_error.py
@@ -64,6 +64,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
from hindsight_client_api.models.include_options import IncludeOptions
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
from hindsight_client_api.models.list_tags_response import ListTagsResponse
from hindsight_client_api.models.memory_item import MemoryItem
from hindsight_client_api.models.operation_response import OperationResponse
from hindsight_client_api.models.operations_list_response import OperationsListResponse
@@ -76,6 +77,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
from hindsight_client_api.models.reflect_response import ReflectResponse
from hindsight_client_api.models.retain_request import RetainRequest
from hindsight_client_api.models.retain_response import RetainResponse
from hindsight_client_api.models.tag_item import TagItem
from hindsight_client_api.models.token_usage import TokenUsage
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
from hindsight_client_api.models.validation_error import ValidationError
@@ -939,6 +939,7 @@ class BanksApi:
async def get_agent_stats(
self,
bank_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
@@ -958,6 +959,8 @@ class BanksApi:
:param bank_id: (required)
:type bank_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
@@ -982,6 +985,7 @@ class BanksApi:
_param = self._get_agent_stats_serialize(
bank_id=bank_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
@@ -1007,6 +1011,7 @@ class BanksApi:
async def get_agent_stats_with_http_info(
self,
bank_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
@@ -1026,6 +1031,8 @@ class BanksApi:
:param bank_id: (required)
:type bank_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
@@ -1050,6 +1057,7 @@ class BanksApi:
_param = self._get_agent_stats_serialize(
bank_id=bank_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
@@ -1075,6 +1083,7 @@ class BanksApi:
async def get_agent_stats_without_preload_content(
self,
bank_id: StrictStr,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
Annotated[StrictFloat, Field(gt=0)],
@@ -1094,6 +1103,8 @@ class BanksApi:
:param bank_id: (required)
:type bank_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
@@ -1118,6 +1129,7 @@ class BanksApi:
_param = self._get_agent_stats_serialize(
bank_id=bank_id,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
_headers=_headers,
@@ -1138,6 +1150,7 @@ class BanksApi:
def _get_agent_stats_serialize(
self,
bank_id,
authorization,
_request_auth,
_content_type,
_headers,
@@ -1163,6 +1176,8 @@ class BanksApi:
_path_params['bank_id'] = bank_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
@@ -338,6 +338,7 @@ class EntitiesApi:
self,
bank_id: StrictStr,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of entities to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
@@ -354,12 +355,14 @@ class EntitiesApi:
) -> EntityListResponse:
"""List entities
List all entities (people, organizations, etc.) known by the bank, ordered by mention count.
List all entities (people, organizations, etc.) known by the bank, ordered by mention count. Supports pagination.
:param bank_id: (required)
:type bank_id: str
:param limit: Maximum number of entities 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
@@ -387,6 +390,7 @@ class EntitiesApi:
_param = self._list_entities_serialize(
bank_id=bank_id,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
@@ -414,6 +418,7 @@ class EntitiesApi:
self,
bank_id: StrictStr,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of entities to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
@@ -430,12 +435,14 @@ class EntitiesApi:
) -> ApiResponse[EntityListResponse]:
"""List entities
List all entities (people, organizations, etc.) known by the bank, ordered by mention count.
List all entities (people, organizations, etc.) known by the bank, ordered by mention count. Supports pagination.
:param bank_id: (required)
:type bank_id: str
:param limit: Maximum number of entities 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
@@ -463,6 +470,7 @@ class EntitiesApi:
_param = self._list_entities_serialize(
bank_id=bank_id,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
@@ -490,6 +498,7 @@ class EntitiesApi:
self,
bank_id: StrictStr,
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of entities to return")] = None,
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
authorization: Optional[StrictStr] = None,
_request_timeout: Union[
None,
@@ -506,12 +515,14 @@ class EntitiesApi:
) -> RESTResponseType:
"""List entities
List all entities (people, organizations, etc.) known by the bank, ordered by mention count.
List all entities (people, organizations, etc.) known by the bank, ordered by mention count. Supports pagination.
:param bank_id: (required)
:type bank_id: str
:param limit: Maximum number of entities 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
@@ -539,6 +550,7 @@ class EntitiesApi:
_param = self._list_entities_serialize(
bank_id=bank_id,
limit=limit,
offset=offset,
authorization=authorization,
_request_auth=_request_auth,
_content_type=_content_type,
@@ -561,6 +573,7 @@ class EntitiesApi:
self,
bank_id,
limit,
offset,
authorization,
_request_auth,
_content_type,
@@ -590,6 +603,10 @@ class EntitiesApi:
_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
@@ -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
@@ -17,7 +17,7 @@ import pprint
import re # noqa: F401
import json
from pydantic import BaseModel, ConfigDict
from pydantic import BaseModel, ConfigDict, StrictInt
from typing import Any, ClassVar, Dict, List
from hindsight_client_api.models.entity_list_item import EntityListItem
from typing import Optional, Set
@@ -28,7 +28,10 @@ class EntityListResponse(BaseModel):
Response model for entity list endpoint.
""" # noqa: E501
items: List[EntityListItem]
__properties: ClassVar[List[str]] = ["items"]
total: StrictInt
limit: StrictInt
offset: StrictInt
__properties: ClassVar[List[str]] = ["items", "total", "limit", "offset"]
model_config = ConfigDict(
populate_by_name=True,
@@ -88,7 +91,10 @@ class EntityListResponse(BaseModel):
return cls.model_validate(obj)
_obj = cls.model_validate({
"items": [EntityListItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None
"items": [EntityListItem.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
@@ -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"}
@@ -449,6 +449,38 @@ class TestEntities:
assert response is not None
assert response.items is not None
assert isinstance(response.items, list)
# Verify pagination fields
assert response.total is not None
assert response.limit is not None
assert response.offset is not None
assert response.offset == 0
assert response.limit == 100 # default limit
def test_list_entities_with_pagination(self, client, bank_id):
"""Test listing entities with pagination parameters."""
import asyncio
from hindsight_client_api import ApiClient, Configuration
from hindsight_client_api.api import EntitiesApi
async def do_list_paginated():
config = Configuration(host=HINDSIGHT_API_URL)
api_client = ApiClient(config)
api = EntitiesApi(api_client)
# Test with custom limit
response = await api.list_entities(bank_id=bank_id, limit=5, offset=0)
assert response.limit == 5
assert response.offset == 0
assert len(response.items) <= 5
# Test with offset
response_offset = await api.list_entities(bank_id=bank_id, limit=1, offset=1)
assert response_offset.offset == 1
assert response_offset.limit == 1
return response
asyncio.get_event_loop().run_until_complete(do_list_paginated())
def test_get_entity(self, client, bank_id):
"""Test getting a specific entity."""
+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)
@@ -169,8 +169,6 @@ export const createSseClient = <TData = unknown>({
const { done, value } = await reader.read();
if (done) break;
buffer += value;
// Normalize line endings: CRLF -> LF, then CR -> LF
buffer = buffer.replace(/\r\n/g, "\n").replace(/\r/g, "\n");
const chunks = buffer.split("\n\n");
buffer = chunks.pop() ?? "";
@@ -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
*
@@ -236,7 +256,7 @@ export const getAgentStats = <ThrowOnError extends boolean = false>(
/**
* List entities
*
* List all entities (people, organizations, etc.) known by the bank, ordered by mention count.
* List all entities (people, organizations, etc.) known by the bank, ordered by mention count. Supports pagination.
*/
export const listEntities = <ThrowOnError extends boolean = false>(
options: Options<ListEntitiesData, ThrowOnError>,
@@ -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>;
};
/**
@@ -495,6 +501,18 @@ export type EntityListResponse = {
* Items
*/
items: Array<EntityListItem>;
/**
* Total
*/
total: number;
/**
* Limit
*/
limit: number;
/**
* Offset
*/
offset: number;
};
/**
@@ -654,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
*
@@ -690,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;
};
/**
@@ -779,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";
};
/**
@@ -867,6 +927,10 @@ export type RecallResult = {
* Chunk Id
*/
chunk_id?: string | null;
/**
* Tags
*/
tags?: Array<string> | null;
};
/**
@@ -946,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";
};
/**
@@ -992,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;
};
/**
@@ -1030,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
*
@@ -1213,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?: {
@@ -1320,6 +1460,12 @@ export type ListBanksResponse = ListBanksResponses[keyof ListBanksResponses];
export type GetAgentStatsData = {
body?: never;
headers?: {
/**
* Authorization
*/
authorization?: string | null;
};
path: {
/**
* Bank Id
@@ -1370,6 +1516,12 @@ export type ListEntitiesData = {
* Maximum number of entities to return
*/
limit?: number;
/**
* Offset
*
* Offset for pagination
*/
offset?: number;
};
url: "/v1/default/banks/{bank_id}/entities";
};
@@ -1608,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",
+21 -12
View File
@@ -62,6 +62,7 @@ export interface MemoryItemInput {
metadata?: Record<string, string>;
document_id?: string;
entities?: EntityInput[];
tags?: string[];
}
export class HindsightClient {
@@ -78,6 +79,16 @@ export class HindsightClient {
);
}
/**
* Validates the API response and throws an error if the request failed.
*/
private validateResponse<T>(response: { data?: T; error?: unknown }, operation: string): T {
if (!response.data) {
throw new Error(`${operation} failed: ${JSON.stringify(response.error || 'Unknown error')}`);
}
return response.data;
}
/**
* Retain a single memory for a bank.
*/
@@ -126,19 +137,20 @@ export class HindsightClient {
body: { items: [item], async: options?.async },
});
return response.data!;
return this.validateResponse(response, 'retain');
}
/**
* 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()
@@ -156,11 +168,12 @@ export class HindsightClient {
path: { bank_id: bankId },
body: {
items: itemsWithDocId,
document_tags: options?.documentTags,
async: options?.async,
},
});
return response.data!;
return this.validateResponse(response, 'retainBatch');
}
/**
@@ -198,11 +211,7 @@ export class HindsightClient {
},
});
if (!response.data) {
throw new Error(`API returned no data: ${JSON.stringify(response.error || 'Unknown error')}`);
}
return response.data;
return this.validateResponse(response, 'recall');
}
/**
@@ -223,7 +232,7 @@ export class HindsightClient {
},
});
return response.data!;
return this.validateResponse(response, 'reflect');
}
/**
@@ -244,7 +253,7 @@ export class HindsightClient {
},
});
return response.data!;
return this.validateResponse(response, 'listMemories');
}
/**
@@ -264,7 +273,7 @@ export class HindsightClient {
},
});
return response.data!;
return this.validateResponse(response, 'createBank');
}
/**
@@ -276,7 +285,7 @@ export class HindsightClient {
path: { bank_id: bankId },
});
return response.data!;
return this.validateResponse(response, 'getBankProfile');
}
}
+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"
@@ -11,11 +11,12 @@ export async function GET(request: NextRequest) {
}
const limit = searchParams.get("limit") ? Number(searchParams.get("limit")) : undefined;
const offset = searchParams.get("offset") ? Number(searchParams.get("offset")) : undefined;
const response = await sdk.listEntities({
client: lowLevelClient,
path: { bank_id: bankId },
query: { limit },
query: { limit, offset },
});
if (response.error) {
@@ -1,5 +1,7 @@
import { NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
import { createClient, createConfig, sdk } from "@vectorize-io/hindsight-client";
const HEALTH_CHECK_TIMEOUT_MS = 3000;
export async function GET() {
const status: {
@@ -15,19 +17,37 @@ export async function GET() {
service: "hindsight-control-plane",
};
// Check dataplane connectivity
// Check dataplane connectivity with a short timeout
const dataplaneUrl = process.env.HINDSIGHT_CP_DATAPLANE_API_URL || "http://localhost:8888";
try {
await sdk.listBanks({ client: lowLevelClient });
status.dataplane = {
status: "connected",
url: dataplaneUrl,
};
const controller = new AbortController();
const timeoutId = setTimeout(() => controller.abort(), HEALTH_CHECK_TIMEOUT_MS);
const healthClient = createClient(
createConfig({
baseUrl: dataplaneUrl,
signal: controller.signal,
})
);
try {
await sdk.listBanks({ client: healthClient });
status.dataplane = {
status: "connected",
url: dataplaneUrl,
};
} finally {
clearTimeout(timeoutId);
}
} catch (error) {
let errorMessage = error instanceof Error ? error.message : String(error);
if (error instanceof Error && error.name === "AbortError") {
errorMessage = `Request timed out after ${HEALTH_CHECK_TIMEOUT_MS}ms`;
}
status.dataplane = {
status: "disconnected",
url: dataplaneUrl,
error: error instanceof Error ? error.message : String(error),
error: errorMessage,
};
}
@@ -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
@@ -4,6 +4,7 @@ import { useState, useEffect } from "react";
import { client } from "@/lib/api";
import { useBank } from "@/lib/bank-context";
import { Button } from "@/components/ui/button";
import { ChevronLeft, ChevronRight, ChevronsLeft, ChevronsRight } from "lucide-react";
import {
Table,
TableBody,
@@ -29,6 +30,8 @@ interface EntityDetail extends Entity {
}>;
}
const ITEMS_PER_PAGE = 50;
export function EntitiesView() {
const { currentBank } = useBank();
const [entities, setEntities] = useState<Entity[]>([]);
@@ -37,16 +40,26 @@ export function EntitiesView() {
const [loadingDetail, setLoadingDetail] = useState(false);
const [regenerating, setRegenerating] = useState(false);
const loadEntities = async () => {
// Pagination state
const [currentPage, setCurrentPage] = useState(1);
const [total, setTotal] = useState(0);
const totalPages = Math.ceil(total / ITEMS_PER_PAGE);
const offset = (currentPage - 1) * ITEMS_PER_PAGE;
const loadEntities = async (page: number = 1) => {
if (!currentBank) return;
setLoading(true);
try {
const result: any = await client.listEntities({
const pageOffset = (page - 1) * ITEMS_PER_PAGE;
const result = await client.listEntities({
bank_id: currentBank,
limit: 100,
limit: ITEMS_PER_PAGE,
offset: pageOffset,
});
setEntities(result.items || []);
setTotal(result.total || 0);
} catch (error) {
console.error("Error loading entities:", error);
alert("Error loading entities: " + (error as Error).message);
@@ -86,9 +99,16 @@ export function EntitiesView() {
}
};
// Handle page change
const handlePageChange = (newPage: number) => {
setCurrentPage(newPage);
loadEntities(newPage);
};
useEffect(() => {
if (currentBank) {
loadEntities();
setCurrentPage(1);
loadEntities(1);
setSelectedEntity(null);
}
}, [currentBank]);
@@ -105,13 +125,15 @@ export function EntitiesView() {
{loading ? (
<div className="flex items-center justify-center py-20">
<div className="text-center">
<div className="text-4xl mb-2"></div>
<div className="text-4xl mb-2">...</div>
<div className="text-sm text-muted-foreground">Loading entities...</div>
</div>
</div>
) : entities.length > 0 ? (
<>
<div className="mb-4 text-sm text-muted-foreground">{entities.length} entities</div>
<div className="mb-4 text-sm text-muted-foreground">
{total} {total === 1 ? "entity" : "entities"}
</div>
<div className="overflow-x-auto">
<Table>
<TableHeader>
@@ -146,11 +168,61 @@ export function EntitiesView() {
</TableBody>
</Table>
</div>
{/* Pagination Controls */}
{totalPages > 1 && (
<div className="flex items-center justify-between mt-3 pt-3 border-t">
<div className="text-xs text-muted-foreground">
{offset + 1}-{Math.min(offset + ITEMS_PER_PAGE, total)} of {total}
</div>
<div className="flex items-center gap-1">
<Button
variant="outline"
size="sm"
onClick={() => handlePageChange(1)}
disabled={currentPage === 1 || loading}
className="h-7 w-7 p-0"
>
<ChevronsLeft className="h-3 w-3" />
</Button>
<Button
variant="outline"
size="sm"
onClick={() => handlePageChange(currentPage - 1)}
disabled={currentPage === 1 || loading}
className="h-7 w-7 p-0"
>
<ChevronLeft className="h-3 w-3" />
</Button>
<span className="text-xs px-2">
{currentPage} / {totalPages}
</span>
<Button
variant="outline"
size="sm"
onClick={() => handlePageChange(currentPage + 1)}
disabled={currentPage === totalPages || loading}
className="h-7 w-7 p-0"
>
<ChevronRight className="h-3 w-3" />
</Button>
<Button
variant="outline"
size="sm"
onClick={() => handlePageChange(totalPages)}
disabled={currentPage === totalPages || loading}
className="h-7 w-7 p-0"
>
<ChevronsRight className="h-3 w-3" />
</Button>
</div>
</div>
)}
</>
) : (
<div className="flex items-center justify-center py-20">
<div className="text-center">
<div className="text-4xl mb-2">👥</div>
<div className="text-4xl mb-2">...</div>
<div className="text-sm text-muted-foreground">No entities found</div>
<div className="text-xs text-muted-foreground mt-1">
Entities are extracted from facts when memories are added.
@@ -178,7 +250,7 @@ export function EntitiesView() {
onClick={() => setSelectedEntity(null)}
className="h-8 w-8 p-0"
>
<span className="text-lg">×</span>
<span className="text-lg">x</span>
</Button>
</div>
@@ -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>
+32 -2
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",
@@ -127,11 +131,17 @@ export class ControlPlaneClient {
/**
* List entities
*/
async listEntities(params: { bank_id: string; limit?: number }) {
async listEntities(params: { bank_id: string; limit?: number; offset?: number }) {
const queryParams = new URLSearchParams();
queryParams.append("bank_id", params.bank_id);
if (params.limit) queryParams.append("limit", params.limit.toString());
return this.fetchApi(`/api/entities?${queryParams}`);
if (params.offset) queryParams.append("offset", params.offset.toString());
return this.fetchApi<{
items: any[];
total: number;
limit: number;
offset: number;
}>(`/api/entities?${queryParams}`);
}
/**
@@ -203,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)
+81 -13
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,10 +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
@@ -205,16 +233,26 @@ 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/...
```
### Server
| Variable | Description | Default |
|----------|-------------|---------|
| `HINDSIGHT_API_HOST` | Bind address | `0.0.0.0` |
| `HINDSIGHT_API_PORT` | Server port | `8888` |
| `HINDSIGHT_API_LOG_LEVEL` | Log level: `debug`, `info`, `warning`, `error` | `info` |
| `HINDSIGHT_API_MCP_ENABLED` | Enable MCP server at `/mcp/{bank_id}/` | `true` |
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
@@ -239,11 +277,29 @@ Requests without a valid API key receive a `401 Unauthorized` response.
For advanced authentication (JWT, OAuth, multi-tenant schemas), implement a custom `TenantExtension`. See the [Extensions documentation](./extensions.md) for details.
:::
### Server
| Variable | Description | Default |
|----------|-------------|---------|
| `HINDSIGHT_API_HOST` | Bind address | `0.0.0.0` |
| `HINDSIGHT_API_PORT` | Server port | `8888` |
| `HINDSIGHT_API_WORKERS` | Number of uvicorn worker processes | `1` |
| `HINDSIGHT_API_LOG_LEVEL` | Log level: `debug`, `info`, `warning`, `error` | `info` |
| `HINDSIGHT_API_MCP_ENABLED` | Enable MCP server at `/mcp/{bank_id}/` | `true` |
### Retrieval
| Variable | Description | Default |
|----------|-------------|---------|
| `HINDSIGHT_API_GRAPH_RETRIEVER` | Graph retrieval algorithm: `bfs` or `mpfp` | `bfs` |
| `HINDSIGHT_API_GRAPH_RETRIEVER` | Graph retrieval algorithm: `link_expansion`, `mpfp`, or `bfs` | `link_expansion` |
| `HINDSIGHT_API_RECALL_MAX_CONCURRENT` | Max concurrent recall operations per worker (backpressure) | `32` |
| `HINDSIGHT_API_RERANKER_MAX_CANDIDATES` | Max candidates to rerank per recall (RRF pre-filters the rest) | `300` |
#### Graph Retrieval Algorithms
- **`link_expansion`** (default): Fast, simple graph expansion from semantic seeds via entity co-occurrence and causal links. Target latency under 100ms. Recommended for most use cases.
- **`mpfp`**: Multi-Path Fact Propagation - iterative graph traversal with activation spreading. More thorough but slower.
- **`bfs`**: Breadth-first search from seed facts. Simple but less effective for large graphs.
### Entity Observations
@@ -262,6 +318,17 @@ Controls the retain (memory ingestion) pipeline.
|----------|-------------|---------|
| `HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS` | Max completion tokens for fact extraction LLM calls | `64000` |
| `HINDSIGHT_API_RETAIN_CHUNK_SIZE` | Max characters per chunk for fact extraction. Larger chunks extract fewer LLM calls but may lose context. | `3000` |
| `HINDSIGHT_API_RETAIN_EXTRACTION_MODE` | Fact extraction mode: `concise` (selective, fewer high-quality facts) or `verbose` (detailed, more facts) | `concise` |
| `HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS` | Extract causal relationships between facts | `true` |
| `HINDSIGHT_API_RETAIN_OBSERVATIONS_ASYNC` | Run entity observation generation asynchronously (after retain completes) | `false` |
#### Extraction Modes
The extraction mode controls how aggressively facts are extracted from content:
- **`concise`** (default): Selective extraction that focuses on significant, long-term valuable facts. Filters out greetings, filler, and trivial information. Produces fewer but higher-quality facts with better performance.
- **`verbose`**: Detailed extraction that captures every piece of information with maximum verbosity. Produces more facts with extensive detail but slower performance and higher token usage.
### Local MCP Server
@@ -283,8 +350,9 @@ Controls background task processing for async operations like opinion formation
| Variable | Description | Default |
|----------|-------------|---------|
| `HINDSIGHT_API_TASK_BATCH_SIZE` | Max tasks to process in one batch | `10` |
| `HINDSIGHT_API_TASK_BATCH_INTERVAL` | Interval between batch processing in seconds | `1.0` |
| `HINDSIGHT_API_TASK_BACKEND` | Task backend implementation: `memory` (in-process queue) or `noop` (discard tasks, useful for tests) | `memory` |
| `HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_SIZE` | Max tasks to process in one batch (memory backend only) | `10` |
| `HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_INTERVAL` | Interval between batch processing in seconds (memory backend only) | `1.0` |
### Performance Optimization
@@ -171,4 +171,4 @@ PORT=80 HINDSIGHT_CP_DATAPLANE_API_URL=https://api.hindsight.io npx @vectorize-i
- [Configuration](./configuration.md) — Environment variables and settings
- [Models](./models.md) — ML models and providers
- [Metrics](./metrics.md) — Monitoring and observability
- [Monitoring](./monitoring.md) — Metrics and observability
-95
View File
@@ -1,95 +0,0 @@
# Metrics
Hindsight exposes Prometheus metrics at `/metrics` for monitoring.
```bash
curl http://localhost:8888/metrics
```
## Available Metrics
### Operation Metrics
| Metric | Type | Labels | Description |
|--------|------|--------|-------------|
| `hindsight.operation.duration` | Histogram | operation, bank_id, source, budget, max_tokens, success | Duration of operations in seconds |
| `hindsight.operation.total` | Counter | operation, bank_id, source, budget, max_tokens, success | Total number of operations executed |
**Labels:**
- `operation`: Operation type (`retain`, `recall`, `reflect`)
- `bank_id`: Memory bank identifier
- `source`: Where the operation was triggered from (`api`, `reflect`, `internal`)
- `budget`: Budget level if specified (`low`, `mid`, `high`)
- `max_tokens`: Max tokens if specified
- `success`: Whether the operation succeeded (`true`, `false`)
The `source` label allows distinguishing between:
- `api`: Direct API calls from clients
- `reflect`: Internal recall calls made during reflect operations
- `internal`: Other internal operations
### LLM Metrics
| Metric | Type | Labels | Description |
|--------|------|--------|-------------|
| `hindsight.llm.duration` | Histogram | provider, model, scope, success | Duration of LLM API calls in seconds |
| `hindsight.llm.calls.total` | Counter | provider, model, scope, success | Total number of LLM API calls |
| `hindsight.llm.tokens.input` | Counter | provider, model, scope, success, token_bucket | Input tokens for LLM calls |
| `hindsight.llm.tokens.output` | Counter | provider, model, scope, success, token_bucket | Output tokens from LLM calls |
**Labels:**
- `provider`: LLM provider (`openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`)
- `model`: Model name (e.g., `gpt-4`, `claude-3-sonnet`)
- `scope`: What the LLM call is for (`memory`, `reflect`, `entity_observation`, `answer`)
- `success`: Whether the call succeeded (`true`, `false`)
- `token_bucket`: Token count bucket for cardinality control (`0-100`, `100-500`, `500-1k`, `1k-5k`, `5k-10k`, `10k-50k`, `50k+`)
### Histogram Buckets
Custom bucket boundaries are configured for better percentile accuracy:
**Operation Duration Buckets (seconds):**
```
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
```
**LLM Duration Buckets (seconds):**
```
0.1, 0.25, 0.5, 1.0, 2.0, 3.0, 5.0, 10.0, 15.0, 30.0, 60.0, 120.0
```
## Prometheus Configuration
```yaml
scrape_configs:
- job_name: 'hindsight'
static_configs:
- targets: ['localhost:8888']
```
## Example Queries
### Average operation latency by type
```promql
rate(hindsight_operation_duration_sum[5m]) / rate(hindsight_operation_duration_count[5m])
```
### LLM calls per minute by provider
```promql
rate(hindsight_llm_calls_total[1m]) * 60
```
### P95 LLM latency
```promql
histogram_quantile(0.95, rate(hindsight_llm_duration_bucket[5m]))
```
### Total tokens consumed by model
```promql
sum by (model) (hindsight_llm_tokens_input_total + hindsight_llm_tokens_output_total)
```
### Internal vs API recall operations
```promql
sum by (source) (rate(hindsight_operation_total{operation="recall"}[5m]))
```
+199
View File
@@ -0,0 +1,199 @@
# Monitoring
Hindsight provides comprehensive monitoring through Prometheus metrics and pre-built Grafana dashboards.
## Local Development
For local metrics visualization, a convenience script downloads and runs Prometheus and Grafana:
```bash
./scripts/dev/start-monitoring.sh
```
This will start:
- **Grafana**: http://localhost:8890 (anonymous access enabled)
- **Prometheus**: http://localhost:8889
- **API Metrics**: http://localhost:8888/metrics
:::note Production Deployment
The local monitoring script is for development only. In production, you need to install and configure Prometheus and Grafana separately, then point Prometheus to scrape your Hindsight API's `/metrics` endpoint.
:::
## Grafana Dashboards
Pre-built dashboards are available in [`monitoring/grafana/dashboards/`](https://github.com/anthropics/hindsight/tree/main/monitoring/grafana/dashboards). Import these JSON files into your Grafana instance:
| Dashboard | Description |
|-----------|-------------|
| **Hindsight Operations** | Operation rates, latency percentiles, per-bank metrics |
| **Hindsight LLM Metrics** | LLM calls, token usage, latency by scope/provider |
| **Hindsight API Service** | HTTP requests, error rates, DB pool, process metrics |
The dashboards are automatically provisioned when using the monitoring stack script.
## Metrics Endpoint
Hindsight exposes Prometheus metrics at `/metrics`:
```bash
curl http://localhost:8888/metrics
```
## Available Metrics
### Operation Metrics
| Metric | Type | Labels | Description |
|--------|------|--------|-------------|
| `hindsight.operation.duration` | Histogram | operation, bank_id, source, budget, max_tokens, success | Duration of operations in seconds |
| `hindsight.operation.total` | Counter | operation, bank_id, source, budget, max_tokens, success | Total number of operations executed |
**Labels:**
- `operation`: Operation type (`retain`, `recall`, `reflect`)
- `bank_id`: Memory bank identifier
- `source`: Where the operation was triggered from (`api`, `reflect`, `internal`)
- `budget`: Budget level if specified (`low`, `mid`, `high`)
- `max_tokens`: Max tokens if specified
- `success`: Whether the operation succeeded (`true`, `false`)
The `source` label allows distinguishing between:
- `api`: Direct API calls from clients
- `reflect`: Internal recall calls made during reflect operations
- `internal`: Other internal operations
### LLM Metrics
| Metric | Type | Labels | Description |
|--------|------|--------|-------------|
| `hindsight.llm.duration` | Histogram | provider, model, scope, success | Duration of LLM API calls in seconds |
| `hindsight.llm.calls.total` | Counter | provider, model, scope, success | Total number of LLM API calls |
| `hindsight.llm.tokens.input` | Counter | provider, model, scope, success, token_bucket | Input tokens for LLM calls |
| `hindsight.llm.tokens.output` | Counter | provider, model, scope, success, token_bucket | Output tokens from LLM calls |
**Labels:**
- `provider`: LLM provider (`openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`)
- `model`: Model name (e.g., `gpt-4`, `claude-3-sonnet`)
- `scope`: What the LLM call is for (`memory`, `reflect`, `entity_observation`, `answer`)
- `success`: Whether the call succeeded (`true`, `false`)
- `token_bucket`: Token count bucket for cardinality control (`0-100`, `100-500`, `500-1k`, `1k-5k`, `5k-10k`, `10k-50k`, `50k+`)
### HTTP Request Metrics
| Metric | Type | Labels | Description |
|--------|------|--------|-------------|
| `hindsight.http.duration` | Histogram | method, endpoint, status_code, status_class | Duration of HTTP requests in seconds |
| `hindsight.http.requests.total` | Counter | method, endpoint, status_code, status_class | Total number of HTTP requests |
| `hindsight.http.requests.in_progress` | UpDownCounter | method, endpoint | Number of HTTP requests currently being processed |
**Labels:**
- `method`: HTTP method (`GET`, `POST`, `PUT`, `DELETE`)
- `endpoint`: Request path (normalized to reduce cardinality - UUIDs replaced with `{id}`)
- `status_code`: HTTP status code (`200`, `400`, `500`, etc.)
- `status_class`: Status code class (`2xx`, `4xx`, `5xx`)
### Database Pool Metrics
| Metric | Type | Labels | Description |
|--------|------|--------|-------------|
| `hindsight.db.pool.size` | Gauge | - | Current number of connections in the pool |
| `hindsight.db.pool.idle` | Gauge | - | Number of idle connections in the pool |
| `hindsight.db.pool.min` | Gauge | - | Minimum pool size |
| `hindsight.db.pool.max` | Gauge | - | Maximum pool size |
### Process Metrics
| Metric | Type | Labels | Description |
|--------|------|--------|-------------|
| `hindsight.process.cpu.seconds` | Gauge | type | Process CPU time in seconds |
| `hindsight.process.memory.bytes` | Gauge | type | Process memory usage in bytes |
| `hindsight.process.open_fds` | Gauge | - | Number of open file descriptors |
| `hindsight.process.threads` | Gauge | - | Number of active threads |
**Labels:**
- `type` (CPU): `user` or `system`
- `type` (Memory): `rss_max` (maximum resident set size)
### Histogram Buckets
Custom bucket boundaries are configured for better percentile accuracy:
**Operation Duration Buckets (seconds):**
```
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
```
**LLM Duration Buckets (seconds):**
```
0.1, 0.25, 0.5, 1.0, 2.0, 3.0, 5.0, 10.0, 15.0, 30.0, 60.0, 120.0
```
**HTTP Duration Buckets (seconds):**
```
0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0
```
## Prometheus Configuration
```yaml
scrape_configs:
- job_name: 'hindsight'
static_configs:
- targets: ['localhost:8888']
```
## Example Queries
### Average operation latency by type
```promql
rate(hindsight_operation_duration_sum[5m]) / rate(hindsight_operation_duration_count[5m])
```
### LLM calls per minute by provider
```promql
rate(hindsight_llm_calls_total[1m]) * 60
```
### P95 LLM latency
```promql
histogram_quantile(0.95, rate(hindsight_llm_duration_bucket[5m]))
```
### Total tokens consumed by model
```promql
sum by (model) (hindsight_llm_tokens_input_total + hindsight_llm_tokens_output_total)
```
### Internal vs API recall operations
```promql
sum by (source) (rate(hindsight_operation_total{operation="recall"}[5m]))
```
### HTTP requests per second by endpoint
```promql
sum by (endpoint) (rate(hindsight_http_requests_total[1m]))
```
### HTTP error rate (5xx)
```promql
sum(rate(hindsight_http_requests_total{status_class="5xx"}[5m])) / sum(rate(hindsight_http_requests_total[5m]))
```
### P95 HTTP latency
```promql
histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))
```
### Database pool utilization
```promql
hindsight_db_pool_size / hindsight_db_pool_max
```
### Active database connections
```promql
hindsight_db_pool_size - hindsight_db_pool_idle
```
### CPU usage rate
```promql
rate(hindsight_process_cpu_seconds{type="user"}[1m])
```
+34
View File
@@ -167,6 +167,39 @@ As facts accumulate about an entity, Hindsight synthesizes **observations** —
---
## Tagging Memories
You can tag memories for filtering during recall—useful when one memory bank serves multiple users but each user should only see relevant memories.
```python
# Tag memories for specific users
client.retain(
bank_id="my-agent",
items=[
{
"content": "Alice prefers morning meetings",
"tags": ["user_alice"]
}
]
)
# Apply tags to all items in a batch
client.retain(
bank_id="my-agent",
document_tags=["session_123", "user_alice"], # Applied to all items
items=[
{"content": "Alice discussed the project timeline"},
{"content": "Alice mentioned she needs help with Python"}
]
)
```
During recall, use `tags_match` to control matching:
- `"any"` (default): OR matching - returns memories where **any** tag overlaps
- `"all"`: AND matching - returns memories containing **all** specified tags
---
## What You Get
After `retain()` completes:
@@ -176,6 +209,7 @@ After `retain()` completes:
- **Knowledge graph** with entity, temporal, semantic, and causal links
- **Temporal grounding** for both historical and recency-based queries
- **Background processing** that generates entity summaries
- **Optional tags** for filtering during recall
All stored in your isolated **memory bank**, ready for `recall()` and `reflect()`.
@@ -134,6 +134,8 @@ Hindsight is built for AI agents, not humans. Traditional search systems return
- `max_tokens`: How much memory content to return (default: 4096 tokens)
- `budget`: Search depth level (low, mid, high)
- `fact_type`: Filter by world, experience, opinion, or all
- `tags`: Filter memories by tags
- `tags_match`: How to match tags - `"any"` for OR (default), `"all"` for AND
### Expanding Context: Chunks and Entity Observations
@@ -229,6 +231,14 @@ Budget and max_tokens control different aspects of recall:
---
## Graph Retrieval Algorithms
Hindsight supports multiple graph traversal algorithms. The default (`link_expansion`) is optimized for fast retrieval with target latency under 100ms.
See [Configuration → Retrieval](./configuration#retrieval) for available algorithms and how to configure them.
---
## Next Steps
- [**Retain**](./retain) — How memories are stored with rich context
+2 -2
View File
@@ -133,8 +133,8 @@ const sidebars: SidebarsConfig = {
},
{
type: 'doc',
id: 'developer/metrics',
label: 'Metrics',
id: 'developer/monitoring',
label: 'Monitoring',
},
{
type: 'doc',
+419 -4
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": [
@@ -454,6 +520,22 @@
"type": "string",
"title": "Bank Id"
}
},
{
"name": "authorization",
"in": "header",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Authorization"
}
}
],
"responses": {
@@ -486,7 +568,7 @@
"Entities"
],
"summary": "List entities",
"description": "List all entities (people, organizations, etc.) known by the bank, ordered by mention count.",
"description": "List all entities (people, organizations, etc.) known by the bank, ordered by mention count. Supports pagination.",
"operationId": "list_entities",
"parameters": [
{
@@ -510,6 +592,18 @@
},
"description": "Maximum number of entities 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",
@@ -916,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": [
@@ -2191,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",
@@ -2212,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"
}
},
@@ -2407,11 +2614,26 @@
},
"type": "array",
"title": "Items"
},
"total": {
"type": "integer",
"title": "Total"
},
"limit": {
"type": "integer",
"title": "Limit"
},
"offset": {
"type": "integer",
"title": "Offset"
}
},
"type": "object",
"required": [
"items"
"items",
"total",
"limit",
"offset"
],
"title": "EntityListResponse",
"description": "Response model for entity list endpoint.",
@@ -2424,7 +2646,10 @@
"last_seen": "2024-02-01T14:00:00Z",
"mention_count": 15
}
]
],
"limit": 100,
"offset": 0,
"total": 150
}
},
"EntityObservationResponse": {
@@ -2706,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": {
@@ -2775,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",
@@ -2800,6 +3091,10 @@
"channel": "engineering",
"source": "slack"
},
"tags": [
"user_a",
"user_b"
],
"timestamp": "2024-01-15T10:30:00Z"
}
},
@@ -2953,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",
@@ -2971,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",
@@ -3192,6 +3518,20 @@
}
],
"title": "Chunk Id"
},
"tags": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Tags"
}
},
"type": "object",
@@ -3216,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"
}
@@ -3358,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",
@@ -3391,7 +3762,11 @@
"key_points"
],
"type": "object"
}
},
"tags": [
"user_a"
],
"tags_match": "any"
}
},
"ReflectResponse": {
@@ -3481,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",
@@ -3491,6 +3881,10 @@
"description": "Request model for retain endpoint.",
"example": {
"async": false,
"document_tags": [
"user_a",
"user_b"
],
"items": [
{
"content": "Alice works at Google",
@@ -3569,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"

Some files were not shown because too many files have changed in this diff Show More