Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3b343bc6aa | ||
|
|
1907a33edd | ||
|
|
de132501c6 | ||
|
|
a75dcfebf5 | ||
|
|
20c8f8b06a | ||
|
|
f5f3fca4ad | ||
|
|
d47c8a28cc | ||
|
|
1ffc2a418c | ||
|
|
fa53917c63 | ||
|
|
59913086be |
@@ -174,6 +174,25 @@ This runs the same checks as the pre-commit hook (Ruff for Python, ESLint/Pretti
|
||||
- Multi-bank queries are client responsibility to orchestrate
|
||||
- Disposition traits only affect reflect, not recall
|
||||
|
||||
### Control Plane API Routes
|
||||
|
||||
When adding or modifying parameters in the dataplane API (hindsight-api), you must also update the control plane routes that proxy to it:
|
||||
|
||||
1. **API Routes** (`hindsight-control-plane/src/app/api/`):
|
||||
- `recall/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/recall`
|
||||
- `reflect/route.ts` - proxies to `/v1/default/banks/{bank_id}/reflect`
|
||||
- `memories/retain/route.ts` - proxies to `/v1/default/banks/{bank_id}/memories/retain`
|
||||
- Other routes follow the same pattern
|
||||
|
||||
2. **Client types** (`hindsight-control-plane/src/lib/api.ts`):
|
||||
- Update the TypeScript type definitions for `recall()`, `reflect()`, `retain()` etc.
|
||||
|
||||
3. **Checklist when adding new API parameters**:
|
||||
- Add parameter extraction in the route handler (destructure from `body`)
|
||||
- Pass the parameter to the SDK call
|
||||
- Update the client type definition in `lib/api.ts`
|
||||
- Update any UI components that need to use the new parameter
|
||||
|
||||
### Python Style
|
||||
- Python 3.11+, type hints required
|
||||
- Async throughout (asyncpg, async FastAPI)
|
||||
|
||||
@@ -2,8 +2,8 @@ apiVersion: v2
|
||||
name: hindsight
|
||||
description: Hindsight helm chart
|
||||
type: application
|
||||
version: 0.2.1
|
||||
appVersion: "0.2.1"
|
||||
version: 0.3.0
|
||||
appVersion: "0.3.0"
|
||||
keywords:
|
||||
- ai
|
||||
- memory
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
"""add_tags_column
|
||||
|
||||
Revision ID: g2a3b4c5d6e7
|
||||
Revises: f1a2b3c4d5e6
|
||||
Create Date: 2025-01-13
|
||||
|
||||
Add tags column to memory_units and documents tables for visibility scoping.
|
||||
Tags enable filtering memories by scope (e.g., user IDs, session IDs) during recall/reflect.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "g2a3b4c5d6e7"
|
||||
down_revision: str | Sequence[str] | None = "f1a2b3c4d5e6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def _get_schema_prefix() -> str:
|
||||
"""Get schema prefix for table names (e.g., 'tenant_x.' or '' for public)."""
|
||||
schema = context.config.get_main_option("target_schema")
|
||||
return f'"{schema}".' if schema else ""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Add tags column to memory_units and documents tables."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
# Add tags column to memory_units table
|
||||
op.execute(f"ALTER TABLE {schema}memory_units ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
|
||||
|
||||
# Create GIN index for efficient array containment queries (tags && ARRAY['x'])
|
||||
op.execute(f"CREATE INDEX IF NOT EXISTS idx_memory_units_tags ON {schema}memory_units USING GIN (tags)")
|
||||
|
||||
# Add tags column to documents table for document-level tags
|
||||
op.execute(f"ALTER TABLE {schema}documents ADD COLUMN IF NOT EXISTS tags VARCHAR[] NOT NULL DEFAULT '{{}}'")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Remove tags columns and index."""
|
||||
schema = _get_schema_prefix()
|
||||
|
||||
op.execute(f"DROP INDEX IF EXISTS {schema}idx_memory_units_tags")
|
||||
op.execute(f"ALTER TABLE {schema}memory_units DROP COLUMN IF EXISTS tags")
|
||||
op.execute(f"ALTER TABLE {schema}documents DROP COLUMN IF EXISTS tags")
|
||||
@@ -37,6 +37,7 @@ from hindsight_api import MemoryEngine
|
||||
from hindsight_api.engine.db_utils import acquire_with_retry
|
||||
from hindsight_api.engine.memory_engine import Budget, fq_table
|
||||
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES, TokenUsage
|
||||
from hindsight_api.engine.search.tags import TagsMatch
|
||||
from hindsight_api.extensions import HttpExtension, OperationValidationError, load_extension
|
||||
from hindsight_api.metrics import create_metrics_collector, get_metrics_collector, initialize_metrics
|
||||
from hindsight_api.models import RequestContext
|
||||
@@ -81,6 +82,8 @@ class RecallRequest(BaseModel):
|
||||
"trace": True,
|
||||
"query_timestamp": "2023-05-30T23:40:00",
|
||||
"include": {"entities": {"max_tokens": 500}},
|
||||
"tags": ["user_a"],
|
||||
"tags_match": "any",
|
||||
}
|
||||
}
|
||||
)
|
||||
@@ -99,6 +102,15 @@ class RecallRequest(BaseModel):
|
||||
default_factory=IncludeOptions,
|
||||
description="Options for including additional data (entities are included by default)",
|
||||
)
|
||||
tags: list[str] | None = Field(
|
||||
default=None,
|
||||
description="Filter memories by tags. If not specified, all memories are returned.",
|
||||
)
|
||||
tags_match: TagsMatch = Field(
|
||||
default="any",
|
||||
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
|
||||
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
|
||||
)
|
||||
|
||||
|
||||
class RecallResult(BaseModel):
|
||||
@@ -119,6 +131,7 @@ class RecallResult(BaseModel):
|
||||
"document_id": "session_abc123",
|
||||
"metadata": {"source": "slack"},
|
||||
"chunk_id": "456e7890-e12b-34d5-a678-901234567890",
|
||||
"tags": ["user_a", "user_b"],
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -134,6 +147,7 @@ class RecallResult(BaseModel):
|
||||
document_id: str | None = None # Document this memory belongs to
|
||||
metadata: dict[str, str] | None = None # User-defined metadata
|
||||
chunk_id: str | None = None # Chunk this fact was extracted from
|
||||
tags: list[str] | None = None # Visibility scope tags
|
||||
|
||||
|
||||
class EntityObservationResponse(BaseModel):
|
||||
@@ -306,6 +320,7 @@ class MemoryItem(BaseModel):
|
||||
"metadata": {"source": "slack", "channel": "engineering"},
|
||||
"document_id": "meeting_notes_2024_01_15",
|
||||
"entities": [{"text": "Alice"}, {"text": "ML model", "type": "CONCEPT"}],
|
||||
"tags": ["user_a", "user_b"],
|
||||
}
|
||||
},
|
||||
)
|
||||
@@ -319,6 +334,10 @@ class MemoryItem(BaseModel):
|
||||
default=None,
|
||||
description="Optional entities to combine with auto-extracted entities.",
|
||||
)
|
||||
tags: list[str] | None = Field(
|
||||
default=None,
|
||||
description="Optional tags for visibility scoping. Memories with tags can be filtered during recall.",
|
||||
)
|
||||
|
||||
@field_validator("timestamp", mode="before")
|
||||
@classmethod
|
||||
@@ -353,6 +372,7 @@ class RetainRequest(BaseModel):
|
||||
},
|
||||
],
|
||||
"async": False,
|
||||
"document_tags": ["user_a", "user_b"],
|
||||
}
|
||||
}
|
||||
)
|
||||
@@ -363,6 +383,10 @@ class RetainRequest(BaseModel):
|
||||
alias="async",
|
||||
description="If true, process asynchronously in background. If false, wait for completion (default: false)",
|
||||
)
|
||||
document_tags: list[str] | None = Field(
|
||||
default=None,
|
||||
description="Tags applied to all items in this request. These are merged with any item-level tags.",
|
||||
)
|
||||
|
||||
|
||||
class RetainResponse(BaseModel):
|
||||
@@ -431,6 +455,8 @@ class ReflectRequest(BaseModel):
|
||||
},
|
||||
"required": ["summary", "key_points"],
|
||||
},
|
||||
"tags": ["user_a"],
|
||||
"tags_match": "any",
|
||||
}
|
||||
}
|
||||
)
|
||||
@@ -446,6 +472,15 @@ class ReflectRequest(BaseModel):
|
||||
default=None,
|
||||
description="Optional JSON Schema for structured output. When provided, the response will include a 'structured_output' field with the LLM response parsed according to this schema.",
|
||||
)
|
||||
tags: list[str] | None = Field(
|
||||
default=None,
|
||||
description="Filter memories by tags during reflection. If not specified, all memories are considered.",
|
||||
)
|
||||
tags_match: TagsMatch = Field(
|
||||
default="any",
|
||||
description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), "
|
||||
"'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
|
||||
)
|
||||
|
||||
|
||||
class OpinionItem(BaseModel):
|
||||
@@ -728,6 +763,37 @@ class ListDocumentsResponse(BaseModel):
|
||||
offset: int
|
||||
|
||||
|
||||
class TagItem(BaseModel):
|
||||
"""Single tag with usage count."""
|
||||
|
||||
tag: str = Field(description="The tag value")
|
||||
count: int = Field(description="Number of memories with this tag")
|
||||
|
||||
|
||||
class ListTagsResponse(BaseModel):
|
||||
"""Response model for list tags endpoint."""
|
||||
|
||||
model_config = ConfigDict(
|
||||
json_schema_extra={
|
||||
"example": {
|
||||
"items": [
|
||||
{"tag": "user:alice", "count": 42},
|
||||
{"tag": "user:bob", "count": 15},
|
||||
{"tag": "session:abc123", "count": 8},
|
||||
],
|
||||
"total": 25,
|
||||
"limit": 100,
|
||||
"offset": 0,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
items: list[TagItem]
|
||||
total: int
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
|
||||
class DocumentResponse(BaseModel):
|
||||
"""Response model for get document endpoint."""
|
||||
|
||||
@@ -741,6 +807,7 @@ class DocumentResponse(BaseModel):
|
||||
"created_at": "2024-01-15T10:30:00Z",
|
||||
"updated_at": "2024-01-15T10:30:00Z",
|
||||
"memory_unit_count": 15,
|
||||
"tags": ["user_a", "session_123"],
|
||||
}
|
||||
}
|
||||
)
|
||||
@@ -752,6 +819,7 @@ class DocumentResponse(BaseModel):
|
||||
created_at: str
|
||||
updated_at: str
|
||||
memory_unit_count: int
|
||||
tags: list[str] = Field(default_factory=list, description="Tags associated with this document")
|
||||
|
||||
|
||||
class DeleteDocumentResponse(BaseModel):
|
||||
@@ -1179,6 +1247,37 @@ def _register_routes(app: FastAPI):
|
||||
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/list: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.get(
|
||||
"/v1/default/banks/{bank_id}/memories/{memory_id}",
|
||||
summary="Get memory unit",
|
||||
description="Get a single memory unit by ID with all its metadata including entities and tags.",
|
||||
operation_id="get_memory",
|
||||
tags=["Memory"],
|
||||
)
|
||||
async def api_get_memory(
|
||||
bank_id: str,
|
||||
memory_id: str,
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""Get a single memory unit by ID."""
|
||||
try:
|
||||
data = await app.state.memory.get_memory_unit(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
request_context=request_context,
|
||||
)
|
||||
if data is None:
|
||||
raise HTTPException(status_code=404, detail=f"Memory unit '{memory_id}' not found")
|
||||
return data
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
|
||||
logger.error(f"Error in /v1/default/banks/{bank_id}/memories/{memory_id}: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.post(
|
||||
"/v1/default/banks/{bank_id}/memories/recall",
|
||||
response_model=RecallResponse,
|
||||
@@ -1243,6 +1342,8 @@ def _register_routes(app: FastAPI):
|
||||
include_chunks=include_chunks,
|
||||
max_chunk_tokens=max_chunk_tokens,
|
||||
request_context=request_context,
|
||||
tags=request.tags,
|
||||
tags_match=request.tags_match,
|
||||
)
|
||||
|
||||
# Convert core MemoryFact objects to API RecallResult objects (excluding internal metrics)
|
||||
@@ -1258,6 +1359,7 @@ def _register_routes(app: FastAPI):
|
||||
mentioned_at=fact.mentioned_at,
|
||||
document_id=fact.document_id,
|
||||
chunk_id=fact.chunk_id,
|
||||
tags=fact.tags,
|
||||
)
|
||||
for fact in core_result.results
|
||||
]
|
||||
@@ -1350,6 +1452,8 @@ def _register_routes(app: FastAPI):
|
||||
max_tokens=request.max_tokens,
|
||||
response_schema=request.response_schema,
|
||||
request_context=request_context,
|
||||
tags=request.tags,
|
||||
tags_match=request.tags_match,
|
||||
)
|
||||
|
||||
# Convert core MemoryFact objects to API ReflectFact objects if facts are requested
|
||||
@@ -1734,6 +1838,59 @@ def _register_routes(app: FastAPI):
|
||||
logger.error(f"Error in /v1/default/banks/{bank_id}/documents/{document_id}: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.get(
|
||||
"/v1/default/banks/{bank_id}/tags",
|
||||
response_model=ListTagsResponse,
|
||||
summary="List tags",
|
||||
description="List all unique tags in a memory bank with usage counts. "
|
||||
"Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.",
|
||||
operation_id="list_tags",
|
||||
tags=["Memory"],
|
||||
)
|
||||
async def api_list_tags(
|
||||
bank_id: str,
|
||||
q: str | None = Query(
|
||||
default=None,
|
||||
description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). "
|
||||
"Use '*' as wildcard. Case-insensitive.",
|
||||
),
|
||||
limit: int = Query(default=100, description="Maximum number of tags to return"),
|
||||
offset: int = Query(default=0, description="Offset for pagination"),
|
||||
request_context: RequestContext = Depends(get_request_context),
|
||||
):
|
||||
"""
|
||||
List all unique tags in a memory bank.
|
||||
|
||||
Use this endpoint to discover available tags or expand wildcard patterns.
|
||||
Supports '*' wildcards for flexible matching (case-insensitive):
|
||||
- 'user:*' matches user:alice, user:bob
|
||||
- '*-admin' matches role-admin, super-admin
|
||||
- 'env*-prod' matches env-prod, environment-prod
|
||||
|
||||
Args:
|
||||
bank_id: Memory Bank ID (from path)
|
||||
q: Wildcard pattern to filter tags (use '*' as wildcard)
|
||||
limit: Maximum number of tags to return (default: 100)
|
||||
offset: Offset for pagination (default: 0)
|
||||
"""
|
||||
try:
|
||||
data = await app.state.memory.list_tags(
|
||||
bank_id=bank_id,
|
||||
pattern=q,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
request_context=request_context,
|
||||
)
|
||||
return data
|
||||
except (AuthenticationError, HTTPException):
|
||||
raise
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
|
||||
logger.error(f"Error in /v1/default/banks/{bank_id}/tags: {error_detail}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@app.get(
|
||||
"/v1/default/chunks/{chunk_id:path}",
|
||||
response_model=ChunkResponse,
|
||||
@@ -2096,11 +2253,15 @@ def _register_routes(app: FastAPI):
|
||||
content_dict["document_id"] = item.document_id
|
||||
if item.entities:
|
||||
content_dict["entities"] = [{"text": e.text, "type": e.type or "CONCEPT"} for e in item.entities]
|
||||
if item.tags:
|
||||
content_dict["tags"] = item.tags
|
||||
contents.append(content_dict)
|
||||
|
||||
if request.async_:
|
||||
# Async processing: queue task and return immediately
|
||||
result = await app.state.memory.submit_async_retain(bank_id, contents, request_context=request_context)
|
||||
result = await app.state.memory.submit_async_retain(
|
||||
bank_id, contents, document_tags=request.document_tags, request_context=request_context
|
||||
)
|
||||
return RetainResponse.model_validate(
|
||||
{
|
||||
"success": True,
|
||||
@@ -2114,7 +2275,11 @@ def _register_routes(app: FastAPI):
|
||||
# Synchronous processing: wait for completion (record metrics)
|
||||
with metrics.record_operation("retain", bank_id=bank_id, source="api"):
|
||||
result, usage = await app.state.memory.retain_batch_async(
|
||||
bank_id=bank_id, contents=contents, request_context=request_context, return_usage=True
|
||||
bank_id=bank_id,
|
||||
contents=contents,
|
||||
document_tags=request.document_tags,
|
||||
request_context=request_context,
|
||||
return_usage=True,
|
||||
)
|
||||
|
||||
return RetainResponse.model_validate(
|
||||
|
||||
@@ -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,10 +44,19 @@ ENV_EMBEDDINGS_LOCAL_MODEL = "HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL"
|
||||
ENV_EMBEDDINGS_TEI_URL = "HINDSIGHT_API_EMBEDDINGS_TEI_URL"
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY = "HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY"
|
||||
ENV_EMBEDDINGS_OPENAI_MODEL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL"
|
||||
ENV_EMBEDDINGS_OPENAI_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL"
|
||||
|
||||
ENV_COHERE_API_KEY = "HINDSIGHT_API_COHERE_API_KEY"
|
||||
ENV_EMBEDDINGS_COHERE_MODEL = "HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL"
|
||||
ENV_EMBEDDINGS_COHERE_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL"
|
||||
ENV_RERANKER_COHERE_MODEL = "HINDSIGHT_API_RERANKER_COHERE_MODEL"
|
||||
ENV_RERANKER_COHERE_BASE_URL = "HINDSIGHT_API_RERANKER_COHERE_BASE_URL"
|
||||
|
||||
# LiteLLM gateway configuration (for embeddings and reranker via LiteLLM proxy)
|
||||
ENV_LITELLM_API_BASE = "HINDSIGHT_API_LITELLM_API_BASE"
|
||||
ENV_LITELLM_API_KEY = "HINDSIGHT_API_LITELLM_API_KEY"
|
||||
ENV_EMBEDDINGS_LITELLM_MODEL = "HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL"
|
||||
ENV_RERANKER_LITELLM_MODEL = "HINDSIGHT_API_RERANKER_LITELLM_MODEL"
|
||||
|
||||
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
|
||||
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
|
||||
@@ -59,11 +71,13 @@ 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"
|
||||
|
||||
@@ -120,14 +134,21 @@ 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 = "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
|
||||
@@ -180,6 +201,36 @@ 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()
|
||||
@@ -222,6 +273,8 @@ class HindsightConfig:
|
||||
embeddings_provider: str
|
||||
embeddings_local_model: str
|
||||
embeddings_tei_url: str | None
|
||||
embeddings_openai_base_url: str | None
|
||||
embeddings_cohere_base_url: str | None
|
||||
|
||||
# Reranker
|
||||
reranker_provider: str
|
||||
@@ -230,17 +283,20 @@ class HindsightConfig:
|
||||
reranker_tei_batch_size: int
|
||||
reranker_tei_max_concurrent: int
|
||||
reranker_max_candidates: int
|
||||
reranker_cohere_base_url: str | None
|
||||
|
||||
# Server
|
||||
host: str
|
||||
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
|
||||
@@ -297,6 +353,8 @@ class HindsightConfig:
|
||||
embeddings_provider=os.getenv(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER),
|
||||
embeddings_local_model=os.getenv(ENV_EMBEDDINGS_LOCAL_MODEL, DEFAULT_EMBEDDINGS_LOCAL_MODEL),
|
||||
embeddings_tei_url=os.getenv(ENV_EMBEDDINGS_TEI_URL),
|
||||
embeddings_openai_base_url=os.getenv(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None,
|
||||
embeddings_cohere_base_url=os.getenv(ENV_EMBEDDINGS_COHERE_BASE_URL) or None,
|
||||
# Reranker
|
||||
reranker_provider=os.getenv(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER),
|
||||
reranker_local_model=os.getenv(ENV_RERANKER_LOCAL_MODEL, DEFAULT_RERANKER_LOCAL_MODEL),
|
||||
@@ -306,15 +364,20 @@ class HindsightConfig:
|
||||
os.getenv(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT))
|
||||
),
|
||||
reranker_max_candidates=int(os.getenv(ENV_RERANKER_MAX_CANDIDATES, str(DEFAULT_RERANKER_MAX_CANDIDATES))),
|
||||
reranker_cohere_base_url=os.getenv(ENV_RERANKER_COHERE_BASE_URL) or None,
|
||||
# Server
|
||||
host=os.getenv(ENV_HOST, DEFAULT_HOST),
|
||||
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
|
||||
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",
|
||||
@@ -384,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)."""
|
||||
|
||||
@@ -15,18 +15,24 @@ from concurrent.futures import ThreadPoolExecutor
|
||||
import httpx
|
||||
|
||||
from ..config import (
|
||||
DEFAULT_LITELLM_API_BASE,
|
||||
DEFAULT_RERANKER_COHERE_MODEL,
|
||||
DEFAULT_RERANKER_FLASHRANK_CACHE_DIR,
|
||||
DEFAULT_RERANKER_FLASHRANK_MODEL,
|
||||
DEFAULT_RERANKER_LITELLM_MODEL,
|
||||
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
DEFAULT_RERANKER_LOCAL_MODEL,
|
||||
DEFAULT_RERANKER_PROVIDER,
|
||||
DEFAULT_RERANKER_TEI_BATCH_SIZE,
|
||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
|
||||
ENV_COHERE_API_KEY,
|
||||
ENV_LITELLM_API_BASE,
|
||||
ENV_LITELLM_API_KEY,
|
||||
ENV_RERANKER_COHERE_BASE_URL,
|
||||
ENV_RERANKER_COHERE_MODEL,
|
||||
ENV_RERANKER_FLASHRANK_CACHE_DIR,
|
||||
ENV_RERANKER_FLASHRANK_MODEL,
|
||||
ENV_RERANKER_LITELLM_MODEL,
|
||||
ENV_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||
ENV_RERANKER_LOCAL_MODEL,
|
||||
ENV_RERANKER_PROVIDER,
|
||||
@@ -392,6 +398,7 @@ class CohereCrossEncoder(CrossEncoderModel):
|
||||
self,
|
||||
api_key: str,
|
||||
model: str = DEFAULT_RERANKER_COHERE_MODEL,
|
||||
base_url: str | None = None,
|
||||
timeout: float = 60.0,
|
||||
):
|
||||
"""
|
||||
@@ -400,10 +407,12 @@ class CohereCrossEncoder(CrossEncoderModel):
|
||||
Args:
|
||||
api_key: Cohere API key
|
||||
model: Cohere rerank model name (default: rerank-english-v3.0)
|
||||
base_url: Custom base URL for Cohere-compatible API (e.g., Azure-hosted endpoint)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.base_url = base_url
|
||||
self.timeout = timeout
|
||||
self._client = None
|
||||
|
||||
@@ -421,8 +430,14 @@ class CohereCrossEncoder(CrossEncoderModel):
|
||||
except ImportError:
|
||||
raise ImportError("cohere is required for CohereCrossEncoder. Install it with: pip install cohere")
|
||||
|
||||
logger.info(f"Reranker: initializing Cohere provider with model {self.model}")
|
||||
self._client = cohere.Client(api_key=self.api_key, timeout=self.timeout)
|
||||
base_url_msg = f" at {self.base_url}" if self.base_url else ""
|
||||
logger.info(f"Reranker: initializing Cohere provider with model {self.model}{base_url_msg}")
|
||||
|
||||
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
|
||||
client_kwargs = {"api_key": self.api_key, "timeout": self.timeout}
|
||||
if self.base_url:
|
||||
client_kwargs["base_url"] = self.base_url
|
||||
self._client = cohere.Client(**client_kwargs)
|
||||
logger.info("Reranker: Cohere provider initialized")
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
@@ -641,6 +656,116 @@ class FlashRankCrossEncoder(CrossEncoderModel):
|
||||
return await loop.run_in_executor(FlashRankCrossEncoder._executor, self._predict_sync, pairs)
|
||||
|
||||
|
||||
class LiteLLMCrossEncoder(CrossEncoderModel):
|
||||
"""
|
||||
LiteLLM cross-encoder implementation using LiteLLM proxy's /rerank endpoint.
|
||||
|
||||
LiteLLM provides a unified interface for multiple reranking providers via
|
||||
the Cohere-compatible /rerank endpoint.
|
||||
See: https://docs.litellm.ai/docs/rerank
|
||||
|
||||
Supported providers via LiteLLM:
|
||||
- Cohere (rerank-english-v3.0, etc.) - prefix with cohere/
|
||||
- Together AI - prefix with together_ai/
|
||||
- Azure AI - prefix with azure_ai/
|
||||
- Jina AI - prefix with jina_ai/
|
||||
- AWS Bedrock - prefix with bedrock/
|
||||
- Voyage AI - prefix with voyage/
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_base: str = DEFAULT_LITELLM_API_BASE,
|
||||
api_key: str | None = None,
|
||||
model: str = DEFAULT_RERANKER_LITELLM_MODEL,
|
||||
timeout: float = 60.0,
|
||||
):
|
||||
"""
|
||||
Initialize LiteLLM cross-encoder client.
|
||||
|
||||
Args:
|
||||
api_base: Base URL of the LiteLLM proxy (default: http://localhost:4000)
|
||||
api_key: API key for the LiteLLM proxy (optional, depends on proxy config)
|
||||
model: Reranking model name (default: cohere/rerank-english-v3.0)
|
||||
Use provider prefix (e.g., cohere/, together_ai/, voyage/)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
"""
|
||||
self.api_base = api_base.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self._async_client: httpx.AsyncClient | None = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "litellm"
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize the async HTTP client."""
|
||||
if self._async_client is not None:
|
||||
return
|
||||
|
||||
logger.info(f"Reranker: initializing LiteLLM provider at {self.api_base} with model {self.model}")
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
|
||||
self._async_client = httpx.AsyncClient(timeout=self.timeout, headers=headers)
|
||||
logger.info("Reranker: LiteLLM provider initialized")
|
||||
|
||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||
"""
|
||||
Score query-document pairs using the LiteLLM proxy's /rerank endpoint.
|
||||
|
||||
Args:
|
||||
pairs: List of (query, document) tuples to score
|
||||
|
||||
Returns:
|
||||
List of relevance scores
|
||||
"""
|
||||
if self._async_client is None:
|
||||
raise RuntimeError("Reranker not initialized. Call initialize() first.")
|
||||
|
||||
if not pairs:
|
||||
return []
|
||||
|
||||
# Group pairs by query (LiteLLM rerank expects one query with multiple documents)
|
||||
query_groups: dict[str, list[tuple[int, str]]] = {}
|
||||
for idx, (query, text) in enumerate(pairs):
|
||||
if query not in query_groups:
|
||||
query_groups[query] = []
|
||||
query_groups[query].append((idx, text))
|
||||
|
||||
all_scores = [0.0] * len(pairs)
|
||||
|
||||
for query, indexed_texts in query_groups.items():
|
||||
texts = [text for _, text in indexed_texts]
|
||||
indices = [idx for idx, _ in indexed_texts]
|
||||
|
||||
# LiteLLM /rerank follows Cohere API format
|
||||
response = await self._async_client.post(
|
||||
f"{self.api_base}/rerank",
|
||||
json={
|
||||
"model": self.model,
|
||||
"query": query,
|
||||
"documents": texts,
|
||||
"top_n": len(texts), # Return all scores
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
# Map scores back to original positions
|
||||
# Response format: {"results": [{"index": 0, "relevance_score": 0.9}, ...]}
|
||||
for item in result.get("results", []):
|
||||
original_idx = item["index"]
|
||||
score = item.get("relevance_score", item.get("score", 0.0))
|
||||
all_scores[indices[original_idx]] = score
|
||||
|
||||
return all_scores
|
||||
|
||||
|
||||
def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
"""
|
||||
Create a CrossEncoderModel instance based on environment variables.
|
||||
@@ -671,14 +796,20 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||
if not api_key:
|
||||
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_RERANKER_PROVIDER} is 'cohere'")
|
||||
model = os.environ.get(ENV_RERANKER_COHERE_MODEL, DEFAULT_RERANKER_COHERE_MODEL)
|
||||
return CohereCrossEncoder(api_key=api_key, model=model)
|
||||
base_url = os.environ.get(ENV_RERANKER_COHERE_BASE_URL) or None
|
||||
return CohereCrossEncoder(api_key=api_key, model=model, base_url=base_url)
|
||||
elif provider == "flashrank":
|
||||
model = os.environ.get(ENV_RERANKER_FLASHRANK_MODEL, DEFAULT_RERANKER_FLASHRANK_MODEL)
|
||||
cache_dir = os.environ.get(ENV_RERANKER_FLASHRANK_CACHE_DIR, DEFAULT_RERANKER_FLASHRANK_CACHE_DIR)
|
||||
return FlashRankCrossEncoder(model_name=model, cache_dir=cache_dir)
|
||||
elif provider == "litellm":
|
||||
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
|
||||
api_key = os.environ.get(ENV_LITELLM_API_KEY)
|
||||
model = os.environ.get(ENV_RERANKER_LITELLM_MODEL, DEFAULT_RERANKER_LITELLM_MODEL)
|
||||
return LiteLLMCrossEncoder(api_base=api_base, api_key=api_key, model=model)
|
||||
elif provider == "rrf":
|
||||
return RRFPassthroughCrossEncoder()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'flashrank', 'rrf'"
|
||||
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei', 'cohere', 'flashrank', 'litellm', 'rrf'"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,284 @@
|
||||
"""
|
||||
Database connection budget management.
|
||||
|
||||
Limits concurrent database connections per operation to prevent
|
||||
a single operation (e.g., recall with parallel queries) from
|
||||
exhausting the connection pool.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, AsyncIterator
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import asyncpg
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class OperationBudget:
|
||||
"""
|
||||
Tracks connection budget for a single operation.
|
||||
|
||||
Each operation gets a semaphore limiting its concurrent connections.
|
||||
"""
|
||||
|
||||
operation_id: str
|
||||
max_connections: int
|
||||
semaphore: asyncio.Semaphore = field(init=False)
|
||||
active_count: int = field(default=0, init=False)
|
||||
|
||||
def __post_init__(self):
|
||||
self.semaphore = asyncio.Semaphore(self.max_connections)
|
||||
|
||||
|
||||
class ConnectionBudgetManager:
|
||||
"""
|
||||
Manages per-operation connection budgets.
|
||||
|
||||
Usage:
|
||||
manager = ConnectionBudgetManager(default_budget=4)
|
||||
|
||||
# Start an operation
|
||||
async with manager.operation(max_connections=2) as op:
|
||||
# Acquire connections within the budget
|
||||
async with op.acquire(pool) as conn:
|
||||
await conn.fetch(...)
|
||||
|
||||
# Multiple connections respect the budget
|
||||
async with op.acquire(pool) as conn1, op.acquire(pool) as conn2:
|
||||
# At most 2 concurrent connections for this operation
|
||||
...
|
||||
"""
|
||||
|
||||
def __init__(self, default_budget: int = 4):
|
||||
"""
|
||||
Initialize the budget manager.
|
||||
|
||||
Args:
|
||||
default_budget: Default max connections per operation
|
||||
"""
|
||||
self.default_budget = default_budget
|
||||
self._operations: dict[str, OperationBudget] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
@asynccontextmanager
|
||||
async def operation(
|
||||
self,
|
||||
max_connections: int | None = None,
|
||||
operation_id: str | None = None,
|
||||
) -> AsyncIterator["BudgetedOperation"]:
|
||||
"""
|
||||
Create a budgeted operation context.
|
||||
|
||||
Args:
|
||||
max_connections: Max concurrent connections for this operation.
|
||||
Defaults to manager's default_budget.
|
||||
operation_id: Optional custom operation ID. Auto-generated if not provided.
|
||||
|
||||
Yields:
|
||||
BudgetedOperation context for acquiring connections
|
||||
"""
|
||||
op_id = operation_id or f"op-{uuid.uuid4().hex[:12]}"
|
||||
budget = max_connections or self.default_budget
|
||||
|
||||
async with self._lock:
|
||||
if op_id in self._operations:
|
||||
raise ValueError(f"Operation {op_id} already exists")
|
||||
self._operations[op_id] = OperationBudget(op_id, budget)
|
||||
|
||||
try:
|
||||
yield BudgetedOperation(self, op_id)
|
||||
finally:
|
||||
async with self._lock:
|
||||
self._operations.pop(op_id, None)
|
||||
|
||||
def _get_budget(self, operation_id: str) -> OperationBudget:
|
||||
"""Get budget for an operation (internal use)."""
|
||||
budget = self._operations.get(operation_id)
|
||||
if not budget:
|
||||
raise ValueError(f"Operation {operation_id} not found")
|
||||
return budget
|
||||
|
||||
|
||||
class BudgetedOperation:
|
||||
"""
|
||||
A single operation with connection budget.
|
||||
|
||||
Provides methods to acquire connections within the budget.
|
||||
"""
|
||||
|
||||
def __init__(self, manager: ConnectionBudgetManager, operation_id: str):
|
||||
self._manager = manager
|
||||
self.operation_id = operation_id
|
||||
|
||||
@property
|
||||
def budget(self) -> OperationBudget:
|
||||
"""Get the budget for this operation."""
|
||||
return self._manager._get_budget(self.operation_id)
|
||||
|
||||
@asynccontextmanager
|
||||
async def acquire(self, pool: "asyncpg.Pool") -> AsyncIterator["asyncpg.Connection"]:
|
||||
"""
|
||||
Acquire a connection within the operation's budget.
|
||||
|
||||
Blocks if the operation has reached its connection limit.
|
||||
|
||||
Args:
|
||||
pool: asyncpg connection pool
|
||||
|
||||
Yields:
|
||||
Database connection
|
||||
"""
|
||||
budget = self.budget
|
||||
async with budget.semaphore:
|
||||
budget.active_count += 1
|
||||
conn = await pool.acquire()
|
||||
try:
|
||||
yield conn
|
||||
finally:
|
||||
budget.active_count -= 1
|
||||
await pool.release(conn)
|
||||
|
||||
def wrap_pool(self, pool: "asyncpg.Pool") -> "BudgetedPool":
|
||||
"""
|
||||
Wrap a pool with this operation's budget.
|
||||
|
||||
The returned BudgetedPool can be passed to functions expecting a pool,
|
||||
and all acquire() calls will be limited by this operation's budget.
|
||||
|
||||
Args:
|
||||
pool: asyncpg connection pool to wrap
|
||||
|
||||
Returns:
|
||||
BudgetedPool that limits connections to this operation's budget
|
||||
"""
|
||||
return BudgetedPool(pool, self)
|
||||
|
||||
async def acquire_many(
|
||||
self,
|
||||
pool: "asyncpg.Pool",
|
||||
count: int,
|
||||
) -> AsyncIterator[list["asyncpg.Connection"]]:
|
||||
"""
|
||||
Acquire multiple connections within the budget.
|
||||
|
||||
Note: This acquires connections sequentially to respect the budget.
|
||||
For parallel acquisition, use multiple acquire() calls with asyncio.gather().
|
||||
|
||||
Args:
|
||||
pool: asyncpg connection pool
|
||||
count: Number of connections to acquire
|
||||
|
||||
Yields:
|
||||
List of database connections
|
||||
"""
|
||||
connections = []
|
||||
try:
|
||||
for _ in range(count):
|
||||
conn = await pool.acquire()
|
||||
connections.append(conn)
|
||||
yield connections
|
||||
finally:
|
||||
for conn in connections:
|
||||
await pool.release(conn)
|
||||
|
||||
|
||||
# Global default manager instance
|
||||
_default_manager: ConnectionBudgetManager | None = None
|
||||
|
||||
|
||||
def get_budget_manager(default_budget: int = 4) -> ConnectionBudgetManager:
|
||||
"""
|
||||
Get or create the global budget manager.
|
||||
|
||||
Args:
|
||||
default_budget: Default max connections per operation
|
||||
|
||||
Returns:
|
||||
Global ConnectionBudgetManager instance
|
||||
"""
|
||||
global _default_manager
|
||||
if _default_manager is None:
|
||||
_default_manager = ConnectionBudgetManager(default_budget=default_budget)
|
||||
return _default_manager
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def budgeted_operation(
|
||||
max_connections: int | None = None,
|
||||
operation_id: str | None = None,
|
||||
default_budget: int = 4,
|
||||
) -> AsyncIterator[BudgetedOperation]:
|
||||
"""
|
||||
Convenience function to create a budgeted operation.
|
||||
|
||||
Args:
|
||||
max_connections: Max concurrent connections for this operation
|
||||
operation_id: Optional custom operation ID
|
||||
default_budget: Default budget if manager not yet created
|
||||
|
||||
Yields:
|
||||
BudgetedOperation context
|
||||
|
||||
Example:
|
||||
async with budgeted_operation(max_connections=2) as op:
|
||||
async with op.acquire(pool) as conn:
|
||||
await conn.fetch(...)
|
||||
"""
|
||||
manager = get_budget_manager(default_budget)
|
||||
async with manager.operation(max_connections, operation_id) as op:
|
||||
yield op
|
||||
|
||||
|
||||
class BudgetedPool:
|
||||
"""
|
||||
A pool wrapper that limits concurrent connection acquisitions.
|
||||
|
||||
This can be passed to functions expecting a pool, and acquire()
|
||||
calls will be limited by the budget semaphore.
|
||||
|
||||
Usage:
|
||||
async with budgeted_operation(max_connections=4) as op:
|
||||
budgeted_pool = op.wrap_pool(pool)
|
||||
# Pass budgeted_pool to functions that expect a pool
|
||||
await some_function(budgeted_pool, ...)
|
||||
"""
|
||||
|
||||
def __init__(self, pool: "asyncpg.Pool", operation: BudgetedOperation):
|
||||
self._pool = pool
|
||||
self._operation = operation
|
||||
|
||||
async def acquire(self) -> "asyncpg.Connection":
|
||||
"""
|
||||
Acquire a connection within the budget.
|
||||
|
||||
Note: Caller must release the connection when done.
|
||||
Prefer using as context manager via acquire_with_retry or op.acquire().
|
||||
"""
|
||||
budget = self._operation.budget
|
||||
await budget.semaphore.acquire()
|
||||
budget.active_count += 1
|
||||
try:
|
||||
return await self._pool.acquire()
|
||||
except Exception:
|
||||
budget.active_count -= 1
|
||||
budget.semaphore.release()
|
||||
raise
|
||||
|
||||
async def release(self, conn: "asyncpg.Connection") -> None:
|
||||
"""Release a connection back to the pool."""
|
||||
budget = self._operation.budget
|
||||
try:
|
||||
await self._pool.release(conn)
|
||||
finally:
|
||||
budget.active_count -= 1
|
||||
budget.semaphore.release()
|
||||
|
||||
def __getattr__(self, name):
|
||||
"""Proxy other attributes to the underlying pool."""
|
||||
return getattr(self._pool, name)
|
||||
@@ -17,16 +17,23 @@ import httpx
|
||||
|
||||
from ..config import (
|
||||
DEFAULT_EMBEDDINGS_COHERE_MODEL,
|
||||
DEFAULT_EMBEDDINGS_LITELLM_MODEL,
|
||||
DEFAULT_EMBEDDINGS_LOCAL_MODEL,
|
||||
DEFAULT_EMBEDDINGS_OPENAI_MODEL,
|
||||
DEFAULT_EMBEDDINGS_PROVIDER,
|
||||
DEFAULT_LITELLM_API_BASE,
|
||||
ENV_COHERE_API_KEY,
|
||||
ENV_EMBEDDINGS_COHERE_BASE_URL,
|
||||
ENV_EMBEDDINGS_COHERE_MODEL,
|
||||
ENV_EMBEDDINGS_LITELLM_MODEL,
|
||||
ENV_EMBEDDINGS_LOCAL_MODEL,
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY,
|
||||
ENV_EMBEDDINGS_OPENAI_BASE_URL,
|
||||
ENV_EMBEDDINGS_OPENAI_MODEL,
|
||||
ENV_EMBEDDINGS_PROVIDER,
|
||||
ENV_EMBEDDINGS_TEI_URL,
|
||||
ENV_LITELLM_API_BASE,
|
||||
ENV_LITELLM_API_KEY,
|
||||
ENV_LLM_API_KEY,
|
||||
)
|
||||
|
||||
@@ -322,6 +329,7 @@ class OpenAIEmbeddings(Embeddings):
|
||||
self,
|
||||
api_key: str,
|
||||
model: str = DEFAULT_EMBEDDINGS_OPENAI_MODEL,
|
||||
base_url: str | None = None,
|
||||
batch_size: int = 100,
|
||||
max_retries: int = 3,
|
||||
):
|
||||
@@ -331,11 +339,13 @@ class OpenAIEmbeddings(Embeddings):
|
||||
Args:
|
||||
api_key: OpenAI API key
|
||||
model: OpenAI embedding model name (default: text-embedding-3-small)
|
||||
base_url: Custom base URL for OpenAI-compatible API (e.g., Azure OpenAI endpoint)
|
||||
batch_size: Maximum batch size for embedding requests (default: 100)
|
||||
max_retries: Maximum number of retries for failed requests (default: 3)
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.base_url = base_url
|
||||
self.batch_size = batch_size
|
||||
self.max_retries = max_retries
|
||||
self._client = None
|
||||
@@ -361,8 +371,14 @@ class OpenAIEmbeddings(Embeddings):
|
||||
except ImportError:
|
||||
raise ImportError("openai is required for OpenAIEmbeddings. Install it with: pip install openai")
|
||||
|
||||
logger.info(f"Embeddings: initializing OpenAI provider with model {self.model}")
|
||||
self._client = OpenAI(api_key=self.api_key, max_retries=self.max_retries)
|
||||
base_url_msg = f" at {self.base_url}" if self.base_url else ""
|
||||
logger.info(f"Embeddings: initializing OpenAI provider with model {self.model}{base_url_msg}")
|
||||
|
||||
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
|
||||
client_kwargs = {"api_key": self.api_key, "max_retries": self.max_retries}
|
||||
if self.base_url:
|
||||
client_kwargs["base_url"] = self.base_url
|
||||
self._client = OpenAI(**client_kwargs)
|
||||
|
||||
# Try to get dimension from known models, otherwise do a test embedding
|
||||
if self.model in self.MODEL_DIMENSIONS:
|
||||
@@ -435,6 +451,7 @@ class CohereEmbeddings(Embeddings):
|
||||
self,
|
||||
api_key: str,
|
||||
model: str = DEFAULT_EMBEDDINGS_COHERE_MODEL,
|
||||
base_url: str | None = None,
|
||||
batch_size: int = 96,
|
||||
timeout: float = 60.0,
|
||||
input_type: str = "search_document",
|
||||
@@ -445,6 +462,7 @@ class CohereEmbeddings(Embeddings):
|
||||
Args:
|
||||
api_key: Cohere API key
|
||||
model: Cohere embedding model name (default: embed-english-v3.0)
|
||||
base_url: Custom base URL for Cohere-compatible API (e.g., Azure-hosted endpoint)
|
||||
batch_size: Maximum batch size for embedding requests (default: 96, Cohere's limit)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
input_type: Input type for embeddings (default: search_document).
|
||||
@@ -452,6 +470,7 @@ class CohereEmbeddings(Embeddings):
|
||||
"""
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.base_url = base_url
|
||||
self.batch_size = batch_size
|
||||
self.timeout = timeout
|
||||
self.input_type = input_type
|
||||
@@ -478,8 +497,14 @@ class CohereEmbeddings(Embeddings):
|
||||
except ImportError:
|
||||
raise ImportError("cohere is required for CohereEmbeddings. Install it with: pip install cohere")
|
||||
|
||||
logger.info(f"Embeddings: initializing Cohere provider with model {self.model}")
|
||||
self._client = cohere.Client(api_key=self.api_key, timeout=self.timeout)
|
||||
base_url_msg = f" at {self.base_url}" if self.base_url else ""
|
||||
logger.info(f"Embeddings: initializing Cohere provider with model {self.model}{base_url_msg}")
|
||||
|
||||
# Build client kwargs, only including base_url if set (for Azure or custom endpoints)
|
||||
client_kwargs = {"api_key": self.api_key, "timeout": self.timeout}
|
||||
if self.base_url:
|
||||
client_kwargs["base_url"] = self.base_url
|
||||
self._client = cohere.Client(**client_kwargs)
|
||||
|
||||
# Try to get dimension from known models, otherwise do a test embedding
|
||||
if self.model in self.MODEL_DIMENSIONS:
|
||||
@@ -529,6 +554,123 @@ class CohereEmbeddings(Embeddings):
|
||||
return all_embeddings
|
||||
|
||||
|
||||
class LiteLLMEmbeddings(Embeddings):
|
||||
"""
|
||||
LiteLLM embeddings implementation using LiteLLM proxy's /embeddings endpoint.
|
||||
|
||||
LiteLLM provides a unified interface for multiple embedding providers.
|
||||
The proxy exposes an OpenAI-compatible /embeddings endpoint.
|
||||
See: https://docs.litellm.ai/docs/embedding/supported_embedding
|
||||
|
||||
Supported providers via LiteLLM:
|
||||
- OpenAI (text-embedding-3-small, text-embedding-ada-002, etc.)
|
||||
- Cohere (embed-english-v3.0, etc.) - prefix with cohere/
|
||||
- Vertex AI (textembedding-gecko, etc.) - prefix with vertex_ai/
|
||||
- HuggingFace, Mistral, Voyage AI, etc.
|
||||
|
||||
The embedding dimension is auto-detected from the model at initialization.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_base: str = DEFAULT_LITELLM_API_BASE,
|
||||
api_key: str | None = None,
|
||||
model: str = DEFAULT_EMBEDDINGS_LITELLM_MODEL,
|
||||
batch_size: int = 100,
|
||||
timeout: float = 60.0,
|
||||
):
|
||||
"""
|
||||
Initialize LiteLLM embeddings client.
|
||||
|
||||
Args:
|
||||
api_base: Base URL of the LiteLLM proxy (default: http://localhost:4000)
|
||||
api_key: API key for the LiteLLM proxy (optional, depends on proxy config)
|
||||
model: Embedding model name (default: text-embedding-3-small)
|
||||
Use provider prefix for non-OpenAI models (e.g., cohere/embed-english-v3.0)
|
||||
batch_size: Maximum batch size for embedding requests (default: 100)
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
"""
|
||||
self.api_base = api_base.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.batch_size = batch_size
|
||||
self.timeout = timeout
|
||||
self._client: httpx.Client | None = None
|
||||
self._dimension: int | None = None
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
return "litellm"
|
||||
|
||||
@property
|
||||
def dimension(self) -> int:
|
||||
if self._dimension is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
return self._dimension
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""Initialize the HTTP client and detect embedding dimension."""
|
||||
if self._client is not None:
|
||||
return
|
||||
|
||||
logger.info(f"Embeddings: initializing LiteLLM provider at {self.api_base} with model {self.model}")
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
|
||||
self._client = httpx.Client(timeout=self.timeout, headers=headers)
|
||||
|
||||
# Do a test embedding to detect dimension
|
||||
try:
|
||||
response = self._client.post(
|
||||
f"{self.api_base}/embeddings",
|
||||
json={"model": self.model, "input": ["test"]},
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
if result.get("data") and len(result["data"]) > 0:
|
||||
self._dimension = len(result["data"][0]["embedding"])
|
||||
logger.info(f"Embeddings: LiteLLM provider initialized (model: {self.model}, dim: {self._dimension})")
|
||||
except httpx.HTTPError as e:
|
||||
raise RuntimeError(f"Failed to connect to LiteLLM proxy at {self.api_base}: {e}")
|
||||
|
||||
def encode(self, texts: list[str]) -> list[list[float]]:
|
||||
"""
|
||||
Generate embeddings using the LiteLLM proxy.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
Returns:
|
||||
List of embedding vectors
|
||||
"""
|
||||
if self._client is None:
|
||||
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
|
||||
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
all_embeddings = []
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, len(texts), self.batch_size):
|
||||
batch = texts[i : i + self.batch_size]
|
||||
|
||||
response = self._client.post(
|
||||
f"{self.api_base}/embeddings",
|
||||
json={"model": self.model, "input": batch},
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
# Sort by index to ensure correct order
|
||||
batch_embeddings = sorted(result["data"], key=lambda x: x["index"])
|
||||
all_embeddings.extend([e["embedding"] for e in batch_embeddings])
|
||||
|
||||
return all_embeddings
|
||||
|
||||
|
||||
def create_embeddings_from_env() -> Embeddings:
|
||||
"""
|
||||
Create an Embeddings instance based on environment variables.
|
||||
@@ -558,12 +700,21 @@ def create_embeddings_from_env() -> Embeddings:
|
||||
f"when {ENV_EMBEDDINGS_PROVIDER} is 'openai'"
|
||||
)
|
||||
model = os.environ.get(ENV_EMBEDDINGS_OPENAI_MODEL, DEFAULT_EMBEDDINGS_OPENAI_MODEL)
|
||||
return OpenAIEmbeddings(api_key=api_key, model=model)
|
||||
base_url = os.environ.get(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None
|
||||
return OpenAIEmbeddings(api_key=api_key, model=model, base_url=base_url)
|
||||
elif provider == "cohere":
|
||||
api_key = os.environ.get(ENV_COHERE_API_KEY)
|
||||
if not api_key:
|
||||
raise ValueError(f"{ENV_COHERE_API_KEY} is required when {ENV_EMBEDDINGS_PROVIDER} is 'cohere'")
|
||||
model = os.environ.get(ENV_EMBEDDINGS_COHERE_MODEL, DEFAULT_EMBEDDINGS_COHERE_MODEL)
|
||||
return CohereEmbeddings(api_key=api_key, model=model)
|
||||
base_url = os.environ.get(ENV_EMBEDDINGS_COHERE_BASE_URL) or None
|
||||
return CohereEmbeddings(api_key=api_key, model=model, base_url=base_url)
|
||||
elif provider == "litellm":
|
||||
api_base = os.environ.get(ENV_LITELLM_API_BASE, DEFAULT_LITELLM_API_BASE)
|
||||
api_key = os.environ.get(ENV_LITELLM_API_KEY)
|
||||
model = os.environ.get(ENV_EMBEDDINGS_LITELLM_MODEL, DEFAULT_EMBEDDINGS_LITELLM_MODEL)
|
||||
return LiteLLMEmbeddings(api_base=api_base, api_key=api_key, model=model)
|
||||
else:
|
||||
raise ValueError(f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei', 'openai', 'cohere'")
|
||||
raise ValueError(
|
||||
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei', 'openai', 'cohere', 'litellm'"
|
||||
)
|
||||
|
||||
@@ -19,6 +19,7 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
from ..config import get_config
|
||||
from ..metrics import get_metrics_collector
|
||||
from .db_budget import budgeted_operation
|
||||
|
||||
# Context variable for current schema (async-safe, per-task isolation)
|
||||
_current_schema: contextvars.ContextVar[str] = contextvars.ContextVar("current_schema", default="public")
|
||||
@@ -150,6 +151,7 @@ from .retain import bank_utils, embedding_utils
|
||||
from .retain.types import RetainContentDict
|
||||
from .search import observation_utils, think_utils
|
||||
from .search.reranking import CrossEncoderReranker
|
||||
from .search.tags import TagsMatch
|
||||
from .task_backend import AsyncIOQueueBackend, NoopTaskBackend, TaskBackend
|
||||
|
||||
|
||||
@@ -1058,6 +1060,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
document_id: str | None = None,
|
||||
fact_type_override: str | None = None,
|
||||
confidence_score: float | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
return_usage: bool = False,
|
||||
):
|
||||
"""
|
||||
@@ -1190,6 +1193,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch=i == 1, # Only upsert on first batch
|
||||
fact_type_override=fact_type_override,
|
||||
confidence_score=confidence_score,
|
||||
document_tags=document_tags,
|
||||
)
|
||||
all_results.extend(sub_results)
|
||||
total_usage = total_usage + sub_usage
|
||||
@@ -1208,6 +1212,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch=True,
|
||||
fact_type_override=fact_type_override,
|
||||
confidence_score=confidence_score,
|
||||
document_tags=document_tags,
|
||||
)
|
||||
|
||||
# Call post-operation hook if validator is configured
|
||||
@@ -1242,6 +1247,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch: bool = True,
|
||||
fact_type_override: str | None = None,
|
||||
confidence_score: float | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
) -> tuple[list[list[str]], "TokenUsage"]:
|
||||
"""
|
||||
Internal method for batch processing without chunking logic.
|
||||
@@ -1258,6 +1264,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch: Whether this is the first batch (for chunked operations, only delete on first batch)
|
||||
fact_type_override: Override fact type for all facts
|
||||
confidence_score: Confidence score for opinions
|
||||
document_tags: Tags applied to all items in this batch
|
||||
|
||||
Returns:
|
||||
Tuple of (unit ID lists, token usage for fact extraction)
|
||||
@@ -1282,6 +1289,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch=is_first_batch,
|
||||
fact_type_override=fact_type_override,
|
||||
confidence_score=confidence_score,
|
||||
document_tags=document_tags,
|
||||
)
|
||||
|
||||
def recall(
|
||||
@@ -1340,6 +1348,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
include_chunks: bool = False,
|
||||
max_chunk_tokens: int = 8192,
|
||||
request_context: "RequestContext",
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> RecallResultModel:
|
||||
"""
|
||||
Recall memories using N*4-way parallel retrieval (N fact types × 4 retrieval methods).
|
||||
@@ -1365,6 +1375,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_entity_tokens: Maximum tokens for entity observations (default 500)
|
||||
include_chunks: Whether to include raw chunks in the response
|
||||
max_chunk_tokens: Maximum tokens for chunks (default 8192)
|
||||
tags: Optional list of tags for visibility filtering (OR matching - returns
|
||||
memories that have at least one matching tag)
|
||||
|
||||
Returns:
|
||||
RecallResultModel containing:
|
||||
@@ -1437,6 +1449,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_chunk_tokens,
|
||||
request_context,
|
||||
semaphore_wait=semaphore_wait,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
break # Success - exit retry loop
|
||||
except Exception as e:
|
||||
@@ -1555,6 +1569,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_chunk_tokens: int = 8192,
|
||||
request_context: "RequestContext" = None,
|
||||
semaphore_wait: float = 0.0,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> RecallResultModel:
|
||||
"""
|
||||
Search implementation with modular retrieval and reranking.
|
||||
@@ -1584,7 +1600,9 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
# Initialize tracer if requested
|
||||
from .search.tracer import SearchTracer
|
||||
|
||||
tracer = SearchTracer(query, thinking_budget, max_tokens) if enable_trace else None
|
||||
tracer = (
|
||||
SearchTracer(query, thinking_budget, max_tokens, tags=tags, tags_match=tags_match) if enable_trace else None
|
||||
)
|
||||
if tracer:
|
||||
tracer.start()
|
||||
|
||||
@@ -1594,8 +1612,9 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
# Buffer logs for clean output in concurrent scenarios
|
||||
recall_id = f"{bank_id[:8]}-{int(time.time() * 1000) % 100000}"
|
||||
log_buffer = []
|
||||
tags_info = f", tags={tags}, tags_match={tags_match}" if tags else ""
|
||||
log_buffer.append(
|
||||
f"[RECALL {recall_id}] Query: '{query[:50]}...' (budget={thinking_budget}, max_tokens={max_tokens})"
|
||||
f"[RECALL {recall_id}] Query: '{query[:50]}...' (budget={thinking_budget}, max_tokens={max_tokens}{tags_info})"
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1609,38 +1628,42 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
tracer.record_query_embedding(query_embedding)
|
||||
tracer.add_phase_metric("generate_query_embedding", step_duration)
|
||||
|
||||
# Step 2: N*4-Way Parallel Retrieval (N fact types × 4 retrieval methods)
|
||||
# Step 2: Optimized parallel retrieval using batched queries
|
||||
# - Semantic + BM25 combined in 1 CTE query for ALL fact types
|
||||
# - Graph runs per fact type (complex traversal)
|
||||
# - Temporal runs per fact type (if constraint detected)
|
||||
step_start = time.time()
|
||||
query_embedding_str = str(query_embedding)
|
||||
|
||||
from .search.retrieval import get_default_graph_retriever, retrieve_parallel
|
||||
from .search.retrieval import (
|
||||
get_default_graph_retriever,
|
||||
retrieve_all_fact_types_parallel,
|
||||
)
|
||||
|
||||
# Track each retrieval start time
|
||||
retrieval_start = time.time()
|
||||
|
||||
# Temporal extraction now runs IN PARALLEL with other retrievals inside retrieve_parallel
|
||||
# This prevents slow dateparser from blocking semantic/BM25/graph retrieval
|
||||
tc_duration = 0.0 # Will be tracked inside temporal retrieval timing
|
||||
|
||||
# Run retrieval for each fact type in parallel
|
||||
# Each retrieve_parallel uses ~4 connections, so 3 fact types = ~12 concurrent connections
|
||||
retrieval_tasks = [
|
||||
retrieve_parallel(
|
||||
pool,
|
||||
# Run optimized retrieval with connection budget
|
||||
config = get_config()
|
||||
async with budgeted_operation(
|
||||
max_connections=config.recall_connection_budget,
|
||||
operation_id=f"recall-{recall_id}",
|
||||
) as op:
|
||||
budgeted_pool = op.wrap_pool(pool)
|
||||
parallel_start = time.time()
|
||||
multi_result = await retrieve_all_fact_types_parallel(
|
||||
budgeted_pool,
|
||||
query,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
ft,
|
||||
fact_type, # Pass all fact types at once
|
||||
thinking_budget,
|
||||
question_date,
|
||||
self.query_analyzer,
|
||||
temporal_constraint=None, # Extracted in parallel inside retrieve_parallel
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
for ft in fact_type
|
||||
]
|
||||
parallel_start = time.time()
|
||||
all_retrievals = await asyncio.gather(*retrieval_tasks)
|
||||
parallel_duration = time.time() - parallel_start
|
||||
parallel_duration = time.time() - parallel_start
|
||||
|
||||
# Combine all results from all fact types and aggregate timings
|
||||
semantic_results = []
|
||||
@@ -1657,12 +1680,15 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
all_mpfp_timings = []
|
||||
|
||||
detected_temporal_constraint = None
|
||||
max_conn_wait = 0.0
|
||||
for idx, retrieval_result in enumerate(all_retrievals):
|
||||
max_conn_wait = multi_result.max_conn_wait
|
||||
for ft in fact_type:
|
||||
retrieval_result = multi_result.results_by_fact_type.get(ft)
|
||||
if not retrieval_result:
|
||||
continue
|
||||
|
||||
# Log fact types in this retrieval batch
|
||||
ft_name = fact_type[idx] if idx < len(fact_type) else "unknown"
|
||||
logger.debug(
|
||||
f"[RECALL {recall_id}] Fact type '{ft_name}': semantic={len(retrieval_result.semantic)}, bm25={len(retrieval_result.bm25)}, graph={len(retrieval_result.graph)}, temporal={len(retrieval_result.temporal) if retrieval_result.temporal else 0}"
|
||||
f"[RECALL {recall_id}] Fact type '{ft}': semantic={len(retrieval_result.semantic)}, bm25={len(retrieval_result.bm25)}, graph={len(retrieval_result.graph)}, temporal={len(retrieval_result.temporal) if retrieval_result.temporal else 0}"
|
||||
)
|
||||
|
||||
semantic_results.extend(retrieval_result.semantic)
|
||||
@@ -1678,8 +1704,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
detected_temporal_constraint = retrieval_result.temporal_constraint
|
||||
# Collect MPFP timings
|
||||
all_mpfp_timings.extend(retrieval_result.mpfp_timings)
|
||||
# Track max connection wait
|
||||
max_conn_wait = max(max_conn_wait, retrieval_result.max_conn_wait)
|
||||
|
||||
# If no temporal results from any fact type, set to None
|
||||
if not temporal_results:
|
||||
@@ -1711,10 +1735,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
temporal_count = len(temporal_results) if temporal_results else 0
|
||||
timing_parts.append(f"temporal={temporal_count}({aggregated_timings['temporal']:.3f}s)")
|
||||
temporal_info = f" | temporal_range={start_dt.strftime('%Y-%m-%d')} to {end_dt.strftime('%Y-%m-%d')}"
|
||||
# Only tc is sequential setup now (adjacency loads in parallel with retrieval)
|
||||
setup_info = f", tc={tc_duration:.3f}s" if tc_duration > 0.01 else ""
|
||||
log_buffer.append(
|
||||
f" [2] Parallel retrieval ({len(fact_type)} fact_types): {', '.join(timing_parts)} in {parallel_duration:.3f}s{setup_info}{temporal_info}"
|
||||
f" [2] Parallel retrieval ({len(fact_type)} fact_types): {', '.join(timing_parts)} in {parallel_duration:.3f}s{temporal_info}"
|
||||
)
|
||||
|
||||
# Log graph retriever timing breakdown if available
|
||||
@@ -1744,6 +1766,11 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
f"edges={hd.get('edges_loaded', 0)}"
|
||||
)
|
||||
|
||||
# Record temporal constraint in tracer if detected
|
||||
if tracer and detected_temporal_constraint:
|
||||
start_dt, end_dt = detected_temporal_constraint
|
||||
tracer.record_temporal_constraint(start_dt, end_dt)
|
||||
|
||||
# Record retrieval results for tracer - per fact type
|
||||
if tracer:
|
||||
# Convert RetrievalResult to old tuple format for tracer
|
||||
@@ -1751,8 +1778,10 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
return [(r.id, r.__dict__) for r in results]
|
||||
|
||||
# Add retrieval results per fact type (to show parallel execution in UI)
|
||||
for idx, rr in enumerate(all_retrievals):
|
||||
ft_name = fact_type[idx] if idx < len(fact_type) else "unknown"
|
||||
for ft_name in fact_type:
|
||||
rr = multi_result.results_by_fact_type.get(ft_name)
|
||||
if not rr:
|
||||
continue
|
||||
|
||||
# Add semantic retrieval results for this fact type
|
||||
tracer.add_retrieval_results(
|
||||
@@ -1784,14 +1813,22 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
fact_type=ft_name,
|
||||
)
|
||||
|
||||
# Add temporal retrieval results for this fact type (even if empty, to show it ran)
|
||||
if rr.temporal is not None:
|
||||
# Add temporal retrieval results for this fact type
|
||||
# Show temporal even with 0 results if constraint was detected
|
||||
if rr.temporal is not None or rr.temporal_constraint is not None:
|
||||
temporal_metadata = {"budget": thinking_budget}
|
||||
if rr.temporal_constraint:
|
||||
start_dt, end_dt = rr.temporal_constraint
|
||||
temporal_metadata["constraint"] = {
|
||||
"start": start_dt.isoformat() if start_dt else None,
|
||||
"end": end_dt.isoformat() if end_dt else None,
|
||||
}
|
||||
tracer.add_retrieval_results(
|
||||
method_name="temporal",
|
||||
results=to_tuple_format(rr.temporal),
|
||||
results=to_tuple_format(rr.temporal or []),
|
||||
duration_seconds=rr.timings.get("temporal", 0.0),
|
||||
score_field="temporal_score",
|
||||
metadata={"budget": thinking_budget},
|
||||
metadata=temporal_metadata,
|
||||
fact_type=ft_name,
|
||||
)
|
||||
|
||||
@@ -2051,6 +2088,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
mentioned_at=result_dict.get("mentioned_at"),
|
||||
document_id=result_dict.get("document_id"),
|
||||
chunk_id=result_dict.get("chunk_id"),
|
||||
tags=result_dict.get("tags"),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -2266,11 +2304,11 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
doc = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT d.id, d.bank_id, d.original_text, d.content_hash,
|
||||
d.created_at, d.updated_at, COUNT(mu.id) as unit_count
|
||||
d.created_at, d.updated_at, d.tags, COUNT(mu.id) as unit_count
|
||||
FROM {fq_table("documents")} d
|
||||
LEFT JOIN {fq_table("memory_units")} mu ON mu.document_id = d.id
|
||||
WHERE d.id = $1 AND d.bank_id = $2
|
||||
GROUP BY d.id, d.bank_id, d.original_text, d.content_hash, d.created_at, d.updated_at
|
||||
GROUP BY d.id, d.bank_id, d.original_text, d.content_hash, d.created_at, d.updated_at, d.tags
|
||||
""",
|
||||
document_id,
|
||||
bank_id,
|
||||
@@ -2287,6 +2325,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
"memory_unit_count": doc["unit_count"],
|
||||
"created_at": doc["created_at"].isoformat() if doc["created_at"] else None,
|
||||
"updated_at": doc["updated_at"].isoformat() if doc["updated_at"] else None,
|
||||
"tags": list(doc["tags"]) if doc["tags"] else [],
|
||||
}
|
||||
|
||||
async def delete_document(
|
||||
@@ -2775,6 +2814,68 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
|
||||
return {"items": items, "total": total, "limit": limit, "offset": offset}
|
||||
|
||||
async def get_memory_unit(
|
||||
self,
|
||||
bank_id: str,
|
||||
memory_id: str,
|
||||
request_context: "RequestContext",
|
||||
):
|
||||
"""
|
||||
Get a single memory unit by ID.
|
||||
|
||||
Args:
|
||||
bank_id: Bank ID
|
||||
memory_id: Memory unit ID
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with memory unit data or None if not found
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
pool = await self._get_pool()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Get the memory unit
|
||||
row = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, fact_type, document_id, chunk_id, tags
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = $1 AND bank_id = $2
|
||||
""",
|
||||
memory_id,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
if not row:
|
||||
return None
|
||||
|
||||
# Get entity information
|
||||
entities_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT e.canonical_name
|
||||
FROM {fq_table("unit_entities")} ue
|
||||
JOIN {fq_table("entities")} e ON ue.entity_id = e.id
|
||||
WHERE ue.unit_id = $1
|
||||
""",
|
||||
row["id"],
|
||||
)
|
||||
entities = [r["canonical_name"] for r in entities_rows]
|
||||
|
||||
return {
|
||||
"id": str(row["id"]),
|
||||
"text": row["text"],
|
||||
"context": row["context"] if row["context"] else "",
|
||||
"date": row["event_date"].isoformat() if row["event_date"] else "",
|
||||
"type": row["fact_type"],
|
||||
"mentioned_at": row["mentioned_at"].isoformat() if row["mentioned_at"] else None,
|
||||
"occurred_start": row["occurred_start"].isoformat() if row["occurred_start"] else None,
|
||||
"occurred_end": row["occurred_end"].isoformat() if row["occurred_end"] else None,
|
||||
"entities": entities,
|
||||
"document_id": row["document_id"] if row["document_id"] else None,
|
||||
"chunk_id": str(row["chunk_id"]) if row["chunk_id"] else None,
|
||||
"tags": row["tags"] if row["tags"] else [],
|
||||
}
|
||||
|
||||
async def list_documents(
|
||||
self,
|
||||
bank_id: str,
|
||||
@@ -3298,6 +3399,8 @@ Guidelines:
|
||||
max_tokens: int = 4096,
|
||||
response_schema: dict | None = None,
|
||||
request_context: "RequestContext",
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> ReflectResult:
|
||||
"""
|
||||
Reflect and formulate an answer using bank identity, world facts, and opinions.
|
||||
@@ -3365,6 +3468,8 @@ Guidelines:
|
||||
fact_type=["experience", "world", "opinion"],
|
||||
include_entities=True,
|
||||
request_context=request_context,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
recall_time = time.time() - recall_start
|
||||
|
||||
@@ -3678,7 +3783,7 @@ Guidelines:
|
||||
SELECT id, canonical_name, mention_count, first_seen, last_seen, metadata
|
||||
FROM {fq_table("entities")}
|
||||
WHERE bank_id = $1
|
||||
ORDER BY mention_count DESC, last_seen DESC
|
||||
ORDER BY mention_count DESC, last_seen DESC, id ASC
|
||||
LIMIT $2 OFFSET $3
|
||||
""",
|
||||
bank_id,
|
||||
@@ -3717,6 +3822,85 @@ Guidelines:
|
||||
"offset": offset,
|
||||
}
|
||||
|
||||
async def list_tags(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
pattern: str | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List all unique tags for a bank with usage counts.
|
||||
|
||||
Use this to discover available tags or expand wildcard patterns.
|
||||
Supports '*' as wildcard for flexible matching (case-insensitive):
|
||||
- 'user:*' matches user:alice, user:bob
|
||||
- '*-admin' matches role-admin, super-admin
|
||||
- 'env*-prod' matches env-prod, environment-prod
|
||||
|
||||
Args:
|
||||
bank_id: Bank identifier
|
||||
pattern: Wildcard pattern to filter tags (use '*' as wildcard, case-insensitive)
|
||||
limit: Maximum number of tags to return
|
||||
offset: Offset for pagination
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with items (list of {tag, count}), total, limit, offset
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
pool = await self._get_pool()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Build pattern filter if provided (convert * to % for ILIKE)
|
||||
pattern_clause = ""
|
||||
params: list[Any] = [bank_id]
|
||||
if pattern:
|
||||
# Convert wildcard pattern: * -> % for SQL ILIKE
|
||||
sql_pattern = pattern.replace("*", "%")
|
||||
pattern_clause = "AND tag ILIKE $2"
|
||||
params.append(sql_pattern)
|
||||
|
||||
# Get total count of distinct tags matching pattern
|
||||
total_row = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT COUNT(DISTINCT tag) as total
|
||||
FROM {fq_table("memory_units")}, unnest(tags) AS tag
|
||||
WHERE bank_id = $1 AND tags IS NOT NULL AND tags != '{{}}'
|
||||
{pattern_clause}
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
total = total_row["total"] if total_row else 0
|
||||
|
||||
# Get paginated tags with counts, ordered by frequency
|
||||
limit_param = len(params) + 1
|
||||
offset_param = len(params) + 2
|
||||
params.extend([limit, offset])
|
||||
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT tag, COUNT(*) as count
|
||||
FROM {fq_table("memory_units")}, unnest(tags) AS tag
|
||||
WHERE bank_id = $1 AND tags IS NOT NULL AND tags != '{{}}'
|
||||
{pattern_clause}
|
||||
GROUP BY tag
|
||||
ORDER BY count DESC, tag ASC
|
||||
LIMIT ${limit_param} OFFSET ${offset_param}
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
|
||||
items = [{"tag": row["tag"], "count": row["count"]} for row in rows]
|
||||
|
||||
return {
|
||||
"items": items,
|
||||
"total": total,
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
}
|
||||
|
||||
async def get_entity_state(
|
||||
self,
|
||||
bank_id: str,
|
||||
@@ -4361,6 +4545,7 @@ Guidelines:
|
||||
contents: list[dict[str, Any]],
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
document_tags: list[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Submit a batch retain operation to run asynchronously."""
|
||||
await self._authenticate_tenant(request_context)
|
||||
@@ -4384,14 +4569,16 @@ Guidelines:
|
||||
)
|
||||
|
||||
# Submit task to background queue
|
||||
await self._task_backend.submit_task(
|
||||
{
|
||||
"type": "batch_retain",
|
||||
"operation_id": str(operation_id),
|
||||
"bank_id": bank_id,
|
||||
"contents": contents,
|
||||
}
|
||||
)
|
||||
task_payload = {
|
||||
"type": "batch_retain",
|
||||
"operation_id": str(operation_id),
|
||||
"bank_id": bank_id,
|
||||
"contents": contents,
|
||||
}
|
||||
if document_tags:
|
||||
task_payload["document_tags"] = document_tags
|
||||
|
||||
await self._task_backend.submit_task(task_payload)
|
||||
|
||||
logger.info(f"Retain task queued for bank_id={bank_id}, {len(contents)} items, operation_id={operation_id}")
|
||||
|
||||
|
||||
@@ -85,6 +85,7 @@ class MemoryFact(BaseModel):
|
||||
"metadata": {"source": "slack"},
|
||||
"chunk_id": "bank123_session_abc123_0",
|
||||
"activation": 0.95,
|
||||
"tags": ["user_a", "session_123"],
|
||||
}
|
||||
}
|
||||
)
|
||||
@@ -102,6 +103,7 @@ class MemoryFact(BaseModel):
|
||||
chunk_id: str | None = Field(
|
||||
None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)"
|
||||
)
|
||||
tags: list[str] | None = Field(None, description="Visibility scope tags associated with this fact")
|
||||
|
||||
|
||||
class ChunkInfo(BaseModel):
|
||||
|
||||
@@ -1268,6 +1268,7 @@ async def extract_facts_from_contents(
|
||||
# mentioned_at: always the event_date (when the conversation/document occurred)
|
||||
mentioned_at=content.event_date,
|
||||
metadata=content.metadata,
|
||||
tags=content.tags,
|
||||
)
|
||||
|
||||
extracted_facts.append(extracted_fact)
|
||||
|
||||
@@ -45,6 +45,7 @@ async def insert_facts_batch(
|
||||
metadata_jsons = []
|
||||
chunk_ids = []
|
||||
document_ids = []
|
||||
tags_list = []
|
||||
|
||||
for fact in facts:
|
||||
fact_texts.append(fact.fact_text)
|
||||
@@ -65,16 +66,31 @@ async def insert_facts_batch(
|
||||
chunk_ids.append(fact.chunk_id)
|
||||
# Use per-fact document_id if available, otherwise fallback to batch-level document_id
|
||||
document_ids.append(fact.document_id if fact.document_id else document_id)
|
||||
# Convert tags to JSON string for proper batch insertion (PostgreSQL unnest doesn't handle 2D arrays well)
|
||||
tags_list.append(json.dumps(fact.tags if fact.tags else []))
|
||||
|
||||
# Batch insert all facts
|
||||
# Note: tags are passed as JSON strings and converted back to varchar[] via jsonb_array_elements_text + array_agg
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id)
|
||||
SELECT $1, * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[]
|
||||
WITH input_data AS (
|
||||
SELECT * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[], $15::jsonb[]
|
||||
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags_json)
|
||||
)
|
||||
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags)
|
||||
SELECT
|
||||
$1,
|
||||
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id,
|
||||
COALESCE(
|
||||
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
|
||||
'{{}}'::varchar[]
|
||||
)
|
||||
FROM input_data
|
||||
RETURNING id
|
||||
""",
|
||||
bank_id,
|
||||
@@ -91,6 +107,7 @@ async def insert_facts_batch(
|
||||
metadata_jsons,
|
||||
chunk_ids,
|
||||
document_ids,
|
||||
tags_list,
|
||||
)
|
||||
|
||||
unit_ids = [str(row["id"]) for row in results]
|
||||
@@ -121,7 +138,13 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
|
||||
|
||||
|
||||
async def handle_document_tracking(
|
||||
conn, bank_id: str, document_id: str, combined_content: str, is_first_batch: bool, retain_params: dict | None = None
|
||||
conn,
|
||||
bank_id: str,
|
||||
document_id: str,
|
||||
combined_content: str,
|
||||
is_first_batch: bool,
|
||||
retain_params: dict | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handle document tracking in the database.
|
||||
@@ -133,6 +156,7 @@ async def handle_document_tracking(
|
||||
combined_content: Combined content text from all content items
|
||||
is_first_batch: Whether this is the first batch (for chunked operations)
|
||||
retain_params: Optional parameters passed during retain (context, event_date, etc.)
|
||||
document_tags: Optional list of tags to associate with the document
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
@@ -149,13 +173,14 @@ async def handle_document_tracking(
|
||||
# Insert document (or update if exists from concurrent operations)
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params)
|
||||
VALUES ($1, $2, $3, $4, $5, $6)
|
||||
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params, tags)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
||||
ON CONFLICT (id, bank_id) DO UPDATE
|
||||
SET original_text = EXCLUDED.original_text,
|
||||
content_hash = EXCLUDED.content_hash,
|
||||
metadata = EXCLUDED.metadata,
|
||||
retain_params = EXCLUDED.retain_params,
|
||||
tags = EXCLUDED.tags,
|
||||
updated_at = NOW()
|
||||
""",
|
||||
document_id,
|
||||
@@ -164,4 +189,5 @@ async def handle_document_tracking(
|
||||
content_hash,
|
||||
json.dumps({}), # Empty metadata dict
|
||||
json.dumps(retain_params) if retain_params else None,
|
||||
document_tags or [],
|
||||
)
|
||||
|
||||
@@ -49,6 +49,7 @@ async def retain_batch(
|
||||
is_first_batch: bool = True,
|
||||
fact_type_override: str | None = None,
|
||||
confidence_score: float | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
) -> tuple[list[list[str]], TokenUsage]:
|
||||
"""
|
||||
Process a batch of content through the retain pipeline.
|
||||
@@ -67,6 +68,7 @@ async def retain_batch(
|
||||
is_first_batch: Whether this is the first batch
|
||||
fact_type_override: Override fact type for all facts
|
||||
confidence_score: Confidence score for opinions
|
||||
document_tags: Tags applied to all items in this batch
|
||||
|
||||
Returns:
|
||||
Tuple of (unit ID lists, token usage for fact extraction)
|
||||
@@ -88,12 +90,16 @@ async def retain_batch(
|
||||
# Convert dicts to RetainContent objects
|
||||
contents = []
|
||||
for item in contents_dicts:
|
||||
# Merge item-level tags with document-level tags
|
||||
item_tags = item.get("tags", []) or []
|
||||
merged_tags = list(set(item_tags + (document_tags or [])))
|
||||
content = RetainContent(
|
||||
content=item["content"],
|
||||
context=item.get("context", ""),
|
||||
event_date=item.get("event_date") or utcnow(),
|
||||
metadata=item.get("metadata", {}),
|
||||
entities=item.get("entities", []),
|
||||
tags=merged_tags,
|
||||
)
|
||||
contents.append(content)
|
||||
|
||||
@@ -131,7 +137,7 @@ async def retain_batch(
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
else:
|
||||
# Check for per-item document_ids
|
||||
@@ -159,7 +165,7 @@ async def retain_batch(
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params
|
||||
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
|
||||
total_time = time.time() - start_time
|
||||
@@ -225,7 +231,7 @@ async def retain_batch(
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
document_ids_added.append(document_id)
|
||||
doc_id_mapping[None] = document_id # For backwards compatibility
|
||||
@@ -269,7 +275,13 @@ async def retain_batch(
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, actual_doc_id, combined_content, is_first_batch, retain_params
|
||||
conn,
|
||||
bank_id,
|
||||
actual_doc_id,
|
||||
combined_content,
|
||||
is_first_batch,
|
||||
retain_params,
|
||||
document_tags,
|
||||
)
|
||||
document_ids_added.append(actual_doc_id)
|
||||
|
||||
|
||||
@@ -21,6 +21,7 @@ class RetainContentDict(TypedDict, total=False):
|
||||
metadata: Custom key-value metadata (optional)
|
||||
document_id: Document ID for this content item (optional)
|
||||
entities: User-provided entities to merge with extracted entities (optional)
|
||||
tags: Visibility scope tags for this content item (optional)
|
||||
"""
|
||||
|
||||
content: str # Required
|
||||
@@ -29,6 +30,7 @@ class RetainContentDict(TypedDict, total=False):
|
||||
metadata: dict[str, str]
|
||||
document_id: str
|
||||
entities: list[dict[str, str]] # [{"text": "...", "type": "..."}]
|
||||
tags: list[str] # Visibility scope tags
|
||||
|
||||
|
||||
def _now_utc() -> datetime:
|
||||
@@ -49,6 +51,7 @@ class RetainContent:
|
||||
event_date: datetime = field(default_factory=_now_utc)
|
||||
metadata: dict[str, str] = field(default_factory=dict)
|
||||
entities: list[dict[str, str]] = field(default_factory=list) # User-provided entities
|
||||
tags: list[str] = field(default_factory=list) # Visibility scope tags
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -113,6 +116,7 @@ class ExtractedFact:
|
||||
context: str = ""
|
||||
mentioned_at: datetime | None = None
|
||||
metadata: dict[str, str] = field(default_factory=dict)
|
||||
tags: list[str] = field(default_factory=list) # Visibility scope tags
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -158,6 +162,9 @@ class ProcessedFact:
|
||||
# Track which content this fact came from (for user entity merging)
|
||||
content_index: int = 0
|
||||
|
||||
# Visibility scope tags
|
||||
tags: list[str] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def is_duplicate(self) -> bool:
|
||||
"""Check if this fact was marked as a duplicate."""
|
||||
@@ -201,6 +208,7 @@ class ProcessedFact:
|
||||
causal_relations=extracted_fact.causal_relations,
|
||||
chunk_id=chunk_id,
|
||||
content_index=extracted_fact.content_index,
|
||||
tags=extracted_fact.tags,
|
||||
)
|
||||
|
||||
|
||||
@@ -232,6 +240,7 @@ class RetainBatch:
|
||||
document_id: str | None = None
|
||||
fact_type_override: str | None = None
|
||||
confidence_score: float | None = None
|
||||
document_tags: list[str] = field(default_factory=list) # Tags applied to all items
|
||||
|
||||
# Extracted data (populated during processing)
|
||||
extracted_facts: list[ExtractedFact] = field(default_factory=list)
|
||||
|
||||
@@ -11,6 +11,7 @@ from abc import ABC, abstractmethod
|
||||
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .tags import TagsMatch, filter_results_by_tags
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -43,6 +44,8 @@ class GraphRetriever(ABC):
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # TypedAdjacency, optional pre-loaded graph
|
||||
tags: list[str] | None = None, # Visibility scope tags for filtering
|
||||
tags_match: TagsMatch = "any", # How to match tags: 'any' (OR) or 'all' (AND)
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve relevant facts via graph traversal.
|
||||
@@ -57,6 +60,7 @@ class GraphRetriever(ABC):
|
||||
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
|
||||
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
|
||||
adjacency: Pre-loaded typed adjacency graph (optional, for MPFP)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
Tuple of (List of RetrievalResult with activation scores, optional timing info)
|
||||
@@ -114,6 +118,8 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # Not used by BFS
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts using BFS spreading activation.
|
||||
@@ -129,7 +135,9 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
for interface compatibility but not used.
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
results = await self._retrieve_with_conn(conn, query_embedding_str, bank_id, fact_type, budget)
|
||||
results = await self._retrieve_with_conn(
|
||||
conn, query_embedding_str, bank_id, fact_type, budget, tags=tags, tags_match=tags_match
|
||||
)
|
||||
return results, None
|
||||
|
||||
async def _retrieve_with_conn(
|
||||
@@ -139,33 +147,46 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[RetrievalResult]:
|
||||
"""Internal implementation with connection."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, self.entry_point_threshold, self.entry_point_limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
# Step 1: Find entry points
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
self.entry_point_threshold,
|
||||
self.entry_point_limit,
|
||||
*params,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
logger.debug(
|
||||
f"[BFS] No entry points found for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
return []
|
||||
|
||||
logger.debug(
|
||||
f"[BFS] Found {len(entry_points)} entry points for fact_type={fact_type} "
|
||||
f"(tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Step 2: BFS spreading activation
|
||||
visited = set()
|
||||
results = []
|
||||
@@ -196,7 +217,7 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.occurred_end,
|
||||
mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type,
|
||||
mu.document_id, mu.chunk_id,
|
||||
mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight, ml.link_type, ml.from_unit_id
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||
@@ -236,4 +257,8 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
neighbor_result = RetrievalResult.from_db_row(dict(n))
|
||||
queue.append((neighbor_result, new_activation))
|
||||
|
||||
# Apply tags filtering (BFS may traverse into memories that don't match tags criteria)
|
||||
if tags:
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
return results
|
||||
|
||||
@@ -18,6 +18,7 @@ import time
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .graph_retrieval import GraphRetriever
|
||||
from .tags import TagsMatch, filter_results_by_tags
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -30,26 +31,32 @@ async def _find_semantic_seeds(
|
||||
fact_type: str,
|
||||
limit: int = 20,
|
||||
threshold: float = 0.3,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[RetrievalResult]:
|
||||
"""Find semantic seeds via embedding search."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
threshold,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
|
||||
|
||||
@@ -95,6 +102,8 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts by expanding links from seeds.
|
||||
@@ -109,6 +118,7 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
semantic_seeds: Pre-computed semantic entry points
|
||||
temporal_seeds: Pre-computed temporal entry points
|
||||
adjacency: Unused, kept for interface compatibility
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
Tuple of (results, timings)
|
||||
@@ -125,15 +135,27 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
else:
|
||||
seeds_start = time.time()
|
||||
all_seeds = await _find_semantic_seeds(
|
||||
conn, query_embedding_str, bank_id, fact_type, limit=20, threshold=0.3
|
||||
conn,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit=20,
|
||||
threshold=0.3,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
timings.seeds_time = time.time() - seeds_start
|
||||
logger.debug(
|
||||
f"[LinkExpansion] Found {len(all_seeds)} semantic seeds for fact_type={fact_type} "
|
||||
f"(tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Add temporal seeds if provided
|
||||
if temporal_seeds:
|
||||
all_seeds.extend(temporal_seeds)
|
||||
|
||||
if not all_seeds:
|
||||
logger.debug("[LinkExpansion] No seeds found, returning empty results")
|
||||
return [], timings
|
||||
|
||||
seed_ids = list({s.id for s in all_seeds})
|
||||
@@ -147,7 +169,7 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
SELECT
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
COUNT(*)::float AS score
|
||||
FROM {fq_table("unit_entities")} seed_ue
|
||||
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
|
||||
@@ -172,7 +194,7 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
SELECT DISTINCT ON (mu.id)
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight + 1.0 AS score
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||
@@ -219,6 +241,10 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
result.activation = row["score"]
|
||||
results.append(result)
|
||||
|
||||
# Apply tags filtering (graph expansion may reach untagged memories)
|
||||
if tags:
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
timings.result_count = len(results)
|
||||
timings.traverse = time.time() - start_time
|
||||
|
||||
|
||||
@@ -23,6 +23,7 @@ from dataclasses import dataclass, field
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .graph_retrieval import GraphRetriever
|
||||
from .tags import TagsMatch
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -448,7 +449,7 @@ async def fetch_memory_units_by_ids(
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1::uuid[])
|
||||
AND fact_type = $2
|
||||
@@ -503,6 +504,8 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # Ignored - kept for interface compatibility
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts using MPFP algorithm with lazy edge loading.
|
||||
@@ -517,6 +520,7 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
semantic_seeds: Pre-computed semantic entry points
|
||||
temporal_seeds: Pre-computed temporal entry points
|
||||
adjacency: Ignored (kept for interface compatibility)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
Tuple of (List of RetrievalResult with activation scores, MPFPTimings)
|
||||
@@ -532,8 +536,13 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
# If no semantic seeds provided, fall back to finding our own
|
||||
if not semantic_seed_nodes:
|
||||
seeds_start = time.time()
|
||||
semantic_seed_nodes = await self._find_semantic_seeds(pool, query_embedding_str, bank_id, fact_type)
|
||||
semantic_seed_nodes = await self._find_semantic_seeds(
|
||||
pool, query_embedding_str, bank_id, fact_type, tags=tags, tags_match=tags_match
|
||||
)
|
||||
timings.seeds_time = time.time() - seeds_start
|
||||
logger.debug(
|
||||
f"[MPFP] Found {len(semantic_seed_nodes)} semantic seeds for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Collect all pattern jobs
|
||||
pattern_jobs = []
|
||||
@@ -549,6 +558,9 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
pattern_jobs.append((temporal_seed_nodes, pattern))
|
||||
|
||||
if not pattern_jobs:
|
||||
logger.debug(
|
||||
f"[MPFP] No pattern jobs (semantic_seeds={len(semantic_seed_nodes)}, temporal_seeds={len(temporal_seed_nodes)})"
|
||||
)
|
||||
return [], timings
|
||||
|
||||
timings.pattern_count = len(pattern_jobs)
|
||||
@@ -587,6 +599,7 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
timings.fusion = time.time() - step_start
|
||||
|
||||
if not fused:
|
||||
logger.debug(f"[MPFP] No fused results after RRF fusion (pattern_count={len(pattern_results)})")
|
||||
return [], timings
|
||||
|
||||
# Get top result IDs
|
||||
@@ -596,6 +609,13 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
step_start = time.time()
|
||||
results = await fetch_memory_units_by_ids(pool, result_ids, fact_type)
|
||||
timings.fetch = time.time() - step_start
|
||||
|
||||
# Filter results by tags (graph traversal may have picked up unfiltered memories)
|
||||
if tags:
|
||||
from .tags import filter_results_by_tags
|
||||
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
timings.result_count = len(results)
|
||||
|
||||
# Add activation scores from fusion
|
||||
@@ -634,8 +654,17 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
fact_type: str,
|
||||
limit: int = 20,
|
||||
threshold: float = 0.3,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[SeedNode]:
|
||||
"""Fallback: find semantic seeds via embedding search."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
@@ -645,14 +674,11 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
threshold,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
|
||||
return [SeedNode(node_id=str(r["id"]), score=r["similarity"]) for r in rows]
|
||||
|
||||
@@ -20,6 +20,7 @@ from ..memory_engine import fq_table
|
||||
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
|
||||
from .link_expansion_retrieval import LinkExpansionRetriever
|
||||
from .mpfp_retrieval import MPFPGraphRetriever
|
||||
from .tags import TagsMatch, build_tags_where_clause_simple
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -39,6 +40,18 @@ class ParallelRetrievalResult:
|
||||
max_conn_wait: float = 0.0 # Maximum connection acquisition wait time across all methods
|
||||
|
||||
|
||||
@dataclass
|
||||
class MultiFactTypeRetrievalResult:
|
||||
"""Result from retrieval across all fact types."""
|
||||
|
||||
# Results per fact type
|
||||
results_by_fact_type: dict[str, ParallelRetrievalResult]
|
||||
# Aggregate timings
|
||||
timings: dict[str, float] = field(default_factory=dict)
|
||||
# Max connection wait across all operations
|
||||
max_conn_wait: float = 0.0
|
||||
|
||||
|
||||
# Default graph retriever instance (can be overridden)
|
||||
_default_graph_retriever: GraphRetriever | None = None
|
||||
|
||||
@@ -73,7 +86,12 @@ def set_default_graph_retriever(retriever: GraphRetriever) -> None:
|
||||
|
||||
|
||||
async def retrieve_semantic(
|
||||
conn, query_emb_str: str, bank_id: str, fact_type: str, limit: int
|
||||
conn,
|
||||
query_emb_str: str,
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
limit: int,
|
||||
tags: list[str] | None = None,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Semantic retrieval via vector similarity.
|
||||
@@ -84,31 +102,44 @@ async def retrieve_semantic(
|
||||
agent_id: bank ID
|
||||
fact_type: Fact type to filter
|
||||
limit: Maximum results to return
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
List of RetrievalResult objects
|
||||
"""
|
||||
from .tags import TagsMatch, build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 5)
|
||||
params = [query_emb_str, bank_id, fact_type, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= 0.3
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $4
|
||||
""",
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in results]
|
||||
|
||||
|
||||
async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, limit: int) -> list[RetrievalResult]:
|
||||
async def retrieve_bm25(
|
||||
conn,
|
||||
query_text: str,
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
limit: int,
|
||||
tags: list[str] | None = None,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
BM25 keyword retrieval via full-text search.
|
||||
|
||||
@@ -118,12 +149,15 @@ async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, lim
|
||||
agent_id: bank ID
|
||||
fact_type: Fact type to filter
|
||||
limit: Maximum results to return
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
List of RetrievalResult objects
|
||||
"""
|
||||
import re
|
||||
|
||||
from .tags import TagsMatch, build_tags_where_clause_simple
|
||||
|
||||
# Sanitize query text: remove special characters that have meaning in tsquery
|
||||
# Keep only alphanumeric characters and spaces
|
||||
sanitized_text = re.sub(r"[^\w\s]", " ", query_text.lower())
|
||||
@@ -139,25 +173,394 @@ async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, lim
|
||||
# This prevents empty results when some terms are missing
|
||||
query_tsquery = " | ".join(tokens)
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 5)
|
||||
params = [query_tsquery, bank_id, fact_type, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
ts_rank_cd(search_vector, to_tsquery('english', $1)) AS bm25_score
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND fact_type = $3
|
||||
AND search_vector @@ to_tsquery('english', $1)
|
||||
{tags_clause}
|
||||
ORDER BY bm25_score DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
query_tsquery,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in results]
|
||||
|
||||
|
||||
async def retrieve_semantic_bm25_combined(
|
||||
conn,
|
||||
query_emb_str: str,
|
||||
query_text: str,
|
||||
bank_id: str,
|
||||
fact_types: list[str],
|
||||
limit: int,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]]:
|
||||
"""
|
||||
Combined semantic + BM25 retrieval for multiple fact types in a single query.
|
||||
|
||||
Uses CTEs with window functions to get top-N results per fact type per method,
|
||||
all in one database round-trip.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
query_emb_str: Query embedding as string
|
||||
query_text: Query text for BM25
|
||||
bank_id: Bank ID
|
||||
fact_types: List of fact types to retrieve
|
||||
limit: Maximum results per method per fact type
|
||||
|
||||
Returns:
|
||||
Dict mapping fact_type -> (semantic_results, bm25_results)
|
||||
"""
|
||||
import re
|
||||
|
||||
# Sanitize query text for BM25 (same as retrieve_bm25)
|
||||
sanitized_text = re.sub(r"[^\w\s]", " ", query_text.lower())
|
||||
tokens = [token for token in sanitized_text.split() if token]
|
||||
|
||||
# If no valid tokens for BM25, just run semantic
|
||||
if not tokens:
|
||||
tags_clause = build_tags_where_clause_simple(tags, 5, match=tags_match)
|
||||
params = [query_emb_str, bank_id, fact_types, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
WITH semantic_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
NULL::float AS bm25_score,
|
||||
'semantic' AS source,
|
||||
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY embedding <=> $1::vector) AS rn
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = ANY($3)
|
||||
AND (1 - (embedding <=> $1::vector)) >= 0.3
|
||||
{tags_clause}
|
||||
)
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM semantic_ranked
|
||||
WHERE rn <= $4
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
# Group by fact_type
|
||||
result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = {
|
||||
ft: ([], []) for ft in fact_types
|
||||
}
|
||||
for r in results:
|
||||
row = dict(r)
|
||||
ft = row.get("fact_type")
|
||||
row.pop("source", None)
|
||||
if ft in result_dict:
|
||||
result_dict[ft][0].append(RetrievalResult.from_db_row(row))
|
||||
return result_dict
|
||||
|
||||
query_tsquery = " | ".join(tokens)
|
||||
|
||||
# Build tags clause - param 6 if tags provided
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_emb_str, bank_id, fact_types, limit, query_tsquery]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
# Combined CTE query for both semantic and BM25 across all fact types
|
||||
# Uses window functions to limit per fact_type per method
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
WITH semantic_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
NULL::float AS bm25_score,
|
||||
'semantic' AS source,
|
||||
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY embedding <=> $1::vector) AS rn
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = ANY($3)
|
||||
AND (1 - (embedding <=> $1::vector)) >= 0.3
|
||||
{tags_clause}
|
||||
),
|
||||
bm25_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
NULL::float AS similarity,
|
||||
ts_rank_cd(search_vector, to_tsquery('english', $5)) AS bm25_score,
|
||||
'bm25' AS source,
|
||||
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY ts_rank_cd(search_vector, to_tsquery('english', $5)) DESC) AS rn
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND fact_type = ANY($3)
|
||||
AND search_vector @@ to_tsquery('english', $5)
|
||||
{tags_clause}
|
||||
),
|
||||
semantic AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM semantic_ranked WHERE rn <= $4
|
||||
),
|
||||
bm25 AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM bm25_ranked WHERE rn <= $4
|
||||
)
|
||||
SELECT * FROM semantic
|
||||
UNION ALL
|
||||
SELECT * FROM bm25
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
|
||||
# Group results by fact_type and source
|
||||
result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = {ft: ([], []) for ft in fact_types}
|
||||
for r in results:
|
||||
row = dict(r)
|
||||
source = row.pop("source", None)
|
||||
ft = row.get("fact_type")
|
||||
if ft in result_dict:
|
||||
if source == "semantic":
|
||||
result_dict[ft][0].append(RetrievalResult.from_db_row(row))
|
||||
else:
|
||||
result_dict[ft][1].append(RetrievalResult.from_db_row(row))
|
||||
|
||||
return result_dict
|
||||
|
||||
|
||||
async def retrieve_temporal_combined(
|
||||
conn,
|
||||
query_emb_str: str,
|
||||
bank_id: str,
|
||||
fact_types: list[str],
|
||||
start_date: datetime,
|
||||
end_date: datetime,
|
||||
budget: int,
|
||||
semantic_threshold: float = 0.1,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> dict[str, list[RetrievalResult]]:
|
||||
"""
|
||||
Temporal retrieval for multiple fact types in a single query.
|
||||
|
||||
Batches the entry point query using window functions to get top-N per fact type,
|
||||
then runs spreading for each fact type.
|
||||
|
||||
Args:
|
||||
conn: Database connection
|
||||
query_emb_str: Query embedding as string
|
||||
bank_id: Bank ID
|
||||
fact_types: List of fact types to retrieve
|
||||
start_date: Start of time range
|
||||
end_date: End of time range
|
||||
budget: Node budget for spreading per fact type
|
||||
semantic_threshold: Minimum semantic similarity to include
|
||||
|
||||
Returns:
|
||||
Dict mapping fact_type -> list of RetrievalResult
|
||||
"""
|
||||
from ..memory_engine import fq_table
|
||||
|
||||
# Ensure dates are timezone-aware
|
||||
if start_date.tzinfo is None:
|
||||
start_date = start_date.replace(tzinfo=UTC)
|
||||
if end_date.tzinfo is None:
|
||||
end_date = end_date.replace(tzinfo=UTC)
|
||||
|
||||
# Build tags clause
|
||||
tags_clause = build_tags_where_clause_simple(tags, 7, match=tags_match)
|
||||
params = [query_emb_str, bank_id, fact_types, start_date, end_date, semantic_threshold]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
# Batch query: Get entry points for ALL fact types at once with window function
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
WITH ranked_entries AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, embedding <=> $1::vector) AS rn
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND fact_type = ANY($3)
|
||||
AND embedding IS NOT NULL
|
||||
AND (
|
||||
(occurred_start IS NOT NULL AND occurred_end IS NOT NULL
|
||||
AND occurred_start <= $5 AND occurred_end >= $4)
|
||||
OR
|
||||
(mentioned_at IS NOT NULL AND mentioned_at BETWEEN $4 AND $5)
|
||||
OR
|
||||
(occurred_start IS NOT NULL AND occurred_start BETWEEN $4 AND $5)
|
||||
OR
|
||||
(occurred_end IS NOT NULL AND occurred_end BETWEEN $4 AND $5)
|
||||
)
|
||||
AND (1 - (embedding <=> $1::vector)) >= $6
|
||||
{tags_clause}
|
||||
)
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags, similarity
|
||||
FROM ranked_entries
|
||||
WHERE rn <= 10
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
return {ft: [] for ft in fact_types}
|
||||
|
||||
# Group entry points by fact type
|
||||
entries_by_ft: dict[str, list] = {ft: [] for ft in fact_types}
|
||||
for ep in entry_points:
|
||||
ft = ep["fact_type"]
|
||||
if ft in entries_by_ft:
|
||||
entries_by_ft[ft].append(ep)
|
||||
|
||||
# Calculate shared temporal parameters
|
||||
total_days = (end_date - start_date).total_seconds() / 86400
|
||||
mid_date = start_date + (end_date - start_date) / 2
|
||||
|
||||
# Process each fact type (spreading needs to stay per fact type due to link filtering)
|
||||
results_by_ft: dict[str, list[RetrievalResult]] = {}
|
||||
|
||||
for ft in fact_types:
|
||||
ft_entry_points = entries_by_ft.get(ft, [])
|
||||
if not ft_entry_points:
|
||||
results_by_ft[ft] = []
|
||||
continue
|
||||
|
||||
results = []
|
||||
visited = set()
|
||||
node_scores = {}
|
||||
|
||||
# Process entry points
|
||||
for ep in ft_entry_points:
|
||||
unit_id = str(ep["id"])
|
||||
visited.add(unit_id)
|
||||
|
||||
# Calculate temporal proximity
|
||||
best_date = None
|
||||
if ep["occurred_start"] is not None and ep["occurred_end"] is not None:
|
||||
best_date = ep["occurred_start"] + (ep["occurred_end"] - ep["occurred_start"]) / 2
|
||||
elif ep["occurred_start"] is not None:
|
||||
best_date = ep["occurred_start"]
|
||||
elif ep["occurred_end"] is not None:
|
||||
best_date = ep["occurred_end"]
|
||||
elif ep["mentioned_at"] is not None:
|
||||
best_date = ep["mentioned_at"]
|
||||
|
||||
if best_date:
|
||||
days_from_mid = abs((best_date - mid_date).total_seconds() / 86400)
|
||||
temporal_proximity = 1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
|
||||
else:
|
||||
temporal_proximity = 0.5
|
||||
|
||||
ep_result = RetrievalResult.from_db_row(dict(ep))
|
||||
ep_result.temporal_score = temporal_proximity
|
||||
ep_result.temporal_proximity = temporal_proximity
|
||||
results.append(ep_result)
|
||||
node_scores[unit_id] = (ep["similarity"], 1.0)
|
||||
|
||||
# Spreading through temporal links (same as single-fact-type version)
|
||||
frontier = list(node_scores.keys())
|
||||
budget_remaining = budget - len(ft_entry_points)
|
||||
batch_size = 20
|
||||
|
||||
# Build tags clause for spreading (use param 6 since 1-5 are used)
|
||||
spreading_tags_clause = build_tags_where_clause_simple(tags, 6, table_alias="mu.", match=tags_match)
|
||||
|
||||
while frontier and budget_remaining > 0:
|
||||
batch_ids = frontier[:batch_size]
|
||||
frontier = frontier[batch_size:]
|
||||
|
||||
spreading_params = [query_emb_str, batch_ids, ft, semantic_threshold, batch_size * 10]
|
||||
if tags:
|
||||
spreading_params.append(tags)
|
||||
|
||||
neighbors = await conn.fetch(
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight, ml.link_type, ml.from_unit_id,
|
||||
1 - (mu.embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||
WHERE ml.from_unit_id = ANY($2::uuid[])
|
||||
AND ml.link_type IN ('temporal', 'causes', 'caused_by', 'enables', 'prevents')
|
||||
AND ml.weight >= 0.1
|
||||
AND mu.fact_type = $3
|
||||
AND mu.embedding IS NOT NULL
|
||||
AND (1 - (mu.embedding <=> $1::vector)) >= $4
|
||||
{spreading_tags_clause}
|
||||
ORDER BY ml.weight DESC
|
||||
LIMIT $5
|
||||
""",
|
||||
*spreading_params,
|
||||
)
|
||||
|
||||
for n in neighbors:
|
||||
neighbor_id = str(n["id"])
|
||||
if neighbor_id in visited:
|
||||
continue
|
||||
|
||||
visited.add(neighbor_id)
|
||||
budget_remaining -= 1
|
||||
|
||||
parent_id = str(n["from_unit_id"])
|
||||
_, parent_temporal_score = node_scores.get(parent_id, (0.5, 0.5))
|
||||
|
||||
neighbor_best_date = None
|
||||
if n["occurred_start"] is not None and n["occurred_end"] is not None:
|
||||
neighbor_best_date = n["occurred_start"] + (n["occurred_end"] - n["occurred_start"]) / 2
|
||||
elif n["occurred_start"] is not None:
|
||||
neighbor_best_date = n["occurred_start"]
|
||||
elif n["occurred_end"] is not None:
|
||||
neighbor_best_date = n["occurred_end"]
|
||||
elif n["mentioned_at"] is not None:
|
||||
neighbor_best_date = n["mentioned_at"]
|
||||
|
||||
if neighbor_best_date:
|
||||
days_from_mid = abs((neighbor_best_date - mid_date).total_seconds() / 86400)
|
||||
neighbor_temporal_proximity = (
|
||||
1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
|
||||
)
|
||||
else:
|
||||
neighbor_temporal_proximity = 0.3
|
||||
|
||||
link_type = n["link_type"]
|
||||
if link_type in ("causes", "caused_by"):
|
||||
causal_boost = 2.0
|
||||
elif link_type in ("enables", "prevents"):
|
||||
causal_boost = 1.5
|
||||
else:
|
||||
causal_boost = 1.0
|
||||
|
||||
propagated_temporal = parent_temporal_score * n["weight"] * causal_boost * 0.7
|
||||
combined_temporal = max(neighbor_temporal_proximity, propagated_temporal)
|
||||
|
||||
neighbor_result = RetrievalResult.from_db_row(dict(n))
|
||||
neighbor_result.temporal_score = combined_temporal
|
||||
neighbor_result.temporal_proximity = neighbor_temporal_proximity
|
||||
results.append(neighbor_result)
|
||||
|
||||
if budget_remaining > 0 and combined_temporal > 0.2:
|
||||
node_scores[neighbor_id] = (n["similarity"], combined_temporal)
|
||||
frontier.append(neighbor_id)
|
||||
|
||||
if budget_remaining <= 0:
|
||||
break
|
||||
|
||||
results_by_ft[ft] = results
|
||||
|
||||
return results_by_ft
|
||||
|
||||
|
||||
async def retrieve_temporal(
|
||||
conn,
|
||||
query_emb_str: str,
|
||||
@@ -167,6 +570,7 @@ async def retrieve_temporal(
|
||||
end_date: datetime,
|
||||
budget: int,
|
||||
semantic_threshold: float = 0.1,
|
||||
tags: list[str] | None = None,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Temporal retrieval with spreading activation.
|
||||
@@ -185,6 +589,7 @@ async def retrieve_temporal(
|
||||
end_date: End of time range
|
||||
budget: Node budget for spreading
|
||||
semantic_threshold: Minimum semantic similarity to include
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
List of RetrievalResult objects with temporal scores
|
||||
@@ -196,9 +601,16 @@ async def retrieve_temporal(
|
||||
if end_date.tzinfo is None:
|
||||
end_date = end_date.replace(tzinfo=UTC)
|
||||
|
||||
from .tags import TagsMatch, build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 7)
|
||||
params = [query_emb_str, bank_id, fact_type, start_date, end_date, semantic_threshold]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
@@ -218,15 +630,11 @@ async def retrieve_temporal(
|
||||
(occurred_end IS NOT NULL AND occurred_end BETWEEN $4 AND $5)
|
||||
)
|
||||
AND (1 - (embedding <=> $1::vector)) >= $6
|
||||
{tags_clause}
|
||||
ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, (embedding <=> $1::vector) ASC
|
||||
LIMIT 10
|
||||
""",
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
start_date,
|
||||
end_date,
|
||||
semantic_threshold,
|
||||
*params,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
@@ -378,6 +786,7 @@ async def retrieve_parallel(
|
||||
query_analyzer: Optional["QueryAnalyzer"] = None,
|
||||
graph_retriever: GraphRetriever | None = None,
|
||||
temporal_constraint: tuple | None = None, # Pre-extracted temporal constraint
|
||||
tags: list[str] | None = None, # Visibility scope tags for filtering
|
||||
) -> ParallelRetrievalResult:
|
||||
"""
|
||||
Run 3-way or 4-way parallel retrieval (adds temporal if detected).
|
||||
@@ -393,6 +802,7 @@ async def retrieve_parallel(
|
||||
query_analyzer: Query analyzer to use (defaults to TransformerQueryAnalyzer)
|
||||
graph_retriever: Graph retrieval strategy (defaults to configured retriever)
|
||||
temporal_constraint: Pre-extracted temporal constraint (optional)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
ParallelRetrievalResult with semantic, bm25, graph, temporal results and timings
|
||||
@@ -413,6 +823,7 @@ async def retrieve_parallel(
|
||||
retriever,
|
||||
question_date,
|
||||
query_analyzer,
|
||||
tags=tags,
|
||||
)
|
||||
else:
|
||||
# For BFS, extract temporal constraint upfront (legacy path)
|
||||
@@ -423,7 +834,15 @@ async def retrieve_parallel(
|
||||
query_text, reference_date=question_date, analyzer=query_analyzer
|
||||
)
|
||||
return await _retrieve_parallel_bfs(
|
||||
pool, query_text, query_embedding_str, bank_id, fact_type, thinking_budget, temporal_constraint, retriever
|
||||
pool,
|
||||
query_text,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
thinking_budget,
|
||||
temporal_constraint,
|
||||
retriever,
|
||||
tags=tags,
|
||||
)
|
||||
|
||||
|
||||
@@ -447,6 +866,7 @@ async def _retrieve_parallel_mpfp(
|
||||
retriever: GraphRetriever,
|
||||
question_date: datetime | None = None,
|
||||
query_analyzer=None,
|
||||
tags: list[str] | None = None,
|
||||
) -> ParallelRetrievalResult:
|
||||
"""
|
||||
MPFP retrieval with true parallelization.
|
||||
@@ -468,7 +888,9 @@ async def _retrieve_parallel_mpfp(
|
||||
acquire_start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
conn_wait = time.time() - acquire_start
|
||||
results = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
|
||||
results = await retrieve_semantic(
|
||||
conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget, tags=tags
|
||||
)
|
||||
return _TimedResult(results, time.time() - start, conn_wait)
|
||||
|
||||
async def run_bm25() -> _TimedResult:
|
||||
@@ -477,7 +899,7 @@ async def _retrieve_parallel_mpfp(
|
||||
acquire_start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
conn_wait = time.time() - acquire_start
|
||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget)
|
||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget, tags=tags)
|
||||
return _TimedResult(results, time.time() - start, conn_wait)
|
||||
|
||||
async def run_graph() -> tuple[list[RetrievalResult], float, MPFPTimings | None]:
|
||||
@@ -495,6 +917,7 @@ async def _retrieve_parallel_mpfp(
|
||||
query_text=query_text,
|
||||
semantic_seeds=None, # Let MPFP find its own seeds
|
||||
temporal_seeds=None, # Don't wait for temporal extraction
|
||||
tags=tags,
|
||||
)
|
||||
return results, time.time() - start, mpfp_timing
|
||||
|
||||
@@ -666,6 +1089,7 @@ async def _retrieve_parallel_bfs(
|
||||
thinking_budget: int,
|
||||
temporal_constraint: tuple | None,
|
||||
retriever: GraphRetriever,
|
||||
tags: list[str] | None = None,
|
||||
) -> ParallelRetrievalResult:
|
||||
"""BFS retrieval: all methods run in parallel (original behavior)."""
|
||||
import time
|
||||
@@ -673,13 +1097,15 @@ async def _retrieve_parallel_bfs(
|
||||
async def run_semantic() -> _TimedResult:
|
||||
start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
results = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
|
||||
results = await retrieve_semantic(
|
||||
conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget, tags=tags
|
||||
)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
async def run_bm25() -> _TimedResult:
|
||||
start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget)
|
||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget, tags=tags)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
async def run_graph() -> _TimedResult:
|
||||
@@ -691,6 +1117,7 @@ async def _retrieve_parallel_bfs(
|
||||
fact_type=fact_type,
|
||||
budget=thinking_budget,
|
||||
query_text=query_text,
|
||||
tags=tags,
|
||||
)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
@@ -706,6 +1133,7 @@ async def _retrieve_parallel_bfs(
|
||||
tc_end,
|
||||
budget=thinking_budget,
|
||||
semantic_threshold=0.1,
|
||||
tags=tags,
|
||||
)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
@@ -748,3 +1176,171 @@ async def _retrieve_parallel_bfs(
|
||||
},
|
||||
temporal_constraint=None,
|
||||
)
|
||||
|
||||
|
||||
async def retrieve_all_fact_types_parallel(
|
||||
pool,
|
||||
query_text: str,
|
||||
query_embedding_str: str,
|
||||
bank_id: str,
|
||||
fact_types: list[str],
|
||||
thinking_budget: int,
|
||||
question_date: datetime | None = None,
|
||||
query_analyzer: Optional["QueryAnalyzer"] = None,
|
||||
graph_retriever: GraphRetriever | None = None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> MultiFactTypeRetrievalResult:
|
||||
"""
|
||||
Optimized retrieval for multiple fact types using batched queries.
|
||||
|
||||
This reduces database round-trips by:
|
||||
1. Combining semantic + BM25 into one CTE query for ALL fact types (1 query instead of 2N)
|
||||
2. Running graph retrieval per fact type in parallel (N parallel tasks)
|
||||
3. Running temporal retrieval per fact type in parallel (N parallel tasks)
|
||||
|
||||
Args:
|
||||
pool: Database connection pool
|
||||
query_text: Query text
|
||||
query_embedding_str: Query embedding as string
|
||||
bank_id: Bank ID
|
||||
fact_types: List of fact types to retrieve
|
||||
thinking_budget: Budget for graph traversal and retrieval limits
|
||||
question_date: Optional date when question was asked (for temporal filtering)
|
||||
query_analyzer: Query analyzer to use (defaults to TransformerQueryAnalyzer)
|
||||
graph_retriever: Graph retrieval strategy (defaults to configured retriever)
|
||||
|
||||
Returns:
|
||||
MultiFactTypeRetrievalResult with results organized by fact type
|
||||
"""
|
||||
import time
|
||||
|
||||
retriever = graph_retriever or get_default_graph_retriever()
|
||||
start_time = time.time()
|
||||
timings: dict[str, float] = {}
|
||||
|
||||
# Step 1: Extract temporal constraint first (CPU work, no DB)
|
||||
# Do this before DB queries so we know if we need temporal retrieval
|
||||
temporal_extraction_start = time.time()
|
||||
from .temporal_extraction import extract_temporal_constraint
|
||||
|
||||
temporal_constraint = extract_temporal_constraint(query_text, reference_date=question_date, analyzer=query_analyzer)
|
||||
temporal_extraction_time = time.time() - temporal_extraction_start
|
||||
timings["temporal_extraction"] = temporal_extraction_time
|
||||
|
||||
# Step 2: Run semantic + BM25 + temporal combined in ONE connection!
|
||||
# This reduces connection usage from 2 to 1 for these operations
|
||||
semantic_bm25_start = time.time()
|
||||
temporal_results_by_ft: dict[str, list[RetrievalResult]] = {}
|
||||
temporal_time = 0.0
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
conn_wait = time.time() - semantic_bm25_start
|
||||
|
||||
# Semantic + BM25 combined
|
||||
semantic_bm25_results = await retrieve_semantic_bm25_combined(
|
||||
conn,
|
||||
query_embedding_str,
|
||||
query_text,
|
||||
bank_id,
|
||||
fact_types,
|
||||
thinking_budget,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
semantic_bm25_time = time.time() - semantic_bm25_start
|
||||
|
||||
# Temporal combined (if constraint detected) - same connection!
|
||||
if temporal_constraint:
|
||||
tc_start, tc_end = temporal_constraint
|
||||
temporal_start = time.time()
|
||||
temporal_results_by_ft = await retrieve_temporal_combined(
|
||||
conn,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_types,
|
||||
tc_start,
|
||||
tc_end,
|
||||
budget=thinking_budget,
|
||||
semantic_threshold=0.1,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
temporal_time = time.time() - temporal_start
|
||||
|
||||
timings["semantic_bm25_combined"] = semantic_bm25_time
|
||||
timings["temporal_combined"] = temporal_time
|
||||
|
||||
# Step 3: Run graph retrieval for each fact type in parallel
|
||||
async def run_graph_for_fact_type(ft: str) -> tuple[str, list[RetrievalResult], float, MPFPTimings | None]:
|
||||
graph_start = time.time()
|
||||
results, mpfp_timing = await retriever.retrieve(
|
||||
pool=pool,
|
||||
query_embedding_str=query_embedding_str,
|
||||
bank_id=bank_id,
|
||||
fact_type=ft,
|
||||
budget=thinking_budget,
|
||||
query_text=query_text,
|
||||
semantic_seeds=None,
|
||||
temporal_seeds=None,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
return ft, results, time.time() - graph_start, mpfp_timing
|
||||
|
||||
# Run graph for all fact types in parallel
|
||||
graph_tasks = [run_graph_for_fact_type(ft) for ft in fact_types]
|
||||
graph_results_list = await asyncio.gather(*graph_tasks)
|
||||
|
||||
# Organize results by fact type
|
||||
results_by_fact_type: dict[str, ParallelRetrievalResult] = {}
|
||||
max_conn_wait = conn_wait # Single connection for semantic+bm25+temporal
|
||||
all_mpfp_timings: list[MPFPTimings] = []
|
||||
|
||||
for ft in fact_types:
|
||||
# Get semantic + bm25 results for this fact type
|
||||
semantic_results, bm25_results = semantic_bm25_results.get(ft, ([], []))
|
||||
|
||||
# Find graph results for this fact type
|
||||
graph_results = []
|
||||
graph_time = 0.0
|
||||
mpfp_timing = None
|
||||
for gr in graph_results_list:
|
||||
if gr[0] == ft:
|
||||
graph_results = gr[1]
|
||||
graph_time = gr[2]
|
||||
mpfp_timing = gr[3]
|
||||
if mpfp_timing:
|
||||
all_mpfp_timings.append(mpfp_timing)
|
||||
break
|
||||
|
||||
# Get temporal results for this fact type from combined result
|
||||
temporal_results = temporal_results_by_ft.get(ft) if temporal_constraint else None
|
||||
if temporal_results is not None and len(temporal_results) == 0:
|
||||
temporal_results = None
|
||||
|
||||
results_by_fact_type[ft] = ParallelRetrievalResult(
|
||||
semantic=semantic_results,
|
||||
bm25=bm25_results,
|
||||
graph=graph_results,
|
||||
temporal=temporal_results,
|
||||
timings={
|
||||
"semantic": semantic_bm25_time / 2, # Approximate split
|
||||
"bm25": semantic_bm25_time / 2,
|
||||
"graph": graph_time,
|
||||
"temporal": temporal_time, # Same for all fact types (single query)
|
||||
"temporal_extraction": temporal_extraction_time,
|
||||
},
|
||||
temporal_constraint=temporal_constraint,
|
||||
mpfp_timings=[mpfp_timing] if mpfp_timing else [],
|
||||
max_conn_wait=max_conn_wait,
|
||||
)
|
||||
|
||||
total_time = time.time() - start_time
|
||||
timings["total"] = total_time
|
||||
|
||||
return MultiFactTypeRetrievalResult(
|
||||
results_by_fact_type=results_by_fact_type,
|
||||
timings=timings,
|
||||
max_conn_wait=max_conn_wait,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
"""
|
||||
Tags filtering utilities for retrieval.
|
||||
|
||||
Provides SQL building functions for filtering memories by tags.
|
||||
Supports four matching modes via TagsMatch enum:
|
||||
- "any": OR matching, includes untagged memories (default, backward compatible)
|
||||
- "all": AND matching, includes untagged memories
|
||||
- "any_strict": OR matching, excludes untagged memories
|
||||
- "all_strict": AND matching, excludes untagged memories
|
||||
|
||||
OR matching (any/any_strict): Memory matches if ANY of its tags overlap with request tags
|
||||
AND matching (all/all_strict): Memory matches if ALL request tags are present in its tags
|
||||
"""
|
||||
|
||||
from typing import Literal
|
||||
|
||||
TagsMatch = Literal["any", "all", "any_strict", "all_strict"]
|
||||
|
||||
|
||||
def _parse_tags_match(match: TagsMatch) -> tuple[str, bool]:
|
||||
"""
|
||||
Parse TagsMatch into operator and include_untagged flag.
|
||||
|
||||
Returns:
|
||||
Tuple of (operator, include_untagged)
|
||||
- operator: "&&" for any/any_strict, "@>" for all/all_strict
|
||||
- include_untagged: True for any/all, False for any_strict/all_strict
|
||||
"""
|
||||
if match == "any":
|
||||
return "&&", True
|
||||
elif match == "all":
|
||||
return "@>", True
|
||||
elif match == "any_strict":
|
||||
return "&&", False
|
||||
elif match == "all_strict":
|
||||
return "@>", False
|
||||
else:
|
||||
# Default to "any" behavior
|
||||
return "&&", True
|
||||
|
||||
|
||||
def build_tags_where_clause(
|
||||
tags: list[str] | None,
|
||||
param_offset: int = 1,
|
||||
table_alias: str = "",
|
||||
match: TagsMatch = "any",
|
||||
) -> tuple[str, list, int]:
|
||||
"""
|
||||
Build a SQL WHERE clause for filtering by tags.
|
||||
|
||||
Supports four matching modes:
|
||||
- "any" (default): OR matching, includes untagged memories
|
||||
- "all": AND matching, includes untagged memories
|
||||
- "any_strict": OR matching, excludes untagged memories
|
||||
- "all_strict": AND matching, excludes untagged memories
|
||||
|
||||
Args:
|
||||
tags: List of tags to filter by. If None or empty, returns empty clause (no filtering).
|
||||
param_offset: Starting parameter number for SQL placeholders (default 1).
|
||||
table_alias: Optional table alias prefix (e.g., "mu." for "memory_units mu").
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
Tuple of (sql_clause, params, next_param_offset):
|
||||
- sql_clause: SQL WHERE clause string
|
||||
- params: List of parameter values to bind
|
||||
- next_param_offset: Next available parameter number
|
||||
|
||||
Example:
|
||||
>>> clause, params, next_offset = build_tags_where_clause(['user_a'], 3, 'mu.', 'any_strict')
|
||||
>>> print(clause) # "AND mu.tags IS NOT NULL AND mu.tags != '{}' AND mu.tags && $3"
|
||||
"""
|
||||
if not tags:
|
||||
return "", [], param_offset
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(match)
|
||||
|
||||
if include_untagged:
|
||||
# Include untagged memories (NULL or empty array) OR matching tags
|
||||
clause = f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_offset})"
|
||||
else:
|
||||
# Strict: only memories with matching tags (exclude NULL and empty)
|
||||
clause = f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_offset}"
|
||||
|
||||
return clause, [tags], param_offset + 1
|
||||
|
||||
|
||||
def build_tags_where_clause_simple(
|
||||
tags: list[str] | None,
|
||||
param_num: int,
|
||||
table_alias: str = "",
|
||||
match: TagsMatch = "any",
|
||||
) -> str:
|
||||
"""
|
||||
Build a simple SQL WHERE clause for tags filtering.
|
||||
|
||||
This is a convenience version that returns just the clause string,
|
||||
assuming the caller will add the tags array to their params list.
|
||||
|
||||
Args:
|
||||
tags: List of tags to filter by. If None or empty, returns empty string.
|
||||
param_num: Parameter number to use in the clause.
|
||||
table_alias: Optional table alias prefix.
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
SQL clause string or empty string.
|
||||
"""
|
||||
if not tags:
|
||||
return ""
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(match)
|
||||
|
||||
if include_untagged:
|
||||
# Include untagged memories (NULL or empty array) OR matching tags
|
||||
return f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_num})"
|
||||
else:
|
||||
# Strict: only memories with matching tags (exclude NULL and empty)
|
||||
return f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_num}"
|
||||
|
||||
|
||||
def filter_results_by_tags(
|
||||
results: list,
|
||||
tags: list[str] | None,
|
||||
match: TagsMatch = "any",
|
||||
) -> list:
|
||||
"""
|
||||
Filter retrieval results by tags in Python (for post-processing).
|
||||
|
||||
Used when SQL filtering isn't possible (e.g., graph traversal results).
|
||||
|
||||
Args:
|
||||
results: List of RetrievalResult objects with a 'tags' attribute.
|
||||
tags: List of tags to filter by. If None or empty, returns all results.
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
Filtered list of results.
|
||||
"""
|
||||
if not tags:
|
||||
return results
|
||||
|
||||
_, include_untagged = _parse_tags_match(match)
|
||||
is_any_match = match in ("any", "any_strict")
|
||||
|
||||
tags_set = set(tags)
|
||||
filtered = []
|
||||
|
||||
for result in results:
|
||||
result_tags = getattr(result, "tags", None)
|
||||
|
||||
# Check if untagged
|
||||
is_untagged = result_tags is None or len(result_tags) == 0
|
||||
|
||||
if is_untagged:
|
||||
if include_untagged:
|
||||
filtered.append(result)
|
||||
# else: skip untagged
|
||||
else:
|
||||
result_tags_set = set(result_tags)
|
||||
if is_any_match:
|
||||
# Any overlap
|
||||
if result_tags_set & tags_set:
|
||||
filtered.append(result)
|
||||
else:
|
||||
# All tags must be present
|
||||
if tags_set <= result_tags_set:
|
||||
filtered.append(result)
|
||||
|
||||
return filtered
|
||||
@@ -11,6 +11,13 @@ from typing import Any, Literal
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class TemporalConstraint(BaseModel):
|
||||
"""Detected temporal constraint from query analysis."""
|
||||
|
||||
start: datetime | None = Field(default=None, description="Start of temporal range")
|
||||
end: datetime | None = Field(default=None, description="End of temporal range")
|
||||
|
||||
|
||||
class QueryInfo(BaseModel):
|
||||
"""Information about the search query."""
|
||||
|
||||
@@ -19,6 +26,11 @@ class QueryInfo(BaseModel):
|
||||
timestamp: datetime = Field(description="When the query was executed")
|
||||
budget: int = Field(description="Maximum nodes to explore")
|
||||
max_tokens: int = Field(description="Maximum tokens to return in results")
|
||||
tags: list[str] | None = Field(default=None, description="Tags filter applied to recall")
|
||||
tags_match: str | None = Field(default=None, description="Tags matching mode: any, all, any_strict, all_strict")
|
||||
temporal_constraint: TemporalConstraint | None = Field(
|
||||
default=None, description="Detected temporal range from query"
|
||||
)
|
||||
|
||||
|
||||
class EntryPoint(BaseModel):
|
||||
|
||||
@@ -22,6 +22,7 @@ from .trace import (
|
||||
SearchPhaseMetrics,
|
||||
SearchSummary,
|
||||
SearchTrace,
|
||||
TemporalConstraint,
|
||||
WeightComponents,
|
||||
)
|
||||
|
||||
@@ -45,7 +46,14 @@ class SearchTracer:
|
||||
json_output = trace.to_json()
|
||||
"""
|
||||
|
||||
def __init__(self, query: str, budget: int, max_tokens: int):
|
||||
def __init__(
|
||||
self,
|
||||
query: str,
|
||||
budget: int,
|
||||
max_tokens: int,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str | None = None,
|
||||
):
|
||||
"""
|
||||
Initialize tracer.
|
||||
|
||||
@@ -53,10 +61,14 @@ class SearchTracer:
|
||||
query: Search query text
|
||||
budget: Maximum nodes to explore
|
||||
max_tokens: Maximum tokens to return in results
|
||||
tags: Tags filter applied to recall
|
||||
tags_match: Tags matching mode (any, all, any_strict, all_strict)
|
||||
"""
|
||||
self.query_text = query
|
||||
self.budget = budget
|
||||
self.max_tokens = max_tokens
|
||||
self.tags = tags
|
||||
self.tags_match = tags_match
|
||||
|
||||
# Trace data
|
||||
self.query_embedding: list[float] | None = None
|
||||
@@ -66,6 +78,9 @@ class SearchTracer:
|
||||
self.pruned: list[PruningDecision] = []
|
||||
self.phase_metrics: list[SearchPhaseMetrics] = []
|
||||
|
||||
# Temporal constraint detected from query
|
||||
self.temporal_constraint: TemporalConstraint | None = None
|
||||
|
||||
# New 4-way retrieval tracking
|
||||
self.retrieval_results: list[RetrievalMethodResults] = []
|
||||
self.rrf_merged: list[RRFMergeResult] = []
|
||||
@@ -88,6 +103,11 @@ class SearchTracer:
|
||||
"""Record the query embedding."""
|
||||
self.query_embedding = embedding
|
||||
|
||||
def record_temporal_constraint(self, start: datetime | None, end: datetime | None):
|
||||
"""Record the detected temporal constraint from query analysis."""
|
||||
if start is not None or end is not None:
|
||||
self.temporal_constraint = TemporalConstraint(start=start, end=end)
|
||||
|
||||
def add_entry_point(self, node_id: str, text: str, similarity: float, rank: int):
|
||||
"""
|
||||
Record an entry point.
|
||||
@@ -428,6 +448,9 @@ class SearchTracer:
|
||||
timestamp=datetime.now(UTC),
|
||||
budget=self.budget,
|
||||
max_tokens=self.max_tokens,
|
||||
tags=self.tags,
|
||||
tags_match=self.tags_match,
|
||||
temporal_constraint=self.temporal_constraint,
|
||||
)
|
||||
|
||||
# Create summary
|
||||
|
||||
@@ -48,6 +48,7 @@ class RetrievalResult:
|
||||
chunk_id: str | None = None
|
||||
access_count: int = 0
|
||||
embedding: list[float] | None = None
|
||||
tags: list[str] | None = None # Visibility scope tags
|
||||
|
||||
# Retrieval-specific scores (only one will be set depending on retrieval method)
|
||||
similarity: float | None = None # Semantic retrieval
|
||||
@@ -72,6 +73,7 @@ class RetrievalResult:
|
||||
chunk_id=row.get("chunk_id"),
|
||||
access_count=row.get("access_count", 0),
|
||||
embedding=row.get("embedding"),
|
||||
tags=row.get("tags"),
|
||||
similarity=row.get("similarity"),
|
||||
bm25_score=row.get("bm25_score"),
|
||||
activation=row.get("activation"),
|
||||
@@ -156,6 +158,7 @@ class ScoredResult:
|
||||
"chunk_id": self.retrieval.chunk_id,
|
||||
"access_count": self.retrieval.access_count,
|
||||
"embedding": self.retrieval.embedding,
|
||||
"tags": self.retrieval.tags,
|
||||
"semantic_similarity": self.retrieval.similarity,
|
||||
"bm25_score": self.retrieval.bm25_score,
|
||||
}
|
||||
|
||||
@@ -187,19 +187,24 @@ 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,
|
||||
|
||||
@@ -28,6 +28,15 @@ from opentelemetry.sdk.resources import Resource
|
||||
if TYPE_CHECKING:
|
||||
import asyncpg
|
||||
|
||||
|
||||
def _get_tenant() -> str:
|
||||
"""Get current tenant (schema) from context for metrics labeling."""
|
||||
# Import here to avoid circular imports
|
||||
from hindsight_api.engine.memory_engine import get_current_schema
|
||||
|
||||
return get_current_schema()
|
||||
|
||||
|
||||
# Custom bucket boundaries for operation duration (in seconds)
|
||||
# Fine granularity in 0-30s range where most operations complete
|
||||
DURATION_BUCKETS = (0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 20.0, 30.0, 60.0, 120.0)
|
||||
@@ -323,6 +332,7 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
"operation": operation,
|
||||
"bank_id": bank_id,
|
||||
"source": source,
|
||||
"tenant": _get_tenant(),
|
||||
}
|
||||
if budget:
|
||||
attributes["budget"] = budget
|
||||
@@ -373,6 +383,7 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
"model": model,
|
||||
"scope": scope,
|
||||
"success": str(success).lower(),
|
||||
"tenant": _get_tenant(),
|
||||
}
|
||||
|
||||
# Record duration
|
||||
@@ -425,10 +436,14 @@ class MetricsCollector(MetricsCollectorBase):
|
||||
status_code = status_code_getter()
|
||||
status_class = f"{status_code // 100}xx"
|
||||
|
||||
# Get tenant from context (may be set during request processing)
|
||||
tenant = _get_tenant()
|
||||
|
||||
attributes = {
|
||||
**base_attributes,
|
||||
"status_code": str(status_code),
|
||||
"status_class": status_class,
|
||||
"tenant": tenant,
|
||||
}
|
||||
|
||||
# Record duration and count
|
||||
|
||||
@@ -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}'")
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ This module provides the ASGI app for uvicorn import string usage:
|
||||
For CLI usage, use the hindsight-api command instead.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import warnings
|
||||
|
||||
@@ -17,6 +18,12 @@ warnings.filterwarnings("ignore", message="websockets.server.WebSocketServerProt
|
||||
from hindsight_api import MemoryEngine
|
||||
from hindsight_api.api import create_app
|
||||
from hindsight_api.config import get_config
|
||||
from hindsight_api.extensions import (
|
||||
DefaultExtensionContext,
|
||||
OperationValidatorExtension,
|
||||
TenantExtension,
|
||||
load_extension,
|
||||
)
|
||||
|
||||
# Disable tokenizers parallelism to avoid warnings
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
@@ -25,10 +32,33 @@ os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
config = get_config()
|
||||
config.configure_logging()
|
||||
|
||||
# Load operation validator extension if configured
|
||||
operation_validator = load_extension("OPERATION_VALIDATOR", OperationValidatorExtension)
|
||||
if operation_validator:
|
||||
logging.info(f"Loaded operation validator: {operation_validator.__class__.__name__}")
|
||||
|
||||
# Load tenant extension if configured
|
||||
tenant_extension = load_extension("TENANT", TenantExtension)
|
||||
if tenant_extension:
|
||||
logging.info(f"Loaded tenant extension: {tenant_extension.__class__.__name__}")
|
||||
|
||||
# Create app at module level (required for uvicorn import string)
|
||||
# MemoryEngine reads configuration from environment variables automatically
|
||||
# Note: run_migrations=True by default, but migrations are idempotent so safe with workers
|
||||
_memory = MemoryEngine(run_migrations=config.run_migrations_on_startup)
|
||||
_memory = MemoryEngine(
|
||||
operation_validator=operation_validator,
|
||||
tenant_extension=tenant_extension,
|
||||
run_migrations=config.run_migrations_on_startup,
|
||||
)
|
||||
|
||||
# Set extension context on tenant extension (needed for schema provisioning)
|
||||
if tenant_extension:
|
||||
extension_context = DefaultExtensionContext(
|
||||
database_url=config.database_url,
|
||||
memory_engine=_memory,
|
||||
)
|
||||
tenant_extension.set_context(extension_context)
|
||||
logging.info("Extension context set on tenant extension")
|
||||
|
||||
# Create unified app with both HTTP and optionally MCP
|
||||
app = create_app(
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -0,0 +1,883 @@
|
||||
"""
|
||||
Tests for tags-based visibility scoping.
|
||||
|
||||
This module tests the tags feature which allows filtering memories by visibility tags.
|
||||
Use cases:
|
||||
- Multi-user agent: Agent has a single memory bank, users should only see memories from
|
||||
conversations they participated in
|
||||
- Student tracking: Teacher tracks students, students should only see their own data
|
||||
|
||||
The tags use OR-based matching: a memory matches if ANY of its tags overlap with the request tags.
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from hindsight_api.api import create_app
|
||||
from hindsight_api.engine.search.tags import build_tags_where_clause_simple, filter_results_by_tags
|
||||
|
||||
# ============================================================================
|
||||
# Unit Tests for tags SQL builder
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestTagsWhereClauseBuilder:
|
||||
"""Unit tests for the tags WHERE clause SQL builder."""
|
||||
|
||||
def test_no_tags_returns_empty_string(self):
|
||||
"""When tags is None, should return empty string (no filtering)."""
|
||||
result = build_tags_where_clause_simple(None, 5)
|
||||
assert result == ""
|
||||
|
||||
def test_empty_tags_list_returns_empty_string(self):
|
||||
"""When tags is an empty list, should return empty string (no filtering)."""
|
||||
result = build_tags_where_clause_simple([], 5)
|
||||
assert result == ""
|
||||
|
||||
def test_tags_with_different_param_num(self):
|
||||
"""Should use the provided parameter number."""
|
||||
result = build_tags_where_clause_simple(["user_a", "user_b"], 3)
|
||||
# Default is "any" which includes untagged
|
||||
assert "$3" in result
|
||||
|
||||
def test_tags_with_table_alias(self):
|
||||
"""Should include table alias when provided."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, table_alias="mu.")
|
||||
assert "mu.tags" in result
|
||||
|
||||
# ---- Test "any" mode (OR, includes untagged - default) ----
|
||||
|
||||
def test_tags_match_any_includes_untagged(self):
|
||||
"""When match='any', should include untagged memories (NULL or empty)."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, match="any")
|
||||
# Should use OR with NULL/empty check
|
||||
assert "IS NULL" in result
|
||||
assert "= '{}'" in result
|
||||
assert "&&" in result # overlap operator
|
||||
|
||||
def test_tags_match_any_uses_overlap(self):
|
||||
"""When match='any', should use overlap operator (&&)."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, match="any")
|
||||
assert "&&" in result
|
||||
|
||||
# ---- Test "all" mode (AND, includes untagged) ----
|
||||
|
||||
def test_tags_match_all_includes_untagged(self):
|
||||
"""When match='all', should include untagged memories (NULL or empty)."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, match="all")
|
||||
# Should use OR with NULL/empty check
|
||||
assert "IS NULL" in result
|
||||
assert "= '{}'" in result
|
||||
assert "@>" in result # contains operator
|
||||
|
||||
def test_tags_match_all_uses_contains(self):
|
||||
"""When match='all', should use contains operator (@>)."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, match="all")
|
||||
assert "@>" in result
|
||||
|
||||
# ---- Test "any_strict" mode (OR, excludes untagged) ----
|
||||
|
||||
def test_tags_match_any_strict_excludes_untagged(self):
|
||||
"""When match='any_strict', should exclude untagged memories."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, match="any_strict")
|
||||
# Should require tags to be NOT NULL and not empty
|
||||
assert "IS NOT NULL" in result
|
||||
assert "!= '{}'" in result
|
||||
assert "&&" in result # overlap operator
|
||||
|
||||
def test_tags_match_any_strict_uses_overlap(self):
|
||||
"""When match='any_strict', should use overlap operator (&&)."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, match="any_strict")
|
||||
assert "&&" in result
|
||||
# Should NOT include untagged
|
||||
assert "IS NULL" not in result or "IS NOT NULL" in result
|
||||
|
||||
# ---- Test "all_strict" mode (AND, excludes untagged) ----
|
||||
|
||||
def test_tags_match_all_strict_excludes_untagged(self):
|
||||
"""When match='all_strict', should exclude untagged memories."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, match="all_strict")
|
||||
# Should require tags to be NOT NULL and not empty
|
||||
assert "IS NOT NULL" in result
|
||||
assert "!= '{}'" in result
|
||||
assert "@>" in result # contains operator
|
||||
|
||||
def test_tags_match_all_strict_uses_contains(self):
|
||||
"""When match='all_strict', should use contains operator (@>)."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 5, match="all_strict")
|
||||
assert "@>" in result
|
||||
|
||||
# ---- Test table alias with all modes ----
|
||||
|
||||
def test_tags_match_any_with_table_alias(self):
|
||||
"""Should include table alias with any mode."""
|
||||
result = build_tags_where_clause_simple(["user_a"], 3, table_alias="mu.", match="any")
|
||||
assert "mu.tags" in result
|
||||
|
||||
def test_tags_match_all_strict_with_table_alias(self):
|
||||
"""Should include table alias with all_strict mode."""
|
||||
result = build_tags_where_clause_simple(["user_a", "user_b"], 3, table_alias="mu.", match="all_strict")
|
||||
assert "mu.tags" in result
|
||||
assert "@>" in result
|
||||
assert "IS NOT NULL" in result
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Unit Tests for filter_results_by_tags (Python-side filtering)
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class MockResult:
|
||||
"""Mock result object for testing filter_results_by_tags."""
|
||||
|
||||
def __init__(self, tags):
|
||||
self.tags = tags
|
||||
|
||||
|
||||
class TestFilterResultsByTags:
|
||||
"""Unit tests for the Python-side tags filter function."""
|
||||
|
||||
def test_no_tags_returns_all(self):
|
||||
"""When tags is None, should return all results."""
|
||||
results = [MockResult(["a"]), MockResult(["b"]), MockResult(None)]
|
||||
filtered = filter_results_by_tags(results, None)
|
||||
assert len(filtered) == 3
|
||||
|
||||
def test_empty_tags_returns_all(self):
|
||||
"""When tags is empty list, should return all results."""
|
||||
results = [MockResult(["a"]), MockResult(["b"]), MockResult(None)]
|
||||
filtered = filter_results_by_tags(results, [])
|
||||
assert len(filtered) == 3
|
||||
|
||||
# ---- Test "any" mode (OR, includes untagged) ----
|
||||
|
||||
def test_any_mode_includes_matching_tags(self):
|
||||
"""'any' mode should include results with matching tags."""
|
||||
results = [MockResult(["a"]), MockResult(["b"]), MockResult(["c"])]
|
||||
filtered = filter_results_by_tags(results, ["a", "b"], match="any")
|
||||
# "a" and "b" match, "c" doesn't match and isn't untagged, so excluded
|
||||
assert len(filtered) == 2
|
||||
tags_found = [r.tags[0] for r in filtered if r.tags]
|
||||
assert "a" in tags_found
|
||||
assert "b" in tags_found
|
||||
assert "c" not in tags_found
|
||||
|
||||
def test_any_mode_includes_untagged(self):
|
||||
"""'any' mode should include untagged results."""
|
||||
results = [MockResult(["a"]), MockResult(None), MockResult([])]
|
||||
filtered = filter_results_by_tags(results, ["a"], match="any")
|
||||
assert len(filtered) == 3 # a matches, None is untagged, [] is untagged
|
||||
|
||||
def test_any_mode_includes_partial_overlap(self):
|
||||
"""'any' mode should include results with ANY overlapping tag."""
|
||||
results = [MockResult(["a", "x"]), MockResult(["b", "y"])]
|
||||
filtered = filter_results_by_tags(results, ["a"], match="any")
|
||||
# ["a", "x"] matches, ["b", "y"] doesn't, but untagged would be included
|
||||
tags_found = [r.tags for r in filtered]
|
||||
assert ["a", "x"] in tags_found
|
||||
|
||||
# ---- Test "any_strict" mode (OR, excludes untagged) ----
|
||||
|
||||
def test_any_strict_excludes_untagged(self):
|
||||
"""'any_strict' mode should exclude untagged results."""
|
||||
results = [MockResult(["a"]), MockResult(None), MockResult([])]
|
||||
filtered = filter_results_by_tags(results, ["a"], match="any_strict")
|
||||
assert len(filtered) == 1 # Only ["a"] matches
|
||||
assert filtered[0].tags == ["a"]
|
||||
|
||||
def test_any_strict_excludes_non_matching(self):
|
||||
"""'any_strict' mode should exclude non-matching tagged results."""
|
||||
results = [MockResult(["a"]), MockResult(["b"]), MockResult(["c"])]
|
||||
filtered = filter_results_by_tags(results, ["a"], match="any_strict")
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0].tags == ["a"]
|
||||
|
||||
# ---- Test "all" mode (AND, includes untagged) ----
|
||||
|
||||
def test_all_mode_requires_all_tags(self):
|
||||
"""'all' mode should require ALL requested tags to be present."""
|
||||
results = [MockResult(["a", "b"]), MockResult(["a"]), MockResult(["b"])]
|
||||
filtered = filter_results_by_tags(results, ["a", "b"], match="all")
|
||||
# Only ["a", "b"] has both tags, but untagged would also be included
|
||||
tags_found = [r.tags for r in filtered]
|
||||
assert ["a", "b"] in tags_found
|
||||
|
||||
def test_all_mode_includes_untagged(self):
|
||||
"""'all' mode should include untagged results."""
|
||||
results = [MockResult(["a", "b"]), MockResult(None), MockResult([])]
|
||||
filtered = filter_results_by_tags(results, ["a", "b"], match="all")
|
||||
assert len(filtered) == 3 # ["a", "b"] matches, None is untagged, [] is untagged
|
||||
|
||||
# ---- Test "all_strict" mode (AND, excludes untagged) ----
|
||||
|
||||
def test_all_strict_requires_all_tags(self):
|
||||
"""'all_strict' mode should require ALL requested tags."""
|
||||
results = [MockResult(["a", "b"]), MockResult(["a"]), MockResult(["b"])]
|
||||
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0].tags == ["a", "b"]
|
||||
|
||||
def test_all_strict_excludes_untagged(self):
|
||||
"""'all_strict' mode should exclude untagged results."""
|
||||
results = [MockResult(["a", "b"]), MockResult(None), MockResult([])]
|
||||
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0].tags == ["a", "b"]
|
||||
|
||||
def test_all_strict_allows_superset(self):
|
||||
"""'all_strict' mode should allow results with MORE tags than requested."""
|
||||
results = [MockResult(["a", "b", "c"]), MockResult(["a"])]
|
||||
filtered = filter_results_by_tags(results, ["a", "b"], match="all_strict")
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0].tags == ["a", "b", "c"] # Has a, b, AND c
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Integration Tests for tags in retain/recall/reflect
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def api_client(memory):
|
||||
"""Create an async test client for the FastAPI app."""
|
||||
app = create_app(memory, initialize_memory=False)
|
||||
transport = httpx.ASGITransport(app=app)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
yield client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def test_bank_id():
|
||||
"""Provide a unique bank ID for this test run."""
|
||||
return f"tags_test_{datetime.now().timestamp()}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_with_tags(api_client, test_bank_id):
|
||||
"""Test that memories can be stored with tags."""
|
||||
# Store memory with tags
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{
|
||||
"content": "Alice loves hiking in the mountains.",
|
||||
"tags": ["user_alice"]
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
assert result["success"] is True
|
||||
assert result["items_count"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_with_document_tags(api_client, test_bank_id):
|
||||
"""Test that document-level tags are applied to all items."""
|
||||
# Store memories with document-level tags
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"document_tags": ["session_123"],
|
||||
"items": [
|
||||
{"content": "Bob discussed the quarterly report."},
|
||||
{"content": "Charlie mentioned the new product launch."}
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
assert result["success"] is True
|
||||
assert result["items_count"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_merges_document_and_item_tags(api_client, test_bank_id):
|
||||
"""Test that document tags and item tags are merged."""
|
||||
# Store memory with both document and item tags
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"document_tags": ["session_abc"],
|
||||
"items": [
|
||||
{
|
||||
"content": "Dave talked about machine learning.",
|
||||
"tags": ["user_dave"]
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
assert result["success"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_without_tags_returns_all_memories(api_client, test_bank_id):
|
||||
"""Test that recall without tags returns all memories (no filtering)."""
|
||||
# Store memories for different users
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Eve works on natural language processing.", "tags": ["user_eve"]},
|
||||
{"content": "Frank specializes in computer vision.", "tags": ["user_frank"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# Recall without tags - should return all
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/recall",
|
||||
json={"query": "Who works on what?", "budget": "low"}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
results = response.json()["results"]
|
||||
|
||||
# Should find both Eve and Frank
|
||||
texts = [r["text"] for r in results]
|
||||
assert any("Eve" in t for t in texts), "Should find Eve"
|
||||
assert any("Frank" in t for t in texts), "Should find Frank"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_with_tags_filters_memories(api_client, test_bank_id):
|
||||
"""Test that recall with tags only returns matching memories."""
|
||||
# Store memories for different users
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Grace is a data scientist at Google.", "tags": ["user_grace"]},
|
||||
{"content": "Henry is a software engineer at Meta.", "tags": ["user_henry"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# Recall with user_grace tag - should only return Grace's memory
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/recall",
|
||||
json={"query": "Who works at which company?", "budget": "low", "tags": ["user_grace"]}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
results = response.json()["results"]
|
||||
|
||||
# Should find Grace but not Henry
|
||||
texts = [r["text"] for r in results]
|
||||
assert any("Grace" in t for t in texts), "Should find Grace with user_grace tag"
|
||||
# Henry should NOT be found since he has user_henry tag
|
||||
assert not any("Henry" in t for t in texts), "Should NOT find Henry (different tag)"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_with_multiple_tags_uses_or_matching(api_client, test_bank_id):
|
||||
"""Test that multiple tags use OR matching (any match returns the memory)."""
|
||||
# Store memories for different users
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Ivan leads the security team.", "tags": ["user_ivan"]},
|
||||
{"content": "Julia manages the design team.", "tags": ["user_julia"]},
|
||||
{"content": "Karl oversees the marketing team.", "tags": ["user_karl"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# Recall with user_ivan OR user_julia - should return both Ivan and Julia, but not Karl
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/recall",
|
||||
json={"query": "Who leads which team?", "budget": "low", "tags": ["user_ivan", "user_julia"]}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
results = response.json()["results"]
|
||||
|
||||
texts = [r["text"] for r in results]
|
||||
assert any("Ivan" in t for t in texts), "Should find Ivan (tag matches)"
|
||||
assert any("Julia" in t for t in texts), "Should find Julia (tag matches)"
|
||||
assert not any("Karl" in t for t in texts), "Should NOT find Karl (tag doesn't match)"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_returns_memories_with_any_overlapping_tag(api_client, test_bank_id):
|
||||
"""Test that memories with multiple tags are returned if ANY tag matches."""
|
||||
# Store memory with multiple tags
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{
|
||||
"content": "Lisa and Mike discussed the budget in a group chat.",
|
||||
"tags": ["user_lisa", "user_mike"] # Memory visible to both
|
||||
},
|
||||
{"content": "Nancy reviewed the budget alone.", "tags": ["user_nancy"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# Recall with user_lisa - should return the group chat memory
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/recall",
|
||||
json={"query": "What was discussed about the budget?", "budget": "low", "tags": ["user_lisa"]}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
results = response.json()["results"]
|
||||
|
||||
texts = [r["text"] for r in results]
|
||||
assert any("Lisa" in t and "Mike" in t for t in texts), "Should find group chat (Lisa is in tags)"
|
||||
assert not any("Nancy" in t for t in texts), "Should NOT find Nancy's memory"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reflect_with_tags_filters_memories(api_client, test_bank_id):
|
||||
"""Test that reflect with tags only uses matching memories for reasoning."""
|
||||
# Store different memories for different users
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Oscar's favorite color is blue.", "tags": ["user_oscar"]},
|
||||
{"content": "Peter's favorite color is red.", "tags": ["user_peter"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# Reflect with user_oscar tag - should only use Oscar's memories
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/reflect",
|
||||
json={
|
||||
"query": "What is the favorite color?",
|
||||
"budget": "low",
|
||||
"tags": ["user_oscar"],
|
||||
"include": {"facts": {}} # Request facts to verify what was used
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
|
||||
# The response should mention Oscar's color (blue), not Peter's (red)
|
||||
# Note: We can check based_on facts if they're returned
|
||||
if result.get("based_on"):
|
||||
fact_texts = [f["text"] for f in result["based_on"]]
|
||||
# Should use Oscar's memory
|
||||
assert any("Oscar" in t or "blue" in t for t in fact_texts), "Should use Oscar's memory"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recall_with_empty_tags_returns_all(api_client, test_bank_id):
|
||||
"""Test that empty tags list behaves same as no tags (returns all)."""
|
||||
# Store memories
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Quinn studies mathematics.", "tags": ["user_quinn"]},
|
||||
{"content": "Rachel studies physics.", "tags": ["user_rachel"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# Recall with empty tags list - should return all
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/recall",
|
||||
json={"query": "Who studies what?", "budget": "low", "tags": []}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
results = response.json()["results"]
|
||||
|
||||
texts = [r["text"] for r in results]
|
||||
assert any("Quinn" in t for t in texts), "Should find Quinn"
|
||||
assert any("Rachel" in t for t in texts), "Should find Rachel"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multi_user_agent_visibility(api_client):
|
||||
"""
|
||||
Test multi-user agent visibility scoping.
|
||||
|
||||
Scenario:
|
||||
- Agent has one memory bank
|
||||
- Agent chats with User A (room 1) and User B (room 2) separately
|
||||
- Agent also hosts a group chat with both users (room 3)
|
||||
- User A should only see memories from rooms 1 and 3
|
||||
- User B should only see memories from rooms 2 and 3
|
||||
- Agent (no filter) should see all memories
|
||||
"""
|
||||
bank_id = f"multi_user_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories from different chat rooms
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
# Room 1: Agent + User A private chat
|
||||
{"content": "User A said they prefer morning meetings.", "tags": ["user_a"]},
|
||||
# Room 2: Agent + User B private chat
|
||||
{"content": "User B mentioned they like afternoon meetings.", "tags": ["user_b"]},
|
||||
# Room 3: Group chat with both users
|
||||
{"content": "In the group meeting, they agreed to meet at noon.", "tags": ["user_a", "user_b"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# User A queries - should see their private chat and group chat
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories/recall",
|
||||
json={"query": "What meeting time preferences were discussed?", "budget": "low", "tags": ["user_a"]}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
user_a_results = response.json()["results"]
|
||||
user_a_texts = [r["text"] for r in user_a_results]
|
||||
|
||||
assert any("morning" in t for t in user_a_texts), "User A should see their own preference (morning)"
|
||||
assert any("noon" in t for t in user_a_texts), "User A should see group chat (noon)"
|
||||
assert not any("afternoon" in t for t in user_a_texts), "User A should NOT see User B's private preference"
|
||||
|
||||
# User B queries - should see their private chat and group chat
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories/recall",
|
||||
json={"query": "What meeting time preferences were discussed?", "budget": "low", "tags": ["user_b"]}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
user_b_results = response.json()["results"]
|
||||
user_b_texts = [r["text"] for r in user_b_results]
|
||||
|
||||
assert any("afternoon" in t for t in user_b_texts), "User B should see their own preference (afternoon)"
|
||||
assert any("noon" in t for t in user_b_texts), "User B should see group chat (noon)"
|
||||
assert not any("morning" in t for t in user_b_texts), "User B should NOT see User A's private preference"
|
||||
|
||||
# Agent queries (no filter) - should see everything
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories/recall",
|
||||
json={"query": "What meeting time preferences were discussed?", "budget": "low"} # No tags
|
||||
)
|
||||
assert response.status_code == 200
|
||||
agent_results = response.json()["results"]
|
||||
agent_texts = [r["text"] for r in agent_results]
|
||||
|
||||
assert any("morning" in t for t in agent_texts), "Agent should see User A's preference"
|
||||
assert any("afternoon" in t for t in agent_texts), "Agent should see User B's preference"
|
||||
assert any("noon" in t for t in agent_texts), "Agent should see group chat"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_student_tracking_visibility(api_client):
|
||||
"""
|
||||
Test student tracking visibility scoping.
|
||||
|
||||
Scenario:
|
||||
- Teacher bot has one memory bank
|
||||
- Teacher records observations for Student A, Student B
|
||||
- Student A should only see their own data
|
||||
- Teacher (no filter) should see all student data
|
||||
"""
|
||||
bank_id = f"student_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories for different students
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Student A showed improvement in algebra today.", "tags": ["student_a"]},
|
||||
{"content": "Student B struggled with geometry concepts.", "tags": ["student_b"]},
|
||||
{"content": "Student A participated actively in class discussion.", "tags": ["student_a"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# Student A queries - should only see their own data
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories/recall",
|
||||
json={"query": "How am I doing in class?", "budget": "low", "tags": ["student_a"]}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
student_a_results = response.json()["results"]
|
||||
student_a_texts = [r["text"] for r in student_a_results]
|
||||
|
||||
assert any("algebra" in t for t in student_a_texts), "Student A should see their algebra progress"
|
||||
assert any("participated" in t for t in student_a_texts), "Student A should see their participation"
|
||||
assert not any("Student B" in t or "geometry" in t for t in student_a_texts), "Student A should NOT see Student B's data"
|
||||
|
||||
# Teacher queries (no filter) - should see all students
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories/recall",
|
||||
json={"query": "Which students need help?", "budget": "low"} # No tags
|
||||
)
|
||||
assert response.status_code == 200
|
||||
teacher_results = response.json()["results"]
|
||||
teacher_texts = [r["text"] for r in teacher_results]
|
||||
|
||||
assert any("Student A" in t for t in teacher_texts), "Teacher should see Student A's data"
|
||||
assert any("Student B" in t for t in teacher_texts), "Teacher should see Student B's data"
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for list_tags API endpoint
|
||||
# ============================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tags_returns_all_tags(api_client):
|
||||
"""Test that list_tags returns all unique tags with counts."""
|
||||
bank_id = f"list_tags_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories with various tags
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Memory 1 for user alice.", "tags": ["user:alice"]},
|
||||
{"content": "Memory 2 for user alice.", "tags": ["user:alice"]},
|
||||
{"content": "Memory 3 for user bob.", "tags": ["user:bob"]},
|
||||
{"content": "Memory 4 in session 123.", "tags": ["session:123"]},
|
||||
{"content": "Memory 5 for alice in session 456.", "tags": ["user:alice", "session:456"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# List all tags
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
|
||||
# Verify structure
|
||||
assert "items" in result
|
||||
assert "total" in result
|
||||
assert "limit" in result
|
||||
assert "offset" in result
|
||||
|
||||
# Verify tags and counts
|
||||
tags_map = {item["tag"]: item["count"] for item in result["items"]}
|
||||
assert "user:alice" in tags_map
|
||||
assert tags_map["user:alice"] == 3 # 3 memories have this tag
|
||||
assert "user:bob" in tags_map
|
||||
assert tags_map["user:bob"] == 1
|
||||
assert "session:123" in tags_map
|
||||
assert tags_map["session:123"] == 1
|
||||
assert "session:456" in tags_map
|
||||
assert tags_map["session:456"] == 1
|
||||
|
||||
assert result["total"] == 4 # 4 unique tags
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tags_with_wildcard_prefix(api_client):
|
||||
"""Test that list_tags filters with prefix wildcard pattern (user:*)."""
|
||||
bank_id = f"list_tags_wildcard_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories with various tags
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Memory for alice who works at tech.", "tags": ["user:alice"]},
|
||||
{"content": "Memory for bob who is an engineer.", "tags": ["user:bob"]},
|
||||
{"content": "Memory for charlie the designer.", "tags": ["user:charlie"]},
|
||||
{"content": "Session memory about the meeting.", "tags": ["session:abc"]},
|
||||
{"content": "Room memory for conference room.", "tags": ["room:123"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# List tags with 'user:*' wildcard pattern
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "user:*"})
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
|
||||
# Should only return user:* tags
|
||||
tags = [item["tag"] for item in result["items"]]
|
||||
assert "user:alice" in tags
|
||||
assert "user:bob" in tags
|
||||
assert "user:charlie" in tags
|
||||
assert "session:abc" not in tags
|
||||
assert "room:123" not in tags
|
||||
assert result["total"] == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tags_with_wildcard_suffix(api_client):
|
||||
"""Test that list_tags filters with suffix wildcard pattern (*-admin)."""
|
||||
bank_id = f"list_tags_suffix_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories with various tags
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Admin role memory for super admin.", "tags": ["role-admin"]},
|
||||
{"content": "Super admin memory about permissions.", "tags": ["super-admin"]},
|
||||
{"content": "User memory for standard users.", "tags": ["role-user"]},
|
||||
{"content": "Guest memory for visitors.", "tags": ["role-guest"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# List tags with '*-admin' wildcard pattern (suffix match)
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "*-admin"})
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
|
||||
# Should only return *-admin tags
|
||||
tags = [item["tag"] for item in result["items"]]
|
||||
assert "role-admin" in tags
|
||||
assert "super-admin" in tags
|
||||
assert "role-user" not in tags
|
||||
assert "role-guest" not in tags
|
||||
assert result["total"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tags_with_wildcard_middle(api_client):
|
||||
"""Test that list_tags filters with middle wildcard pattern (env*-prod)."""
|
||||
bank_id = f"list_tags_middle_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories with various tags - use meaningful content for fact extraction
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "The production environment is configured with high availability and uses AWS infrastructure.", "tags": ["env-prod"]},
|
||||
{"content": "The enterprise environment for production runs on dedicated servers with 24/7 monitoring.", "tags": ["environment-prod"]},
|
||||
{"content": "The staging environment mirrors production but uses smaller instance sizes.", "tags": ["env-staging"]},
|
||||
{"content": "The development environment allows developers to test their code locally.", "tags": ["env-dev"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# List tags with 'env*-prod' wildcard pattern (middle match)
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "env*-prod"})
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
|
||||
# Should only return env*-prod tags
|
||||
tags = [item["tag"] for item in result["items"]]
|
||||
assert "env-prod" in tags
|
||||
assert "environment-prod" in tags
|
||||
assert "env-staging" not in tags
|
||||
assert "env-dev" not in tags
|
||||
assert result["total"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tags_case_insensitive(api_client):
|
||||
"""Test that list_tags wildcard matching is case-insensitive."""
|
||||
bank_id = f"list_tags_case_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories with mixed case tags - use meaningful content
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Alice is a software engineer who specializes in machine learning algorithms.", "tags": ["User:Alice"]},
|
||||
{"content": "Bob works as a data scientist at a large technology company.", "tags": ["user:bob"]},
|
||||
{"content": "Charlie is the lead designer responsible for the user interface.", "tags": ["USER:CHARLIE"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# List tags with lowercase pattern - should match all cases
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"q": "user:*"})
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
|
||||
# Should match all user tags regardless of case
|
||||
tags = [item["tag"] for item in result["items"]]
|
||||
assert len(tags) == 3
|
||||
assert result["total"] == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tags_pagination(api_client):
|
||||
"""Test that list_tags supports pagination."""
|
||||
bank_id = f"list_tags_pagination_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories with many tags - use meaningful content for fact extraction
|
||||
names = ["Alice", "Bob", "Charlie", "Diana", "Eve", "Frank", "Grace", "Henry", "Ivan", "Julia"]
|
||||
items = [
|
||||
{"content": f"{name} works as a software engineer at company {i}.", "tags": [f"tag:{i:03d}"]}
|
||||
for i, name in enumerate(names)
|
||||
]
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={"items": items}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# Get first page (limit 3)
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"limit": 3, "offset": 0})
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
assert len(result["items"]) == 3
|
||||
assert result["total"] == 10
|
||||
assert result["limit"] == 3
|
||||
assert result["offset"] == 0
|
||||
|
||||
# Get second page
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags", params={"limit": 3, "offset": 3})
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
assert len(result["items"]) == 3
|
||||
assert result["offset"] == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tags_empty_bank(api_client):
|
||||
"""Test that list_tags returns empty for bank with no tags."""
|
||||
bank_id = f"list_tags_empty_test_{datetime.now().timestamp()}"
|
||||
|
||||
# List tags without storing anything
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
|
||||
assert result["items"] == []
|
||||
assert result["total"] == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tags_ordered_by_count(api_client):
|
||||
"""Test that list_tags returns tags ordered by frequency (most used first)."""
|
||||
bank_id = f"list_tags_order_test_{datetime.now().timestamp()}"
|
||||
|
||||
# Store memories with tags having different frequencies - use meaningful content
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{"content": "Alice works at a startup company as a developer.", "tags": ["rare"]},
|
||||
{"content": "Bob is a senior engineer at Google.", "tags": ["common"]},
|
||||
{"content": "Charlie manages the marketing team at Microsoft.", "tags": ["common"]},
|
||||
{"content": "Diana leads the design department at Apple.", "tags": ["common"]},
|
||||
{"content": "Eve is a data scientist at Amazon.", "tags": ["medium"]},
|
||||
{"content": "Frank handles customer support at Meta.", "tags": ["medium"]},
|
||||
]
|
||||
}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
# List tags - should be ordered by count descending
|
||||
response = await api_client.get(f"/v1/default/banks/{bank_id}/tags")
|
||||
assert response.status_code == 200
|
||||
result = response.json()
|
||||
|
||||
tags = [item["tag"] for item in result["items"]]
|
||||
# common (3) should come before medium (2) which should come before rare (1)
|
||||
assert tags.index("common") < tags.index("medium")
|
||||
assert tags.index("medium") < tags.index("rare")
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "hindsight-cli"
|
||||
version = "0.2.1"
|
||||
version = "0.3.0"
|
||||
edition = "2021"
|
||||
authors = ["Hindsight Team"]
|
||||
description = "A beautiful CLI for Hindsight - semantic memory system"
|
||||
|
||||
@@ -5,7 +5,7 @@ use crossterm::{
|
||||
execute,
|
||||
terminal::{disable_raw_mode, enable_raw_mode, EnterAlternateScreen, LeaveAlternateScreen},
|
||||
};
|
||||
use hindsight_client::types::{BankListItem, RecallResult, EntityListItem, Budget};
|
||||
use hindsight_client::types::{BankListItem, RecallResult, EntityListItem, Budget, TagsMatch};
|
||||
use serde_json::{Map, Value};
|
||||
use ratatui::{
|
||||
backend::{Backend, CrosstermBackend},
|
||||
@@ -341,6 +341,8 @@ impl App {
|
||||
trace: false,
|
||||
query_timestamp: None,
|
||||
include: None,
|
||||
tags: None,
|
||||
tags_match: TagsMatch::Any,
|
||||
};
|
||||
|
||||
let result = client.recall(&bank_id, &request, false)
|
||||
@@ -357,6 +359,8 @@ impl App {
|
||||
max_tokens: 4096,
|
||||
include: None,
|
||||
response_schema: None,
|
||||
tags: None,
|
||||
tags_match: TagsMatch::Any,
|
||||
};
|
||||
|
||||
let result = client.reflect(&bank_id, &request, false)
|
||||
|
||||
@@ -9,7 +9,7 @@ use crate::output::{self, OutputFormat};
|
||||
use crate::ui;
|
||||
|
||||
// Import types from generated client
|
||||
use hindsight_client::types::{Budget, ChunkIncludeOptions, IncludeOptions};
|
||||
use hindsight_client::types::{Budget, ChunkIncludeOptions, IncludeOptions, TagsMatch};
|
||||
use serde_json;
|
||||
|
||||
// Helper function to parse budget string to Budget enum
|
||||
@@ -60,6 +60,8 @@ pub fn recall(
|
||||
trace,
|
||||
query_timestamp: None,
|
||||
include,
|
||||
tags: None,
|
||||
tags_match: TagsMatch::Any,
|
||||
};
|
||||
|
||||
let response = client.recall(agent_id, &request, verbose);
|
||||
@@ -116,6 +118,8 @@ pub fn reflect(
|
||||
max_tokens: max_tokens.unwrap_or(4096),
|
||||
include: None,
|
||||
response_schema,
|
||||
tags: None,
|
||||
tags_match: TagsMatch::Any,
|
||||
};
|
||||
|
||||
let response = client.reflect(agent_id, &request, verbose);
|
||||
@@ -162,11 +166,13 @@ pub fn retain(
|
||||
timestamp: None,
|
||||
document_id: Some(doc_id.clone()),
|
||||
entities: None,
|
||||
tags: None,
|
||||
};
|
||||
|
||||
let request = RetainRequest {
|
||||
items: vec![item],
|
||||
async_: r#async,
|
||||
document_tags: None,
|
||||
};
|
||||
|
||||
let response = client.retain(agent_id, &request, r#async, verbose);
|
||||
@@ -272,6 +278,7 @@ pub fn retain_files(
|
||||
timestamp: None,
|
||||
document_id: Some(doc_id),
|
||||
entities: None,
|
||||
tags: None,
|
||||
});
|
||||
|
||||
pb.inc(1);
|
||||
@@ -288,6 +295,7 @@ pub fn retain_files(
|
||||
let request = RetainRequest {
|
||||
items,
|
||||
async_: r#async,
|
||||
document_tags: None,
|
||||
};
|
||||
|
||||
let response = client.retain(agent_id, &request, r#async, verbose);
|
||||
|
||||
@@ -39,6 +39,7 @@ hindsight_client_api/models/http_validation_error.py
|
||||
hindsight_client_api/models/include_options.py
|
||||
hindsight_client_api/models/list_documents_response.py
|
||||
hindsight_client_api/models/list_memory_units_response.py
|
||||
hindsight_client_api/models/list_tags_response.py
|
||||
hindsight_client_api/models/memory_item.py
|
||||
hindsight_client_api/models/operation_response.py
|
||||
hindsight_client_api/models/operations_list_response.py
|
||||
@@ -51,6 +52,7 @@ hindsight_client_api/models/reflect_request.py
|
||||
hindsight_client_api/models/reflect_response.py
|
||||
hindsight_client_api/models/retain_request.py
|
||||
hindsight_client_api/models/retain_response.py
|
||||
hindsight_client_api/models/tag_item.py
|
||||
hindsight_client_api/models/token_usage.py
|
||||
hindsight_client_api/models/update_disposition_request.py
|
||||
hindsight_client_api/models/validation_error.py
|
||||
|
||||
@@ -64,6 +64,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
|
||||
from hindsight_client_api.models.include_options import IncludeOptions
|
||||
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.list_tags_response import ListTagsResponse
|
||||
from hindsight_client_api.models.memory_item import MemoryItem
|
||||
from hindsight_client_api.models.operation_response import OperationResponse
|
||||
from hindsight_client_api.models.operations_list_response import OperationsListResponse
|
||||
@@ -76,6 +77,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
|
||||
from hindsight_client_api.models.reflect_response import ReflectResponse
|
||||
from hindsight_client_api.models.retain_request import RetainRequest
|
||||
from hindsight_client_api.models.retain_response import RetainResponse
|
||||
from hindsight_client_api.models.tag_item import TagItem
|
||||
from hindsight_client_api.models.token_usage import TokenUsage
|
||||
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
|
||||
from hindsight_client_api.models.validation_error import ValidationError
|
||||
|
||||
@@ -17,11 +17,12 @@ from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from typing_extensions import Annotated
|
||||
|
||||
from pydantic import Field, StrictInt, StrictStr
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
from typing_extensions import Annotated
|
||||
from hindsight_client_api.models.delete_response import DeleteResponse
|
||||
from hindsight_client_api.models.graph_data_response import GraphDataResponse
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.list_tags_response import ListTagsResponse
|
||||
from hindsight_client_api.models.recall_request import RecallRequest
|
||||
from hindsight_client_api.models.recall_response import RecallResponse
|
||||
from hindsight_client_api.models.reflect_request import ReflectRequest
|
||||
@@ -654,6 +655,299 @@ class MemoryApi:
|
||||
|
||||
|
||||
|
||||
@validate_call
|
||||
async def get_memory(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
memory_id: StrictStr,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> object:
|
||||
"""Get memory unit
|
||||
|
||||
Get a single memory unit by ID with all its metadata including entities and tags.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param memory_id: (required)
|
||||
:type memory_id: str
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._get_memory_serialize(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "object",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
await response_data.read()
|
||||
return self.api_client.response_deserialize(
|
||||
response_data=response_data,
|
||||
response_types_map=_response_types_map,
|
||||
).data
|
||||
|
||||
|
||||
@validate_call
|
||||
async def get_memory_with_http_info(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
memory_id: StrictStr,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> ApiResponse[object]:
|
||||
"""Get memory unit
|
||||
|
||||
Get a single memory unit by ID with all its metadata including entities and tags.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param memory_id: (required)
|
||||
:type memory_id: str
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._get_memory_serialize(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "object",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
await response_data.read()
|
||||
return self.api_client.response_deserialize(
|
||||
response_data=response_data,
|
||||
response_types_map=_response_types_map,
|
||||
)
|
||||
|
||||
|
||||
@validate_call
|
||||
async def get_memory_without_preload_content(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
memory_id: StrictStr,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> RESTResponseType:
|
||||
"""Get memory unit
|
||||
|
||||
Get a single memory unit by ID with all its metadata including entities and tags.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param memory_id: (required)
|
||||
:type memory_id: str
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._get_memory_serialize(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "object",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
return response_data.response
|
||||
|
||||
|
||||
def _get_memory_serialize(
|
||||
self,
|
||||
bank_id,
|
||||
memory_id,
|
||||
authorization,
|
||||
_request_auth,
|
||||
_content_type,
|
||||
_headers,
|
||||
_host_index,
|
||||
) -> RequestSerialized:
|
||||
|
||||
_host = None
|
||||
|
||||
_collection_formats: Dict[str, str] = {
|
||||
}
|
||||
|
||||
_path_params: Dict[str, str] = {}
|
||||
_query_params: List[Tuple[str, str]] = []
|
||||
_header_params: Dict[str, Optional[str]] = _headers or {}
|
||||
_form_params: List[Tuple[str, str]] = []
|
||||
_files: Dict[
|
||||
str, Union[str, bytes, List[str], List[bytes], List[Tuple[str, bytes]]]
|
||||
] = {}
|
||||
_body_params: Optional[bytes] = None
|
||||
|
||||
# process the path parameters
|
||||
if bank_id is not None:
|
||||
_path_params['bank_id'] = bank_id
|
||||
if memory_id is not None:
|
||||
_path_params['memory_id'] = memory_id
|
||||
# process the query parameters
|
||||
# process the header parameters
|
||||
if authorization is not None:
|
||||
_header_params['authorization'] = authorization
|
||||
# process the form parameters
|
||||
# process the body parameter
|
||||
|
||||
|
||||
# set the HTTP header `Accept`
|
||||
if 'Accept' not in _header_params:
|
||||
_header_params['Accept'] = self.api_client.select_header_accept(
|
||||
[
|
||||
'application/json'
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# authentication setting
|
||||
_auth_settings: List[str] = [
|
||||
]
|
||||
|
||||
return self.api_client.param_serialize(
|
||||
method='GET',
|
||||
resource_path='/v1/default/banks/{bank_id}/memories/{memory_id}',
|
||||
path_params=_path_params,
|
||||
query_params=_query_params,
|
||||
header_params=_header_params,
|
||||
body=_body_params,
|
||||
post_params=_form_params,
|
||||
files=_files,
|
||||
auth_settings=_auth_settings,
|
||||
collection_formats=_collection_formats,
|
||||
_host=_host,
|
||||
_request_auth=_request_auth
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
@validate_call
|
||||
async def list_memories(
|
||||
self,
|
||||
@@ -1000,6 +1294,335 @@ class MemoryApi:
|
||||
|
||||
|
||||
|
||||
@validate_call
|
||||
async def list_tags(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
|
||||
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
|
||||
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> ListTagsResponse:
|
||||
"""List tags
|
||||
|
||||
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
|
||||
:type q: str
|
||||
:param limit: Maximum number of tags to return
|
||||
:type limit: int
|
||||
:param offset: Offset for pagination
|
||||
:type offset: int
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._list_tags_serialize(
|
||||
bank_id=bank_id,
|
||||
q=q,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "ListTagsResponse",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
await response_data.read()
|
||||
return self.api_client.response_deserialize(
|
||||
response_data=response_data,
|
||||
response_types_map=_response_types_map,
|
||||
).data
|
||||
|
||||
|
||||
@validate_call
|
||||
async def list_tags_with_http_info(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
|
||||
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
|
||||
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> ApiResponse[ListTagsResponse]:
|
||||
"""List tags
|
||||
|
||||
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
|
||||
:type q: str
|
||||
:param limit: Maximum number of tags to return
|
||||
:type limit: int
|
||||
:param offset: Offset for pagination
|
||||
:type offset: int
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._list_tags_serialize(
|
||||
bank_id=bank_id,
|
||||
q=q,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "ListTagsResponse",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
await response_data.read()
|
||||
return self.api_client.response_deserialize(
|
||||
response_data=response_data,
|
||||
response_types_map=_response_types_map,
|
||||
)
|
||||
|
||||
|
||||
@validate_call
|
||||
async def list_tags_without_preload_content(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
|
||||
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
|
||||
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> RESTResponseType:
|
||||
"""List tags
|
||||
|
||||
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
|
||||
:type q: str
|
||||
:param limit: Maximum number of tags to return
|
||||
:type limit: int
|
||||
:param offset: Offset for pagination
|
||||
:type offset: int
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._list_tags_serialize(
|
||||
bank_id=bank_id,
|
||||
q=q,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "ListTagsResponse",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
return response_data.response
|
||||
|
||||
|
||||
def _list_tags_serialize(
|
||||
self,
|
||||
bank_id,
|
||||
q,
|
||||
limit,
|
||||
offset,
|
||||
authorization,
|
||||
_request_auth,
|
||||
_content_type,
|
||||
_headers,
|
||||
_host_index,
|
||||
) -> RequestSerialized:
|
||||
|
||||
_host = None
|
||||
|
||||
_collection_formats: Dict[str, str] = {
|
||||
}
|
||||
|
||||
_path_params: Dict[str, str] = {}
|
||||
_query_params: List[Tuple[str, str]] = []
|
||||
_header_params: Dict[str, Optional[str]] = _headers or {}
|
||||
_form_params: List[Tuple[str, str]] = []
|
||||
_files: Dict[
|
||||
str, Union[str, bytes, List[str], List[bytes], List[Tuple[str, bytes]]]
|
||||
] = {}
|
||||
_body_params: Optional[bytes] = None
|
||||
|
||||
# process the path parameters
|
||||
if bank_id is not None:
|
||||
_path_params['bank_id'] = bank_id
|
||||
# process the query parameters
|
||||
if q is not None:
|
||||
|
||||
_query_params.append(('q', q))
|
||||
|
||||
if limit is not None:
|
||||
|
||||
_query_params.append(('limit', limit))
|
||||
|
||||
if offset is not None:
|
||||
|
||||
_query_params.append(('offset', offset))
|
||||
|
||||
# process the header parameters
|
||||
if authorization is not None:
|
||||
_header_params['authorization'] = authorization
|
||||
# process the form parameters
|
||||
# process the body parameter
|
||||
|
||||
|
||||
# set the HTTP header `Accept`
|
||||
if 'Accept' not in _header_params:
|
||||
_header_params['Accept'] = self.api_client.select_header_accept(
|
||||
[
|
||||
'application/json'
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# authentication setting
|
||||
_auth_settings: List[str] = [
|
||||
]
|
||||
|
||||
return self.api_client.param_serialize(
|
||||
method='GET',
|
||||
resource_path='/v1/default/banks/{bank_id}/tags',
|
||||
path_params=_path_params,
|
||||
query_params=_query_params,
|
||||
header_params=_header_params,
|
||||
body=_body_params,
|
||||
post_params=_form_params,
|
||||
files=_files,
|
||||
auth_settings=_auth_settings,
|
||||
collection_formats=_collection_formats,
|
||||
_host=_host,
|
||||
_request_auth=_request_auth
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
@validate_call
|
||||
async def recall_memories(
|
||||
self,
|
||||
|
||||
@@ -42,6 +42,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
|
||||
from hindsight_client_api.models.include_options import IncludeOptions
|
||||
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.list_tags_response import ListTagsResponse
|
||||
from hindsight_client_api.models.memory_item import MemoryItem
|
||||
from hindsight_client_api.models.operation_response import OperationResponse
|
||||
from hindsight_client_api.models.operations_list_response import OperationsListResponse
|
||||
@@ -54,6 +55,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
|
||||
from hindsight_client_api.models.reflect_response import ReflectResponse
|
||||
from hindsight_client_api.models.retain_request import RetainRequest
|
||||
from hindsight_client_api.models.retain_response import RetainResponse
|
||||
from hindsight_client_api.models.tag_item import TagItem
|
||||
from hindsight_client_api.models.token_usage import TokenUsage
|
||||
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
|
||||
from hindsight_client_api.models.validation_error import ValidationError
|
||||
|
||||
@@ -17,7 +17,7 @@ import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, StrictInt, StrictStr
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
|
||||
from typing import Any, ClassVar, Dict, List, Optional
|
||||
from typing import Optional, Set
|
||||
from typing_extensions import Self
|
||||
@@ -33,7 +33,8 @@ class DocumentResponse(BaseModel):
|
||||
created_at: StrictStr
|
||||
updated_at: StrictStr
|
||||
memory_unit_count: StrictInt
|
||||
__properties: ClassVar[List[str]] = ["id", "bank_id", "original_text", "content_hash", "created_at", "updated_at", "memory_unit_count"]
|
||||
tags: Optional[List[StrictStr]] = Field(default=None, description="Tags associated with this document")
|
||||
__properties: ClassVar[List[str]] = ["id", "bank_id", "original_text", "content_hash", "created_at", "updated_at", "memory_unit_count", "tags"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -97,7 +98,8 @@ class DocumentResponse(BaseModel):
|
||||
"content_hash": obj.get("content_hash"),
|
||||
"created_at": obj.get("created_at"),
|
||||
"updated_at": obj.get("updated_at"),
|
||||
"memory_unit_count": obj.get("memory_unit_count")
|
||||
"memory_unit_count": obj.get("memory_unit_count"),
|
||||
"tags": obj.get("tags")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
Hindsight HTTP API
|
||||
|
||||
HTTP API for Hindsight
|
||||
|
||||
The version of the OpenAPI document: 0.1.0
|
||||
Generated by OpenAPI Generator (https://openapi-generator.tech)
|
||||
|
||||
Do not edit the class manually.
|
||||
""" # noqa: E501
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, StrictInt
|
||||
from typing import Any, ClassVar, Dict, List
|
||||
from hindsight_client_api.models.tag_item import TagItem
|
||||
from typing import Optional, Set
|
||||
from typing_extensions import Self
|
||||
|
||||
class ListTagsResponse(BaseModel):
|
||||
"""
|
||||
Response model for list tags endpoint.
|
||||
""" # noqa: E501
|
||||
items: List[TagItem]
|
||||
total: StrictInt
|
||||
limit: StrictInt
|
||||
offset: StrictInt
|
||||
__properties: ClassVar[List[str]] = ["items", "total", "limit", "offset"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
validate_assignment=True,
|
||||
protected_namespaces=(),
|
||||
)
|
||||
|
||||
|
||||
def to_str(self) -> str:
|
||||
"""Returns the string representation of the model using alias"""
|
||||
return pprint.pformat(self.model_dump(by_alias=True))
|
||||
|
||||
def to_json(self) -> str:
|
||||
"""Returns the JSON representation of the model using alias"""
|
||||
# TODO: pydantic v2: use .model_dump_json(by_alias=True, exclude_unset=True) instead
|
||||
return json.dumps(self.to_dict())
|
||||
|
||||
@classmethod
|
||||
def from_json(cls, json_str: str) -> Optional[Self]:
|
||||
"""Create an instance of ListTagsResponse from a JSON string"""
|
||||
return cls.from_dict(json.loads(json_str))
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Return the dictionary representation of the model using alias.
|
||||
|
||||
This has the following differences from calling pydantic's
|
||||
`self.model_dump(by_alias=True)`:
|
||||
|
||||
* `None` is only added to the output dict for nullable fields that
|
||||
were set at model initialization. Other fields with value `None`
|
||||
are ignored.
|
||||
"""
|
||||
excluded_fields: Set[str] = set([
|
||||
])
|
||||
|
||||
_dict = self.model_dump(
|
||||
by_alias=True,
|
||||
exclude=excluded_fields,
|
||||
exclude_none=True,
|
||||
)
|
||||
# override the default output from pydantic by calling `to_dict()` of each item in items (list)
|
||||
_items = []
|
||||
if self.items:
|
||||
for _item_items in self.items:
|
||||
if _item_items:
|
||||
_items.append(_item_items.to_dict())
|
||||
_dict['items'] = _items
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional[Self]:
|
||||
"""Create an instance of ListTagsResponse from a dict"""
|
||||
if obj is None:
|
||||
return None
|
||||
|
||||
if not isinstance(obj, dict):
|
||||
return cls.model_validate(obj)
|
||||
|
||||
_obj = cls.model_validate({
|
||||
"items": [TagItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None,
|
||||
"total": obj.get("total"),
|
||||
"limit": obj.get("limit"),
|
||||
"offset": obj.get("offset")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -34,7 +34,8 @@ class MemoryItem(BaseModel):
|
||||
metadata: Optional[Dict[str, StrictStr]] = None
|
||||
document_id: Optional[StrictStr] = None
|
||||
entities: Optional[List[EntityInput]] = None
|
||||
__properties: ClassVar[List[str]] = ["content", "timestamp", "context", "metadata", "document_id", "entities"]
|
||||
tags: Optional[List[StrictStr]] = None
|
||||
__properties: ClassVar[List[str]] = ["content", "timestamp", "context", "metadata", "document_id", "entities", "tags"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -107,6 +108,11 @@ class MemoryItem(BaseModel):
|
||||
if self.entities is None and "entities" in self.model_fields_set:
|
||||
_dict['entities'] = None
|
||||
|
||||
# set to None if tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.tags is None and "tags" in self.model_fields_set:
|
||||
_dict['tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -124,7 +130,8 @@ class MemoryItem(BaseModel):
|
||||
"context": obj.get("context"),
|
||||
"metadata": obj.get("metadata"),
|
||||
"document_id": obj.get("document_id"),
|
||||
"entities": [EntityInput.from_dict(_item) for _item in obj["entities"]] if obj.get("entities") is not None else None
|
||||
"entities": [EntityInput.from_dict(_item) for _item in obj["entities"]] if obj.get("entities") is not None else None,
|
||||
"tags": obj.get("tags")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictInt, StrictStr
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictInt, StrictStr, field_validator
|
||||
from typing import Any, ClassVar, Dict, List, Optional
|
||||
from hindsight_client_api.models.budget import Budget
|
||||
from hindsight_client_api.models.include_options import IncludeOptions
|
||||
@@ -35,7 +35,19 @@ class RecallRequest(BaseModel):
|
||||
trace: Optional[StrictBool] = False
|
||||
query_timestamp: Optional[StrictStr] = None
|
||||
include: Optional[IncludeOptions] = Field(default=None, description="Options for including additional data (entities are included by default)")
|
||||
__properties: ClassVar[List[str]] = ["query", "types", "budget", "max_tokens", "trace", "query_timestamp", "include"]
|
||||
tags: Optional[List[StrictStr]] = None
|
||||
tags_match: Optional[StrictStr] = Field(default='any', description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).")
|
||||
__properties: ClassVar[List[str]] = ["query", "types", "budget", "max_tokens", "trace", "query_timestamp", "include", "tags", "tags_match"]
|
||||
|
||||
@field_validator('tags_match')
|
||||
def tags_match_validate_enum(cls, value):
|
||||
"""Validates the enum"""
|
||||
if value is None:
|
||||
return value
|
||||
|
||||
if value not in set(['any', 'all', 'any_strict', 'all_strict']):
|
||||
raise ValueError("must be one of enum values ('any', 'all', 'any_strict', 'all_strict')")
|
||||
return value
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -89,6 +101,11 @@ class RecallRequest(BaseModel):
|
||||
if self.query_timestamp is None and "query_timestamp" in self.model_fields_set:
|
||||
_dict['query_timestamp'] = None
|
||||
|
||||
# set to None if tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.tags is None and "tags" in self.model_fields_set:
|
||||
_dict['tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -107,7 +124,9 @@ class RecallRequest(BaseModel):
|
||||
"max_tokens": obj.get("max_tokens") if obj.get("max_tokens") is not None else 4096,
|
||||
"trace": obj.get("trace") if obj.get("trace") is not None else False,
|
||||
"query_timestamp": obj.get("query_timestamp"),
|
||||
"include": IncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None
|
||||
"include": IncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None,
|
||||
"tags": obj.get("tags"),
|
||||
"tags_match": obj.get("tags_match") if obj.get("tags_match") is not None else 'any'
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -37,7 +37,8 @@ class RecallResult(BaseModel):
|
||||
document_id: Optional[StrictStr] = None
|
||||
metadata: Optional[Dict[str, StrictStr]] = None
|
||||
chunk_id: Optional[StrictStr] = None
|
||||
__properties: ClassVar[List[str]] = ["id", "text", "type", "entities", "context", "occurred_start", "occurred_end", "mentioned_at", "document_id", "metadata", "chunk_id"]
|
||||
tags: Optional[List[StrictStr]] = None
|
||||
__properties: ClassVar[List[str]] = ["id", "text", "type", "entities", "context", "occurred_start", "occurred_end", "mentioned_at", "document_id", "metadata", "chunk_id", "tags"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -123,6 +124,11 @@ class RecallResult(BaseModel):
|
||||
if self.chunk_id is None and "chunk_id" in self.model_fields_set:
|
||||
_dict['chunk_id'] = None
|
||||
|
||||
# set to None if tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.tags is None and "tags" in self.model_fields_set:
|
||||
_dict['tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -145,7 +151,8 @@ class RecallResult(BaseModel):
|
||||
"mentioned_at": obj.get("mentioned_at"),
|
||||
"document_id": obj.get("document_id"),
|
||||
"metadata": obj.get("metadata"),
|
||||
"chunk_id": obj.get("chunk_id")
|
||||
"chunk_id": obj.get("chunk_id"),
|
||||
"tags": obj.get("tags")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr, field_validator
|
||||
from typing import Any, ClassVar, Dict, List, Optional
|
||||
from hindsight_client_api.models.budget import Budget
|
||||
from hindsight_client_api.models.reflect_include_options import ReflectIncludeOptions
|
||||
@@ -34,7 +34,19 @@ class ReflectRequest(BaseModel):
|
||||
max_tokens: Optional[StrictInt] = Field(default=4096, description="Maximum tokens for the response")
|
||||
include: Optional[ReflectIncludeOptions] = Field(default=None, description="Options for including additional data (disabled by default)")
|
||||
response_schema: Optional[Dict[str, Any]] = None
|
||||
__properties: ClassVar[List[str]] = ["query", "budget", "context", "max_tokens", "include", "response_schema"]
|
||||
tags: Optional[List[StrictStr]] = None
|
||||
tags_match: Optional[StrictStr] = Field(default='any', description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).")
|
||||
__properties: ClassVar[List[str]] = ["query", "budget", "context", "max_tokens", "include", "response_schema", "tags", "tags_match"]
|
||||
|
||||
@field_validator('tags_match')
|
||||
def tags_match_validate_enum(cls, value):
|
||||
"""Validates the enum"""
|
||||
if value is None:
|
||||
return value
|
||||
|
||||
if value not in set(['any', 'all', 'any_strict', 'all_strict']):
|
||||
raise ValueError("must be one of enum values ('any', 'all', 'any_strict', 'all_strict')")
|
||||
return value
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -88,6 +100,11 @@ class ReflectRequest(BaseModel):
|
||||
if self.response_schema is None and "response_schema" in self.model_fields_set:
|
||||
_dict['response_schema'] = None
|
||||
|
||||
# set to None if tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.tags is None and "tags" in self.model_fields_set:
|
||||
_dict['tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -105,7 +122,9 @@ class ReflectRequest(BaseModel):
|
||||
"context": obj.get("context"),
|
||||
"max_tokens": obj.get("max_tokens") if obj.get("max_tokens") is not None else 4096,
|
||||
"include": ReflectIncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None,
|
||||
"response_schema": obj.get("response_schema")
|
||||
"response_schema": obj.get("response_schema"),
|
||||
"tags": obj.get("tags"),
|
||||
"tags_match": obj.get("tags_match") if obj.get("tags_match") is not None else 'any'
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictBool
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictStr
|
||||
from typing import Any, ClassVar, Dict, List, Optional
|
||||
from hindsight_client_api.models.memory_item import MemoryItem
|
||||
from typing import Optional, Set
|
||||
@@ -29,7 +29,8 @@ class RetainRequest(BaseModel):
|
||||
""" # noqa: E501
|
||||
items: List[MemoryItem]
|
||||
var_async: Optional[StrictBool] = Field(default=False, description="If true, process asynchronously in background. If false, wait for completion (default: false)", alias="async")
|
||||
__properties: ClassVar[List[str]] = ["items", "async"]
|
||||
document_tags: Optional[List[StrictStr]] = None
|
||||
__properties: ClassVar[List[str]] = ["items", "async", "document_tags"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -77,6 +78,11 @@ class RetainRequest(BaseModel):
|
||||
if _item_items:
|
||||
_items.append(_item_items.to_dict())
|
||||
_dict['items'] = _items
|
||||
# set to None if document_tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.document_tags is None and "document_tags" in self.model_fields_set:
|
||||
_dict['document_tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -90,7 +96,8 @@ class RetainRequest(BaseModel):
|
||||
|
||||
_obj = cls.model_validate({
|
||||
"items": [MemoryItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None,
|
||||
"async": obj.get("async") if obj.get("async") is not None else False
|
||||
"async": obj.get("async") if obj.get("async") is not None else False,
|
||||
"document_tags": obj.get("document_tags")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
Hindsight HTTP API
|
||||
|
||||
HTTP API for Hindsight
|
||||
|
||||
The version of the OpenAPI document: 0.1.0
|
||||
Generated by OpenAPI Generator (https://openapi-generator.tech)
|
||||
|
||||
Do not edit the class manually.
|
||||
""" # noqa: E501
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
|
||||
from typing import Any, ClassVar, Dict, List
|
||||
from typing import Optional, Set
|
||||
from typing_extensions import Self
|
||||
|
||||
class TagItem(BaseModel):
|
||||
"""
|
||||
Single tag with usage count.
|
||||
""" # noqa: E501
|
||||
tag: StrictStr = Field(description="The tag value")
|
||||
count: StrictInt = Field(description="Number of memories with this tag")
|
||||
__properties: ClassVar[List[str]] = ["tag", "count"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
validate_assignment=True,
|
||||
protected_namespaces=(),
|
||||
)
|
||||
|
||||
|
||||
def to_str(self) -> str:
|
||||
"""Returns the string representation of the model using alias"""
|
||||
return pprint.pformat(self.model_dump(by_alias=True))
|
||||
|
||||
def to_json(self) -> str:
|
||||
"""Returns the JSON representation of the model using alias"""
|
||||
# TODO: pydantic v2: use .model_dump_json(by_alias=True, exclude_unset=True) instead
|
||||
return json.dumps(self.to_dict())
|
||||
|
||||
@classmethod
|
||||
def from_json(cls, json_str: str) -> Optional[Self]:
|
||||
"""Create an instance of TagItem from a JSON string"""
|
||||
return cls.from_dict(json.loads(json_str))
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Return the dictionary representation of the model using alias.
|
||||
|
||||
This has the following differences from calling pydantic's
|
||||
`self.model_dump(by_alias=True)`:
|
||||
|
||||
* `None` is only added to the output dict for nullable fields that
|
||||
were set at model initialization. Other fields with value `None`
|
||||
are ignored.
|
||||
"""
|
||||
excluded_fields: Set[str] = set([
|
||||
])
|
||||
|
||||
_dict = self.model_dump(
|
||||
by_alias=True,
|
||||
exclude=excluded_fields,
|
||||
exclude_none=True,
|
||||
)
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional[Self]:
|
||||
"""Create an instance of TagItem from a dict"""
|
||||
if obj is None:
|
||||
return None
|
||||
|
||||
if not isinstance(obj, dict):
|
||||
return cls.model_validate(obj)
|
||||
|
||||
_obj = cls.model_validate({
|
||||
"tag": obj.get("tag"),
|
||||
"count": obj.get("count")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -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"}
|
||||
|
||||
@@ -64,6 +64,7 @@ mod tests {
|
||||
metadata: None,
|
||||
timestamp: None,
|
||||
entities: None,
|
||||
tags: None,
|
||||
},
|
||||
types::MemoryItem {
|
||||
content: "Bob works with Alice on the search team".to_string(),
|
||||
@@ -72,8 +73,10 @@ mod tests {
|
||||
metadata: None,
|
||||
timestamp: None,
|
||||
entities: None,
|
||||
tags: None,
|
||||
},
|
||||
],
|
||||
document_tags: None,
|
||||
};
|
||||
let retain_response = client
|
||||
.retain_memories(&bank_id, None, &retain_request)
|
||||
@@ -90,6 +93,8 @@ mod tests {
|
||||
include: None,
|
||||
query_timestamp: None,
|
||||
types: None,
|
||||
tags: None,
|
||||
tags_match: types::TagsMatch::Any,
|
||||
};
|
||||
let recall_response = client
|
||||
.recall_memories(&bank_id, None, &recall_request)
|
||||
@@ -106,6 +111,8 @@ mod tests {
|
||||
max_tokens: 4096,
|
||||
include: None,
|
||||
response_schema: None,
|
||||
tags: None,
|
||||
tags_match: types::TagsMatch::Any,
|
||||
};
|
||||
let reflect_response = client
|
||||
.reflect(&bank_id, None, &reflect_request)
|
||||
|
||||
@@ -39,6 +39,9 @@ import type {
|
||||
GetGraphData,
|
||||
GetGraphErrors,
|
||||
GetGraphResponses,
|
||||
GetMemoryData,
|
||||
GetMemoryErrors,
|
||||
GetMemoryResponses,
|
||||
HealthEndpointHealthGetData,
|
||||
HealthEndpointHealthGetResponses,
|
||||
ListBanksData,
|
||||
@@ -56,6 +59,9 @@ import type {
|
||||
ListOperationsData,
|
||||
ListOperationsErrors,
|
||||
ListOperationsResponses,
|
||||
ListTagsData,
|
||||
ListTagsErrors,
|
||||
ListTagsResponses,
|
||||
MetricsEndpointMetricsGetData,
|
||||
MetricsEndpointMetricsGetResponses,
|
||||
RecallMemoriesData,
|
||||
@@ -148,6 +154,20 @@ export const listMemories = <ThrowOnError extends boolean = false>(
|
||||
ThrowOnError
|
||||
>({ url: "/v1/default/banks/{bank_id}/memories/list", ...options });
|
||||
|
||||
/**
|
||||
* Get memory unit
|
||||
*
|
||||
* Get a single memory unit by ID with all its metadata including entities and tags.
|
||||
*/
|
||||
export const getMemory = <ThrowOnError extends boolean = false>(
|
||||
options: Options<GetMemoryData, ThrowOnError>,
|
||||
) =>
|
||||
(options.client ?? client).get<
|
||||
GetMemoryResponses,
|
||||
GetMemoryErrors,
|
||||
ThrowOnError
|
||||
>({ url: "/v1/default/banks/{bank_id}/memories/{memory_id}", ...options });
|
||||
|
||||
/**
|
||||
* Recall memory
|
||||
*
|
||||
@@ -329,6 +349,20 @@ export const getDocument = <ThrowOnError extends boolean = false>(
|
||||
ThrowOnError
|
||||
>({ url: "/v1/default/banks/{bank_id}/documents/{document_id}", ...options });
|
||||
|
||||
/**
|
||||
* List tags
|
||||
*
|
||||
* List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
|
||||
*/
|
||||
export const listTags = <ThrowOnError extends boolean = false>(
|
||||
options: Options<ListTagsData, ThrowOnError>,
|
||||
) =>
|
||||
(options.client ?? client).get<
|
||||
ListTagsResponses,
|
||||
ListTagsErrors,
|
||||
ThrowOnError
|
||||
>({ url: "/v1/default/banks/{bank_id}/tags", ...options });
|
||||
|
||||
/**
|
||||
* Get chunk details
|
||||
*
|
||||
|
||||
@@ -377,6 +377,12 @@ export type DocumentResponse = {
|
||||
* Memory Unit Count
|
||||
*/
|
||||
memory_unit_count: number;
|
||||
/**
|
||||
* Tags
|
||||
*
|
||||
* Tags associated with this document
|
||||
*/
|
||||
tags?: Array<string>;
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -666,6 +672,30 @@ export type ListMemoryUnitsResponse = {
|
||||
offset: number;
|
||||
};
|
||||
|
||||
/**
|
||||
* ListTagsResponse
|
||||
*
|
||||
* Response model for list tags endpoint.
|
||||
*/
|
||||
export type ListTagsResponse = {
|
||||
/**
|
||||
* Items
|
||||
*/
|
||||
items: Array<TagItem>;
|
||||
/**
|
||||
* Total
|
||||
*/
|
||||
total: number;
|
||||
/**
|
||||
* Limit
|
||||
*/
|
||||
limit: number;
|
||||
/**
|
||||
* Offset
|
||||
*/
|
||||
offset: number;
|
||||
};
|
||||
|
||||
/**
|
||||
* MemoryItem
|
||||
*
|
||||
@@ -702,6 +732,12 @@ export type MemoryItem = {
|
||||
* Optional entities to combine with auto-extracted entities.
|
||||
*/
|
||||
entities?: Array<EntityInput> | null;
|
||||
/**
|
||||
* Tags
|
||||
*
|
||||
* Optional tags for visibility scoping. Memories with tags can be filtered during recall.
|
||||
*/
|
||||
tags?: Array<string> | null;
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -791,6 +827,18 @@ export type RecallRequest = {
|
||||
* Options for including additional data (entities are included by default)
|
||||
*/
|
||||
include?: IncludeOptions;
|
||||
/**
|
||||
* Tags
|
||||
*
|
||||
* Filter memories by tags. If not specified, all memories are returned.
|
||||
*/
|
||||
tags?: Array<string> | null;
|
||||
/**
|
||||
* Tags Match
|
||||
*
|
||||
* How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).
|
||||
*/
|
||||
tags_match?: "any" | "all" | "any_strict" | "all_strict";
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -879,6 +927,10 @@ export type RecallResult = {
|
||||
* Chunk Id
|
||||
*/
|
||||
chunk_id?: string | null;
|
||||
/**
|
||||
* Tags
|
||||
*/
|
||||
tags?: Array<string> | null;
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -958,6 +1010,18 @@ export type ReflectRequest = {
|
||||
response_schema?: {
|
||||
[key: string]: unknown;
|
||||
} | null;
|
||||
/**
|
||||
* Tags
|
||||
*
|
||||
* Filter memories by tags during reflection. If not specified, all memories are considered.
|
||||
*/
|
||||
tags?: Array<string> | null;
|
||||
/**
|
||||
* Tags Match
|
||||
*
|
||||
* How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).
|
||||
*/
|
||||
tags_match?: "any" | "all" | "any_strict" | "all_strict";
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -1004,6 +1068,12 @@ export type RetainRequest = {
|
||||
* If true, process asynchronously in background. If false, wait for completion (default: false)
|
||||
*/
|
||||
async?: boolean;
|
||||
/**
|
||||
* Document Tags
|
||||
*
|
||||
* Tags applied to all items in this request. These are merged with any item-level tags.
|
||||
*/
|
||||
document_tags?: Array<string> | null;
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -1042,6 +1112,26 @@ export type RetainResponse = {
|
||||
usage?: TokenUsage | null;
|
||||
};
|
||||
|
||||
/**
|
||||
* TagItem
|
||||
*
|
||||
* Single tag with usage count.
|
||||
*/
|
||||
export type TagItem = {
|
||||
/**
|
||||
* Tag
|
||||
*
|
||||
* The tag value
|
||||
*/
|
||||
tag: string;
|
||||
/**
|
||||
* Count
|
||||
*
|
||||
* Number of memories with this tag
|
||||
*/
|
||||
count: number;
|
||||
};
|
||||
|
||||
/**
|
||||
* TokenUsage
|
||||
*
|
||||
@@ -1225,6 +1315,44 @@ export type ListMemoriesResponses = {
|
||||
export type ListMemoriesResponse =
|
||||
ListMemoriesResponses[keyof ListMemoriesResponses];
|
||||
|
||||
export type GetMemoryData = {
|
||||
body?: never;
|
||||
headers?: {
|
||||
/**
|
||||
* Authorization
|
||||
*/
|
||||
authorization?: string | null;
|
||||
};
|
||||
path: {
|
||||
/**
|
||||
* Bank Id
|
||||
*/
|
||||
bank_id: string;
|
||||
/**
|
||||
* Memory Id
|
||||
*/
|
||||
memory_id: string;
|
||||
};
|
||||
query?: never;
|
||||
url: "/v1/default/banks/{bank_id}/memories/{memory_id}";
|
||||
};
|
||||
|
||||
export type GetMemoryErrors = {
|
||||
/**
|
||||
* Validation Error
|
||||
*/
|
||||
422: HttpValidationError;
|
||||
};
|
||||
|
||||
export type GetMemoryError = GetMemoryErrors[keyof GetMemoryErrors];
|
||||
|
||||
export type GetMemoryResponses = {
|
||||
/**
|
||||
* Successful Response
|
||||
*/
|
||||
200: unknown;
|
||||
};
|
||||
|
||||
export type RecallMemoriesData = {
|
||||
body: RecallRequest;
|
||||
headers?: {
|
||||
@@ -1632,6 +1760,61 @@ export type GetDocumentResponses = {
|
||||
export type GetDocumentResponse =
|
||||
GetDocumentResponses[keyof GetDocumentResponses];
|
||||
|
||||
export type ListTagsData = {
|
||||
body?: never;
|
||||
headers?: {
|
||||
/**
|
||||
* Authorization
|
||||
*/
|
||||
authorization?: string | null;
|
||||
};
|
||||
path: {
|
||||
/**
|
||||
* Bank Id
|
||||
*/
|
||||
bank_id: string;
|
||||
};
|
||||
query?: {
|
||||
/**
|
||||
* Q
|
||||
*
|
||||
* Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
|
||||
*/
|
||||
q?: string | null;
|
||||
/**
|
||||
* Limit
|
||||
*
|
||||
* Maximum number of tags to return
|
||||
*/
|
||||
limit?: number;
|
||||
/**
|
||||
* Offset
|
||||
*
|
||||
* Offset for pagination
|
||||
*/
|
||||
offset?: number;
|
||||
};
|
||||
url: "/v1/default/banks/{bank_id}/tags";
|
||||
};
|
||||
|
||||
export type ListTagsErrors = {
|
||||
/**
|
||||
* Validation Error
|
||||
*/
|
||||
422: HttpValidationError;
|
||||
};
|
||||
|
||||
export type ListTagsError = ListTagsErrors[keyof ListTagsErrors];
|
||||
|
||||
export type ListTagsResponses = {
|
||||
/**
|
||||
* Successful Response
|
||||
*/
|
||||
200: ListTagsResponse;
|
||||
};
|
||||
|
||||
export type ListTagsResponse2 = ListTagsResponses[keyof ListTagsResponses];
|
||||
|
||||
export type GetChunkData = {
|
||||
body?: never;
|
||||
headers?: {
|
||||
|
||||
@@ -1,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",
|
||||
|
||||
@@ -62,6 +62,7 @@ export interface MemoryItemInput {
|
||||
metadata?: Record<string, string>;
|
||||
document_id?: string;
|
||||
entities?: EntityInput[];
|
||||
tags?: string[];
|
||||
}
|
||||
|
||||
export class HindsightClient {
|
||||
@@ -142,13 +143,14 @@ export class HindsightClient {
|
||||
/**
|
||||
* Retain multiple memories in batch.
|
||||
*/
|
||||
async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; async?: boolean }): Promise<RetainResponse> {
|
||||
async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; documentTags?: string[]; async?: boolean }): Promise<RetainResponse> {
|
||||
const processedItems = items.map((item) => ({
|
||||
content: item.content,
|
||||
context: item.context,
|
||||
metadata: item.metadata,
|
||||
document_id: item.document_id,
|
||||
entities: item.entities,
|
||||
tags: item.tags,
|
||||
timestamp:
|
||||
item.timestamp instanceof Date
|
||||
? item.timestamp.toISOString()
|
||||
@@ -166,6 +168,7 @@ export class HindsightClient {
|
||||
path: { bank_id: bankId },
|
||||
body: {
|
||||
items: itemsWithDocId,
|
||||
document_tags: options?.documentTags,
|
||||
async: options?.async,
|
||||
},
|
||||
});
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@vectorize-io/hindsight-control-plane",
|
||||
"version": "0.2.1",
|
||||
"version": "0.3.0",
|
||||
"description": "Control plane for Hindsight - Semantic memory system",
|
||||
"bin": {
|
||||
"hindsight-control-plane": "./bin/cli.js"
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
import { NextRequest, NextResponse } from "next/server";
|
||||
|
||||
const DATAPLANE_URL = process.env.HINDSIGHT_CP_DATAPLANE_API_URL || "http://localhost:8888";
|
||||
|
||||
export async function GET(
|
||||
request: NextRequest,
|
||||
{ params }: { params: Promise<{ memoryId: string }> }
|
||||
) {
|
||||
try {
|
||||
const { memoryId } = await params;
|
||||
const searchParams = request.nextUrl.searchParams;
|
||||
const bankId = searchParams.get("bank_id");
|
||||
|
||||
if (!bankId) {
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const response = await fetch(
|
||||
`${DATAPLANE_URL}/v1/default/banks/${bankId}/memories/${memoryId}`,
|
||||
{
|
||||
method: "GET",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
}
|
||||
);
|
||||
|
||||
if (!response.ok) {
|
||||
if (response.status === 404) {
|
||||
return NextResponse.json({ error: "Memory not found" }, { status: 404 });
|
||||
}
|
||||
throw new Error(`API returned ${response.status}`);
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
return NextResponse.json(data, { status: 200 });
|
||||
} catch (error) {
|
||||
console.error("Error fetching memory:", error);
|
||||
return NextResponse.json({ error: "Failed to fetch memory" }, { status: 500 });
|
||||
}
|
||||
}
|
||||
@@ -10,9 +10,12 @@ export async function POST(request: NextRequest) {
|
||||
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
|
||||
}
|
||||
|
||||
const { items, document_id } = body;
|
||||
const { items, document_id, document_tags } = body;
|
||||
|
||||
const response = await hindsightClient.retainBatch(bankId, items, { documentId: document_id });
|
||||
const response = await hindsightClient.retainBatch(bankId, items, {
|
||||
documentId: document_id,
|
||||
documentTags: document_tags,
|
||||
});
|
||||
|
||||
return NextResponse.json(response, { status: 200 });
|
||||
} catch (error) {
|
||||
|
||||
@@ -5,7 +5,18 @@ export async function POST(request: NextRequest) {
|
||||
try {
|
||||
const body = await request.json();
|
||||
const bankId = body.bank_id || body.agent_id || "default";
|
||||
const { query, types, fact_type, max_tokens, trace, budget, include, query_timestamp } = body;
|
||||
const {
|
||||
query,
|
||||
types,
|
||||
fact_type,
|
||||
max_tokens,
|
||||
trace,
|
||||
budget,
|
||||
include,
|
||||
query_timestamp,
|
||||
tags,
|
||||
tags_match,
|
||||
} = body;
|
||||
|
||||
const response = await sdk.recallMemories({
|
||||
client: lowLevelClient,
|
||||
@@ -18,6 +29,8 @@ export async function POST(request: NextRequest) {
|
||||
budget: budget || "mid",
|
||||
include,
|
||||
query_timestamp,
|
||||
tags,
|
||||
tags_match,
|
||||
},
|
||||
});
|
||||
|
||||
|
||||
@@ -5,12 +5,14 @@ export async function POST(request: NextRequest) {
|
||||
try {
|
||||
const body = await request.json();
|
||||
const bankId = body.bank_id || body.agent_id || "default";
|
||||
const { query, context, budget, thinking_budget, include_facts } = body;
|
||||
const { query, context, budget, thinking_budget, include_facts, tags, tags_match } = body;
|
||||
|
||||
const requestBody: any = {
|
||||
query,
|
||||
budget: budget || (thinking_budget ? "mid" : "low"),
|
||||
context: context || undefined,
|
||||
tags,
|
||||
tags_match,
|
||||
};
|
||||
|
||||
// Add include options if specified
|
||||
|
||||
@@ -7,6 +7,7 @@ import { Button } from "@/components/ui/button";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
import { Checkbox } from "@/components/ui/checkbox";
|
||||
import { Tag } from "lucide-react";
|
||||
|
||||
export function AddMemoryView() {
|
||||
const { currentBank } = useBank();
|
||||
@@ -14,6 +15,7 @@ export function AddMemoryView() {
|
||||
const [context, setContext] = useState("");
|
||||
const [eventDate, setEventDate] = useState("");
|
||||
const [documentId, setDocumentId] = useState("");
|
||||
const [tags, setTags] = useState("");
|
||||
const [async, setAsync] = useState(false);
|
||||
const [loading, setLoading] = useState(false);
|
||||
const [result, setResult] = useState<string | null>(null);
|
||||
@@ -23,6 +25,7 @@ export function AddMemoryView() {
|
||||
setContext("");
|
||||
setEventDate("");
|
||||
setDocumentId("");
|
||||
setTags("");
|
||||
setAsync(false);
|
||||
setResult(null);
|
||||
};
|
||||
@@ -37,16 +40,24 @@ export function AddMemoryView() {
|
||||
setResult(null);
|
||||
|
||||
try {
|
||||
// Parse tags from comma-separated string
|
||||
const parsedTags = tags
|
||||
.split(",")
|
||||
.map((t) => t.trim())
|
||||
.filter((t) => t.length > 0);
|
||||
|
||||
const item: any = { content };
|
||||
if (context) item.context = context;
|
||||
// datetime-local gives "2024-01-15T10:30", add seconds for proper ISO format
|
||||
if (eventDate) item.timestamp = eventDate + ":00";
|
||||
if (parsedTags.length > 0) item.tags = parsedTags;
|
||||
|
||||
const data: any = await client.retain({
|
||||
bank_id: currentBank,
|
||||
items: [item],
|
||||
document_id: documentId,
|
||||
async,
|
||||
...(parsedTags.length > 0 && { document_tags: parsedTags }),
|
||||
});
|
||||
|
||||
setResult(data.message as string);
|
||||
@@ -112,6 +123,22 @@ export function AddMemoryView() {
|
||||
</small>
|
||||
</div>
|
||||
|
||||
<div className="mb-4">
|
||||
<label className="font-bold block mb-1 text-card-foreground flex items-center gap-2">
|
||||
<Tag className="h-4 w-4" />
|
||||
Tags
|
||||
</label>
|
||||
<Input
|
||||
type="text"
|
||||
value={tags}
|
||||
onChange={(e) => setTags(e.target.value)}
|
||||
placeholder="user_alice, session_123, project_x"
|
||||
/>
|
||||
<small className="text-muted-foreground text-xs mt-1 block">
|
||||
Comma-separated tags for filtering during recall/reflect. Tags cannot contain commas.
|
||||
</small>
|
||||
</div>
|
||||
|
||||
<div className="mb-4">
|
||||
<div className="flex items-center gap-2">
|
||||
<Checkbox
|
||||
|
||||
@@ -23,7 +23,7 @@ import {
|
||||
DialogFooter,
|
||||
} from "@/components/ui/dialog";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { Check, ChevronsUpDown, Plus, FileText, Moon, Sun, Github } from "lucide-react";
|
||||
import { Check, ChevronsUpDown, Plus, FileText, Moon, Sun, Github, Tag } from "lucide-react";
|
||||
import { useTheme } from "@/lib/theme-context";
|
||||
import Image from "next/image";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
@@ -47,6 +47,7 @@ function BankSelectorInner() {
|
||||
const [docContext, setDocContext] = React.useState("");
|
||||
const [docEventDate, setDocEventDate] = React.useState("");
|
||||
const [docDocumentId, setDocDocumentId] = React.useState("");
|
||||
const [docTags, setDocTags] = React.useState("");
|
||||
const [docAsync, setDocAsync] = React.useState(false);
|
||||
const [isCreatingDoc, setIsCreatingDoc] = React.useState(false);
|
||||
const [docError, setDocError] = React.useState<string | null>(null);
|
||||
@@ -83,10 +84,17 @@ function BankSelectorInner() {
|
||||
setDocError(null);
|
||||
|
||||
try {
|
||||
// Parse tags from comma-separated string
|
||||
const parsedTags = docTags
|
||||
.split(",")
|
||||
.map((t) => t.trim())
|
||||
.filter((t) => t.length > 0);
|
||||
|
||||
const item: any = { content: docContent };
|
||||
if (docContext) item.context = docContext;
|
||||
// datetime-local gives "2024-01-15T10:30", add seconds for proper ISO format
|
||||
if (docEventDate) item.timestamp = docEventDate + ":00";
|
||||
if (parsedTags.length > 0) item.tags = parsedTags;
|
||||
|
||||
const params: any = {
|
||||
bank_id: currentBank,
|
||||
@@ -94,6 +102,7 @@ function BankSelectorInner() {
|
||||
};
|
||||
|
||||
if (docDocumentId) params.document_id = docDocumentId;
|
||||
if (parsedTags.length > 0) params.document_tags = parsedTags;
|
||||
|
||||
if (docAsync) {
|
||||
await client.retain({ ...params, async: true });
|
||||
@@ -107,6 +116,7 @@ function BankSelectorInner() {
|
||||
setDocContext("");
|
||||
setDocEventDate("");
|
||||
setDocDocumentId("");
|
||||
setDocTags("");
|
||||
setDocAsync(false);
|
||||
|
||||
// Navigate to documents view to see the new document
|
||||
@@ -335,6 +345,22 @@ function BankSelectorInner() {
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<label className="font-bold block mb-1 text-sm text-foreground flex items-center gap-2">
|
||||
<Tag className="h-4 w-4" />
|
||||
Tags
|
||||
</label>
|
||||
<Input
|
||||
type="text"
|
||||
value={docTags}
|
||||
onChange={(e) => setDocTags(e.target.value)}
|
||||
placeholder="user_alice, session_123, project_x"
|
||||
/>
|
||||
<p className="text-xs text-muted-foreground mt-1">
|
||||
Comma-separated tags for filtering during recall/reflect
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div className="flex items-center gap-2">
|
||||
<Checkbox
|
||||
id="async-doc"
|
||||
@@ -357,6 +383,7 @@ function BankSelectorInner() {
|
||||
setDocContext("");
|
||||
setDocEventDate("");
|
||||
setDocDocumentId("");
|
||||
setDocTags("");
|
||||
setDocAsync(false);
|
||||
setDocError(null);
|
||||
}}
|
||||
|
||||
@@ -362,6 +362,7 @@ export function DataView({ factType }: DataViewProps) {
|
||||
memory={selectedGraphNode}
|
||||
onClose={() => setSelectedGraphNode(null)}
|
||||
inPanel
|
||||
bankId={currentBank || undefined}
|
||||
/>
|
||||
) : (
|
||||
/* Legend & Controls View */
|
||||
@@ -738,13 +739,20 @@ export function DataView({ factType }: DataViewProps) {
|
||||
memory={selectedTableMemory}
|
||||
onClose={() => setSelectedTableMemory(null)}
|
||||
inPanel
|
||||
bankId={currentBank || undefined}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{viewMode === "timeline" && <TimelineView data={data} filteredRows={filteredTableRows} />}
|
||||
{viewMode === "timeline" && (
|
||||
<TimelineView
|
||||
data={data}
|
||||
filteredRows={filteredTableRows}
|
||||
bankId={currentBank || undefined}
|
||||
/>
|
||||
)}
|
||||
</>
|
||||
) : (
|
||||
<div className="flex items-center justify-center py-20">
|
||||
@@ -761,7 +769,15 @@ export function DataView({ factType }: DataViewProps) {
|
||||
// Timeline View Component - Custom compact timeline with zoom and navigation
|
||||
type Granularity = "year" | "month" | "week" | "day";
|
||||
|
||||
function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }) {
|
||||
function TimelineView({
|
||||
data,
|
||||
filteredRows,
|
||||
bankId,
|
||||
}: {
|
||||
data: any;
|
||||
filteredRows: any[];
|
||||
bankId?: string;
|
||||
}) {
|
||||
const [selectedItem, setSelectedItem] = useState<any>(null);
|
||||
const [granularity, setGranularity] = useState<Granularity>("month");
|
||||
const [currentIndex, setCurrentIndex] = useState(0);
|
||||
@@ -1114,7 +1130,12 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
|
||||
{/* Detail Panel - Fixed on Right */}
|
||||
{selectedItem && (
|
||||
<div className="fixed right-0 top-0 h-screen w-[420px] bg-card border-l-2 border-primary shadow-2xl z-50 overflow-y-auto animate-in slide-in-from-right duration-300 ease-out">
|
||||
<MemoryDetailPanel memory={selectedItem} onClose={() => setSelectedItem(null)} inPanel />
|
||||
<MemoryDetailPanel
|
||||
memory={selectedItem}
|
||||
onClose={() => setSelectedItem(null)}
|
||||
inPanel
|
||||
bankId={bankId}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
@@ -318,6 +318,25 @@ export function DocumentsView() {
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Tags */}
|
||||
{selectedDocument.tags && selectedDocument.tags.length > 0 && (
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Tags
|
||||
</div>
|
||||
<div className="flex flex-wrap gap-2">
|
||||
{selectedDocument.tags.map((tag: string, i: number) => (
|
||||
<span
|
||||
key={i}
|
||||
className="text-sm px-3 py-1.5 rounded-full bg-amber-500/10 text-amber-600 dark:text-amber-400 font-medium"
|
||||
>
|
||||
{tag}
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Delete Button */}
|
||||
<div className="pt-2 border-t border-border">
|
||||
<Button
|
||||
|
||||
@@ -1,15 +1,17 @@
|
||||
"use client";
|
||||
|
||||
import { useState } from "react";
|
||||
import { useState, useEffect } from "react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Copy, Check, X } from "lucide-react";
|
||||
import { Copy, Check, X, Loader2 } from "lucide-react";
|
||||
import { DocumentChunkModal } from "./document-chunk-modal";
|
||||
import { client } from "@/lib/api";
|
||||
|
||||
interface MemoryDetailPanelProps {
|
||||
memory: any;
|
||||
onClose: () => void;
|
||||
compact?: boolean;
|
||||
inPanel?: boolean;
|
||||
bankId?: string;
|
||||
}
|
||||
|
||||
export function MemoryDetailPanel({
|
||||
@@ -17,10 +19,40 @@ export function MemoryDetailPanel({
|
||||
onClose,
|
||||
compact = false,
|
||||
inPanel = false,
|
||||
bankId,
|
||||
}: MemoryDetailPanelProps) {
|
||||
const [copiedId, setCopiedId] = useState<string | null>(null);
|
||||
const [modalType, setModalType] = useState<"document" | "chunk" | null>(null);
|
||||
const [modalId, setModalId] = useState<string | null>(null);
|
||||
const [fullMemory, setFullMemory] = useState<any>(null);
|
||||
const [loading, setLoading] = useState(false);
|
||||
|
||||
// Fetch full memory data when panel opens
|
||||
useEffect(() => {
|
||||
const memoryId = memory?.id || memory?.node_id;
|
||||
if (!memoryId || !bankId) {
|
||||
setFullMemory(null);
|
||||
return;
|
||||
}
|
||||
|
||||
setLoading(true);
|
||||
client
|
||||
.getMemory(memoryId, bankId)
|
||||
.then((data) => {
|
||||
setFullMemory(data);
|
||||
})
|
||||
.catch((err) => {
|
||||
console.error("Failed to fetch memory details:", err);
|
||||
// Fall back to showing the partial data we have
|
||||
setFullMemory(null);
|
||||
})
|
||||
.finally(() => {
|
||||
setLoading(false);
|
||||
});
|
||||
}, [memory?.id, memory?.node_id, bankId]);
|
||||
|
||||
// Use full memory data if available, otherwise fall back to the partial data passed in
|
||||
const displayMemory = fullMemory || memory;
|
||||
|
||||
const copyToClipboard = async (text: string) => {
|
||||
try {
|
||||
@@ -50,7 +82,7 @@ export function MemoryDetailPanel({
|
||||
if (!memory) return null;
|
||||
|
||||
// Handle both 'id' and 'node_id' (trace results use node_id)
|
||||
const memoryId = memory.id || memory.node_id;
|
||||
const memoryId = displayMemory.id || displayMemory.node_id;
|
||||
|
||||
const labelSize = compact ? "text-[10px]" : "text-xs";
|
||||
const textSize = compact ? "text-xs" : "text-sm";
|
||||
@@ -71,123 +103,156 @@ export function MemoryDetailPanel({
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
<div className="space-y-5">
|
||||
{/* Full Text */}
|
||||
<div>
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Full Text
|
||||
</div>
|
||||
<div className="text-sm whitespace-pre-wrap leading-relaxed text-foreground">
|
||||
{memory.text}
|
||||
</div>
|
||||
{loading ? (
|
||||
<div className="flex items-center justify-center py-12">
|
||||
<Loader2 className="h-6 w-6 animate-spin text-muted-foreground" />
|
||||
<span className="ml-2 text-muted-foreground">Loading memory details...</span>
|
||||
</div>
|
||||
|
||||
{/* Context */}
|
||||
{memory.context && (
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Context
|
||||
</div>
|
||||
<div className="text-sm text-foreground">{memory.context}</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Dates */}
|
||||
<div className="grid grid-cols-2 gap-4">
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Occurred
|
||||
</div>
|
||||
<div className="text-sm font-medium text-foreground">
|
||||
{memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"}
|
||||
</div>
|
||||
</div>
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Mentioned
|
||||
</div>
|
||||
<div className="text-sm font-medium text-foreground">
|
||||
{memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Entities */}
|
||||
{memory.entities && (
|
||||
) : (
|
||||
<div className="space-y-5">
|
||||
{/* Full Text */}
|
||||
<div>
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">
|
||||
Entities
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Full Text
|
||||
</div>
|
||||
<div className="flex flex-wrap gap-2">
|
||||
{(Array.isArray(memory.entities)
|
||||
? memory.entities
|
||||
: String(memory.entities).split(", ")
|
||||
).map((entity: any, i: number) => {
|
||||
const entityText =
|
||||
typeof entity === "string" ? entity : entity?.name || JSON.stringify(entity);
|
||||
return (
|
||||
<div className="text-sm whitespace-pre-wrap leading-relaxed text-foreground">
|
||||
{displayMemory.text}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Context */}
|
||||
{displayMemory.context && (
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Context
|
||||
</div>
|
||||
<div className="text-sm text-foreground">{displayMemory.context}</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Dates */}
|
||||
<div className="grid grid-cols-2 gap-4">
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Occurred
|
||||
</div>
|
||||
<div className="text-sm font-medium text-foreground">
|
||||
{displayMemory.occurred_start
|
||||
? new Date(displayMemory.occurred_start).toLocaleString()
|
||||
: "N/A"}
|
||||
</div>
|
||||
</div>
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Mentioned
|
||||
</div>
|
||||
<div className="text-sm font-medium text-foreground">
|
||||
{displayMemory.mentioned_at
|
||||
? new Date(displayMemory.mentioned_at).toLocaleString()
|
||||
: "N/A"}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Entities */}
|
||||
{displayMemory.entities &&
|
||||
(Array.isArray(displayMemory.entities)
|
||||
? displayMemory.entities.length > 0
|
||||
: displayMemory.entities) && (
|
||||
<div>
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">
|
||||
Entities
|
||||
</div>
|
||||
<div className="flex flex-wrap gap-2">
|
||||
{(Array.isArray(displayMemory.entities)
|
||||
? displayMemory.entities
|
||||
: String(displayMemory.entities).split(", ")
|
||||
).map((entity: any, i: number) => {
|
||||
const entityText =
|
||||
typeof entity === "string"
|
||||
? entity
|
||||
: entity?.name || JSON.stringify(entity);
|
||||
return (
|
||||
<span
|
||||
key={i}
|
||||
className="text-sm px-3 py-1.5 rounded-full bg-primary/10 text-primary font-medium"
|
||||
>
|
||||
{entityText}
|
||||
</span>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Tags */}
|
||||
{displayMemory.tags && displayMemory.tags.length > 0 && (
|
||||
<div>
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">Tags</div>
|
||||
<div className="flex flex-wrap gap-2">
|
||||
{displayMemory.tags.map((tag: string, i: number) => (
|
||||
<span
|
||||
key={i}
|
||||
className="text-sm px-3 py-1.5 rounded-full bg-primary/10 text-primary font-medium"
|
||||
className="text-sm px-3 py-1.5 rounded-full bg-amber-500/10 text-amber-600 dark:text-amber-400 font-medium"
|
||||
>
|
||||
{entityText}
|
||||
{tag}
|
||||
</span>
|
||||
);
|
||||
})}
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
)}
|
||||
|
||||
{/* ID */}
|
||||
{memoryId && (
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Memory ID
|
||||
{/* ID */}
|
||||
{memoryId && (
|
||||
<div className="p-4 bg-muted/50 rounded-lg">
|
||||
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
|
||||
Memory ID
|
||||
</div>
|
||||
<div className="flex items-center gap-2">
|
||||
<code className="text-xs font-mono break-all flex-1 text-muted-foreground">
|
||||
{memoryId}
|
||||
</code>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
className="h-8 w-8 p-0 flex-shrink-0"
|
||||
onClick={() => copyToClipboard(memoryId)}
|
||||
>
|
||||
{copiedId === memoryId ? (
|
||||
<Check className="h-4 w-4 text-green-600" />
|
||||
) : (
|
||||
<Copy className="h-4 w-4" />
|
||||
)}
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
<div className="flex items-center gap-2">
|
||||
<code className="text-xs font-mono break-all flex-1 text-muted-foreground">
|
||||
{memoryId}
|
||||
</code>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
className="h-8 w-8 p-0 flex-shrink-0"
|
||||
onClick={() => copyToClipboard(memoryId)}
|
||||
>
|
||||
{copiedId === memoryId ? (
|
||||
<Check className="h-4 w-4 text-green-600" />
|
||||
) : (
|
||||
<Copy className="h-4 w-4" />
|
||||
)}
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
)}
|
||||
|
||||
{/* Document/Chunk buttons */}
|
||||
{(memory.document_id || memory.chunk_id) && (
|
||||
<div className="flex gap-3 pt-2">
|
||||
{memory.document_id && (
|
||||
<Button
|
||||
onClick={() => openDocumentModal(memory.document_id)}
|
||||
variant="secondary"
|
||||
className="flex-1"
|
||||
>
|
||||
View Document
|
||||
</Button>
|
||||
)}
|
||||
{memory.chunk_id && (
|
||||
<Button
|
||||
onClick={() => openChunkModal(memory.chunk_id)}
|
||||
variant="secondary"
|
||||
className="flex-1"
|
||||
>
|
||||
View Chunk
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
{/* Document/Chunk buttons */}
|
||||
{(displayMemory.document_id || displayMemory.chunk_id) && (
|
||||
<div className="flex gap-3 pt-2">
|
||||
{displayMemory.document_id && (
|
||||
<Button
|
||||
onClick={() => openDocumentModal(displayMemory.document_id)}
|
||||
variant="secondary"
|
||||
className="flex-1"
|
||||
>
|
||||
View Document
|
||||
</Button>
|
||||
)}
|
||||
{displayMemory.chunk_id && (
|
||||
<Button
|
||||
onClick={() => openChunkModal(displayMemory.chunk_id)}
|
||||
variant="secondary"
|
||||
className="flex-1"
|
||||
>
|
||||
View Chunk
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* Document/Chunk Modal */}
|
||||
@@ -225,123 +290,158 @@ export function MemoryDetailPanel({
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
<div className={gap}>
|
||||
{/* Full Text */}
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Full Text
|
||||
</div>
|
||||
<div className={`${textSize} whitespace-pre-wrap`}>{memory.text}</div>
|
||||
{loading ? (
|
||||
<div className="flex items-center justify-center py-8">
|
||||
<Loader2 className="h-5 w-5 animate-spin text-muted-foreground" />
|
||||
<span className="ml-2 text-sm text-muted-foreground">Loading...</span>
|
||||
</div>
|
||||
|
||||
{/* Context */}
|
||||
{memory.context && (
|
||||
) : (
|
||||
<div className={gap}>
|
||||
{/* Full Text */}
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Context
|
||||
Full Text
|
||||
</div>
|
||||
<div className={textSize}>{memory.context}</div>
|
||||
<div className={`${textSize} whitespace-pre-wrap`}>{displayMemory.text}</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Dates */}
|
||||
<div className="grid grid-cols-2 gap-2">
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Occurred
|
||||
{/* Context */}
|
||||
{displayMemory.context && (
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Context
|
||||
</div>
|
||||
<div className={textSize}>{displayMemory.context}</div>
|
||||
</div>
|
||||
<div className={textSize}>
|
||||
{memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"}
|
||||
</div>
|
||||
</div>
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Mentioned
|
||||
</div>
|
||||
<div className={textSize}>
|
||||
{memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Entities */}
|
||||
{memory.entities && (
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
|
||||
Entities
|
||||
{/* Dates */}
|
||||
<div className="grid grid-cols-2 gap-2">
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Occurred
|
||||
</div>
|
||||
<div className={textSize}>
|
||||
{displayMemory.occurred_start
|
||||
? new Date(displayMemory.occurred_start).toLocaleString()
|
||||
: "N/A"}
|
||||
</div>
|
||||
</div>
|
||||
<div className="flex flex-wrap gap-1">
|
||||
{(Array.isArray(memory.entities)
|
||||
? memory.entities
|
||||
: String(memory.entities).split(", ")
|
||||
).map((entity: any, i: number) => {
|
||||
const entityText =
|
||||
typeof entity === "string" ? entity : entity?.name || JSON.stringify(entity);
|
||||
return (
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Mentioned
|
||||
</div>
|
||||
<div className={textSize}>
|
||||
{displayMemory.mentioned_at
|
||||
? new Date(displayMemory.mentioned_at).toLocaleString()
|
||||
: "N/A"}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Entities */}
|
||||
{displayMemory.entities &&
|
||||
(Array.isArray(displayMemory.entities)
|
||||
? displayMemory.entities.length > 0
|
||||
: displayMemory.entities) && (
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
|
||||
Entities
|
||||
</div>
|
||||
<div className="flex flex-wrap gap-1">
|
||||
{(Array.isArray(displayMemory.entities)
|
||||
? displayMemory.entities
|
||||
: String(displayMemory.entities).split(", ")
|
||||
).map((entity: any, i: number) => {
|
||||
const entityText =
|
||||
typeof entity === "string"
|
||||
? entity
|
||||
: entity?.name || JSON.stringify(entity);
|
||||
return (
|
||||
<span
|
||||
key={i}
|
||||
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-secondary text-secondary-foreground`}
|
||||
>
|
||||
{entityText}
|
||||
</span>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Tags */}
|
||||
{displayMemory.tags && displayMemory.tags.length > 0 && (
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
|
||||
Tags
|
||||
</div>
|
||||
<div className="flex flex-wrap gap-1">
|
||||
{displayMemory.tags.map((tag: string, i: number) => (
|
||||
<span
|
||||
key={i}
|
||||
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-secondary text-secondary-foreground`}
|
||||
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-amber-500/10 text-amber-600 dark:text-amber-400`}
|
||||
>
|
||||
{entityText}
|
||||
{tag}
|
||||
</span>
|
||||
);
|
||||
})}
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
)}
|
||||
|
||||
{/* ID */}
|
||||
{memoryId && (
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Memory ID
|
||||
{/* ID */}
|
||||
{memoryId && (
|
||||
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
|
||||
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
|
||||
Memory ID
|
||||
</div>
|
||||
<div className="flex items-center gap-2">
|
||||
<span className={`${compact ? "text-[10px]" : "text-sm"} font-mono break-all`}>
|
||||
{memoryId}
|
||||
</span>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
className="h-6 w-6 p-0 flex-shrink-0"
|
||||
onClick={() => copyToClipboard(memoryId)}
|
||||
>
|
||||
{copiedId === memoryId ? (
|
||||
<Check className="h-3 w-3 text-green-600" />
|
||||
) : (
|
||||
<Copy className="h-3 w-3" />
|
||||
)}
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
<div className="flex items-center gap-2">
|
||||
<span className={`${compact ? "text-[10px]" : "text-sm"} font-mono break-all`}>
|
||||
{memoryId}
|
||||
</span>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
className="h-6 w-6 p-0 flex-shrink-0"
|
||||
onClick={() => copyToClipboard(memoryId)}
|
||||
>
|
||||
{copiedId === memoryId ? (
|
||||
<Check className="h-3 w-3 text-green-600" />
|
||||
) : (
|
||||
<Copy className="h-3 w-3" />
|
||||
)}
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
)}
|
||||
|
||||
{/* Document/Chunk buttons */}
|
||||
{(memory.document_id || memory.chunk_id) && (
|
||||
<div className={`flex gap-2 ${compact ? "pt-1" : ""}`}>
|
||||
{memory.document_id && (
|
||||
<Button
|
||||
onClick={() => openDocumentModal(memory.document_id)}
|
||||
size="sm"
|
||||
variant="secondary"
|
||||
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
|
||||
>
|
||||
View Document
|
||||
</Button>
|
||||
)}
|
||||
{memory.chunk_id && (
|
||||
<Button
|
||||
onClick={() => openChunkModal(memory.chunk_id)}
|
||||
size="sm"
|
||||
variant="secondary"
|
||||
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
|
||||
>
|
||||
View Chunk
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
{/* Document/Chunk buttons */}
|
||||
{(displayMemory.document_id || displayMemory.chunk_id) && (
|
||||
<div className={`flex gap-2 ${compact ? "pt-1" : ""}`}>
|
||||
{displayMemory.document_id && (
|
||||
<Button
|
||||
onClick={() => openDocumentModal(displayMemory.document_id)}
|
||||
size="sm"
|
||||
variant="secondary"
|
||||
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
|
||||
>
|
||||
View Document
|
||||
</Button>
|
||||
)}
|
||||
{displayMemory.chunk_id && (
|
||||
<Button
|
||||
onClick={() => openChunkModal(displayMemory.chunk_id)}
|
||||
size="sm"
|
||||
variant="secondary"
|
||||
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
|
||||
>
|
||||
View Chunk
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* Document/Chunk Modal */}
|
||||
|
||||
@@ -25,6 +25,8 @@ import {
|
||||
FileText,
|
||||
Users,
|
||||
ArrowDown,
|
||||
Tag,
|
||||
Calendar,
|
||||
} from "lucide-react";
|
||||
import JsonView from "react18-json-view";
|
||||
import "react18-json-view/src/style.css";
|
||||
@@ -32,6 +34,7 @@ import { MemoryDetailPanel } from "./memory-detail-panel";
|
||||
|
||||
type FactType = "world" | "experience" | "opinion";
|
||||
type Budget = "low" | "mid" | "high";
|
||||
type TagsMatch = "any" | "all" | "any_strict" | "all_strict";
|
||||
type ViewMode = "results" | "trace" | "json";
|
||||
|
||||
export function SearchDebugView() {
|
||||
@@ -45,6 +48,8 @@ export function SearchDebugView() {
|
||||
const [queryDate, setQueryDate] = useState("");
|
||||
const [includeChunks, setIncludeChunks] = useState(false);
|
||||
const [includeEntities, setIncludeEntities] = useState(false);
|
||||
const [tags, setTags] = useState("");
|
||||
const [tagsMatch, setTagsMatch] = useState<TagsMatch>("any");
|
||||
|
||||
// Results state
|
||||
const [results, setResults] = useState<any[] | null>(null);
|
||||
@@ -83,6 +88,14 @@ export function SearchDebugView() {
|
||||
|
||||
const INITIAL_RESULTS_COUNT = 5;
|
||||
|
||||
// Helper to find full memory data from results when clicking trace items
|
||||
const selectMemoryFromTrace = (traceResult: any) => {
|
||||
const nodeId = traceResult.id || traceResult.node_id;
|
||||
// Try to find the full result with all metadata
|
||||
const fullResult = results?.find((r: any) => r.id === nodeId || r.node_id === nodeId);
|
||||
setSelectedMemory(fullResult || traceResult);
|
||||
};
|
||||
|
||||
const runSearch = async () => {
|
||||
if (!currentBank) {
|
||||
alert("Please select a memory bank first");
|
||||
@@ -99,6 +112,12 @@ export function SearchDebugView() {
|
||||
setLoading(true);
|
||||
|
||||
try {
|
||||
// Parse tags from comma-separated string
|
||||
const parsedTags = tags
|
||||
.split(",")
|
||||
.map((t) => t.trim())
|
||||
.filter((t) => t.length > 0);
|
||||
|
||||
const requestBody: any = {
|
||||
bank_id: currentBank,
|
||||
query: query,
|
||||
@@ -111,6 +130,7 @@ export function SearchDebugView() {
|
||||
chunks: includeChunks ? { max_tokens: 8192 } : null,
|
||||
},
|
||||
...(queryDate && { query_timestamp: queryDate }),
|
||||
...(parsedTags.length > 0 && { tags: parsedTags, tags_match: tagsMatch }),
|
||||
};
|
||||
|
||||
const data: any = await client.recall(requestBody);
|
||||
@@ -246,6 +266,31 @@ export function SearchDebugView() {
|
||||
</label>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Tags Filter */}
|
||||
<div className="flex items-center gap-4 mt-4 pt-4 border-t">
|
||||
<Tag className="h-4 w-4 text-muted-foreground" />
|
||||
<div className="flex-1 max-w-md">
|
||||
<Input
|
||||
type="text"
|
||||
value={tags}
|
||||
onChange={(e) => setTags(e.target.value)}
|
||||
placeholder="Filter by tags (comma-separated)"
|
||||
className="h-8"
|
||||
/>
|
||||
</div>
|
||||
<Select value={tagsMatch} onValueChange={(v) => setTagsMatch(v as TagsMatch)}>
|
||||
<SelectTrigger className="w-40 h-8">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem value="any">Any (incl. untagged)</SelectItem>
|
||||
<SelectItem value="all">All (incl. untagged)</SelectItem>
|
||||
<SelectItem value="any_strict">Any (strict)</SelectItem>
|
||||
<SelectItem value="all_strict">All (strict)</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
</CardContent>
|
||||
</Card>
|
||||
|
||||
@@ -507,9 +552,29 @@ export function SearchDebugView() {
|
||||
}}
|
||||
>
|
||||
<div className="flex items-center justify-between mb-1">
|
||||
<span className="font-medium text-sm text-foreground capitalize">
|
||||
{method.method_name}
|
||||
</span>
|
||||
<div className="flex items-center gap-2">
|
||||
<span className="font-medium text-sm text-foreground capitalize">
|
||||
{method.method_name}
|
||||
</span>
|
||||
{/* Show temporal range inline */}
|
||||
{method.method_name === "temporal" &&
|
||||
method.metadata?.constraint && (
|
||||
<span className="flex items-center gap-1 text-[10px] text-muted-foreground">
|
||||
<Calendar className="h-3 w-3" />
|
||||
{method.metadata.constraint.start
|
||||
? new Date(
|
||||
method.metadata.constraint.start
|
||||
).toLocaleDateString()
|
||||
: "any"}
|
||||
{" → "}
|
||||
{method.metadata.constraint.end
|
||||
? new Date(
|
||||
method.metadata.constraint.end
|
||||
).toLocaleDateString()
|
||||
: "any"}
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
{isMethodExpanded ? (
|
||||
<ChevronDown className="h-3 w-3 text-muted-foreground" />
|
||||
) : (
|
||||
@@ -546,7 +611,7 @@ export function SearchDebugView() {
|
||||
className="p-2 bg-background rounded cursor-pointer hover:bg-muted/50 transition-colors border border-border"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation();
|
||||
setSelectedMemory(r);
|
||||
selectMemoryFromTrace(r);
|
||||
}}
|
||||
>
|
||||
<div className="flex items-start gap-2">
|
||||
@@ -684,7 +749,7 @@ export function SearchDebugView() {
|
||||
className="p-3 bg-muted/30 rounded-lg cursor-pointer hover:bg-muted/50 transition-colors"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation();
|
||||
setSelectedMemory(r);
|
||||
selectMemoryFromTrace(r);
|
||||
}}
|
||||
>
|
||||
<div className="flex items-start gap-3">
|
||||
@@ -792,7 +857,7 @@ export function SearchDebugView() {
|
||||
className="p-3 bg-muted/30 rounded-lg cursor-pointer hover:bg-muted/50 transition-colors"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation();
|
||||
setSelectedMemory(r);
|
||||
selectMemoryFromTrace(r);
|
||||
}}
|
||||
>
|
||||
<div className="flex items-start gap-3">
|
||||
@@ -931,6 +996,7 @@ export function SearchDebugView() {
|
||||
memory={selectedMemory}
|
||||
onClose={() => setSelectedMemory(null)}
|
||||
inPanel
|
||||
bankId={currentBank || undefined}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
|
||||
@@ -15,10 +15,12 @@ import {
|
||||
} from "@/components/ui/select";
|
||||
import { Checkbox } from "@/components/ui/checkbox";
|
||||
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card";
|
||||
import { Sparkles, Info } from "lucide-react";
|
||||
import { Sparkles, Info, Tag } from "lucide-react";
|
||||
import JsonView from "react18-json-view";
|
||||
import "react18-json-view/src/style.css";
|
||||
|
||||
type TagsMatch = "any" | "all" | "any_strict" | "all_strict";
|
||||
|
||||
export function ThinkView() {
|
||||
const { currentBank } = useBank();
|
||||
const [query, setQuery] = useState("");
|
||||
@@ -28,6 +30,8 @@ export function ThinkView() {
|
||||
const [showRawJson, setShowRawJson] = useState(false);
|
||||
const [result, setResult] = useState<any>(null);
|
||||
const [loading, setLoading] = useState(false);
|
||||
const [tags, setTags] = useState("");
|
||||
const [tagsMatch, setTagsMatch] = useState<TagsMatch>("any");
|
||||
|
||||
const runReflect = async () => {
|
||||
if (!currentBank || !query) return;
|
||||
@@ -35,12 +39,19 @@ export function ThinkView() {
|
||||
setLoading(true);
|
||||
setShowRawJson(false);
|
||||
try {
|
||||
// Parse tags from comma-separated string
|
||||
const parsedTags = tags
|
||||
.split(",")
|
||||
.map((t) => t.trim())
|
||||
.filter((t) => t.length > 0);
|
||||
|
||||
const data: any = await client.reflect({
|
||||
bank_id: currentBank,
|
||||
query,
|
||||
budget,
|
||||
context: context || undefined,
|
||||
include_facts: includeFacts,
|
||||
...(parsedTags.length > 0 && { tags: parsedTags, tags_match: tagsMatch }),
|
||||
});
|
||||
setResult(data);
|
||||
} catch (error) {
|
||||
@@ -103,6 +114,29 @@ export function ThinkView() {
|
||||
rows={3}
|
||||
/>
|
||||
</div>
|
||||
<div className="flex items-center gap-4 mt-4 pt-4 border-t">
|
||||
<Tag className="h-4 w-4 text-muted-foreground" />
|
||||
<div className="flex-1 max-w-md">
|
||||
<Input
|
||||
type="text"
|
||||
value={tags}
|
||||
onChange={(e) => setTags(e.target.value)}
|
||||
placeholder="Filter by tags (comma-separated)"
|
||||
className="h-8"
|
||||
/>
|
||||
</div>
|
||||
<Select value={tagsMatch} onValueChange={(v) => setTagsMatch(v as TagsMatch)}>
|
||||
<SelectTrigger className="w-40 h-8">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
<SelectItem value="any">Any (incl. untagged)</SelectItem>
|
||||
<SelectItem value="all">All (incl. untagged)</SelectItem>
|
||||
<SelectItem value="any_strict">Any (strict)</SelectItem>
|
||||
<SelectItem value="all_strict">All (strict)</SelectItem>
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
</CardContent>
|
||||
</Card>
|
||||
|
||||
|
||||
@@ -53,6 +53,8 @@ export class ControlPlaneClient {
|
||||
chunks?: { max_tokens: number } | null;
|
||||
};
|
||||
query_timestamp?: string;
|
||||
tags?: string[];
|
||||
tags_match?: "any" | "all" | "any_strict" | "all_strict";
|
||||
}) {
|
||||
return this.fetchApi("/api/recall", {
|
||||
method: "POST",
|
||||
@@ -69,6 +71,8 @@ export class ControlPlaneClient {
|
||||
budget?: string;
|
||||
context?: string;
|
||||
include_facts?: boolean;
|
||||
tags?: string[];
|
||||
tags_match?: "any" | "all" | "any_strict" | "all_strict";
|
||||
}) {
|
||||
return this.fetchApi("/api/reflect", {
|
||||
method: "POST",
|
||||
@@ -209,6 +213,26 @@ export class ControlPlaneClient {
|
||||
return this.fetchApi(`/api/chunks/${chunkId}`);
|
||||
}
|
||||
|
||||
/**
|
||||
* Get a single memory by ID
|
||||
*/
|
||||
async getMemory(memoryId: string, bankId: string) {
|
||||
return this.fetchApi<{
|
||||
id: string;
|
||||
text: string;
|
||||
context: string;
|
||||
date: string;
|
||||
type: string;
|
||||
mentioned_at: string | null;
|
||||
occurred_start: string | null;
|
||||
occurred_end: string | null;
|
||||
entities: string[];
|
||||
document_id: string | null;
|
||||
chunk_id: string | null;
|
||||
tags: string[];
|
||||
}>(`/api/memories/${memoryId}?bank_id=${bankId}`);
|
||||
}
|
||||
|
||||
/**
|
||||
* Get bank profile
|
||||
*/
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -139,13 +139,18 @@ export HINDSIGHT_API_REFLECT_LLM_MODEL=llama-3.3-70b-versatile
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `HINDSIGHT_API_EMBEDDINGS_PROVIDER` | Provider: `local`, `tei`, `openai`, or `cohere` | `local` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_PROVIDER` | Provider: `local`, `tei`, `openai`, `cohere`, or `litellm` | `local` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL` | Model for local provider | `BAAI/bge-small-en-v1.5` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_TEI_URL` | TEI server URL | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY` | OpenAI API key (falls back to `HINDSIGHT_API_LLM_API_KEY`) | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL` | OpenAI embedding model | `text-embedding-3-small` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL` | Custom base URL for OpenAI-compatible API (e.g., Azure OpenAI) | - |
|
||||
| `HINDSIGHT_API_COHERE_API_KEY` | Cohere API key (shared for embeddings and reranker) | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL` | Cohere embedding model | `embed-english-v3.0` |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL` | Custom base URL for Cohere-compatible API (e.g., Azure-hosted) | - |
|
||||
| `HINDSIGHT_API_LITELLM_API_BASE` | LiteLLM proxy base URL (shared for embeddings and reranker) | `http://localhost:4000` |
|
||||
| `HINDSIGHT_API_LITELLM_API_KEY` | LiteLLM proxy API key (optional, depends on proxy config) | - |
|
||||
| `HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL` | LiteLLM embedding model (use provider prefix, e.g., `cohere/embed-english-v3.0`) | `text-embedding-3-small` |
|
||||
|
||||
```bash
|
||||
# Local (default) - uses SentenceTransformers
|
||||
@@ -157,6 +162,12 @@ export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=sk-xxxxxxxxxxxx # or reuses HINDSIGHT_API_LLM_API_KEY
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL=text-embedding-3-small # 1536 dimensions
|
||||
|
||||
# Azure OpenAI - embeddings via Azure endpoint
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=your-azure-api-key
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL=text-embedding-3-small
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL=https://your-resource.openai.azure.com/openai/deployments/your-deployment
|
||||
|
||||
# TEI - HuggingFace Text Embeddings Inference (recommended for production)
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=tei
|
||||
export HINDSIGHT_API_EMBEDDINGS_TEI_URL=http://localhost:8080
|
||||
@@ -165,6 +176,18 @@ export HINDSIGHT_API_EMBEDDINGS_TEI_URL=http://localhost:8080
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-api-key
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL=embed-english-v3.0 # 1024 dimensions
|
||||
|
||||
# Azure-hosted Cohere - embeddings via custom endpoint
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-azure-api-key
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL=embed-english-v3.0
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL=https://your-azure-cohere-endpoint.com
|
||||
|
||||
# LiteLLM proxy - unified gateway for multiple providers
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=litellm
|
||||
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
|
||||
export HINDSIGHT_API_LITELLM_API_KEY=your-litellm-key # optional
|
||||
export HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL=text-embedding-3-small # or cohere/embed-english-v3.0
|
||||
```
|
||||
|
||||
#### Embedding Dimensions
|
||||
@@ -187,13 +210,15 @@ Supported OpenAI embedding dimensions:
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `HINDSIGHT_API_RERANKER_PROVIDER` | Provider: `local`, `tei`, or `cohere` | `local` |
|
||||
| `HINDSIGHT_API_RERANKER_PROVIDER` | Provider: `local`, `tei`, `cohere`, `flashrank`, `litellm`, or `rrf` | `local` |
|
||||
| `HINDSIGHT_API_RERANKER_LOCAL_MODEL` | Model for local provider | `cross-encoder/ms-marco-MiniLM-L-6-v2` |
|
||||
| `HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT` | Max concurrent local reranking (prevents CPU thrashing under load) | `4` |
|
||||
| `HINDSIGHT_API_RERANKER_TEI_URL` | TEI server URL | - |
|
||||
| `HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE` | Batch size for TEI reranking | `128` |
|
||||
| `HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT` | Max concurrent TEI reranking requests | `8` |
|
||||
| `HINDSIGHT_API_RERANKER_COHERE_MODEL` | Cohere rerank model | `rerank-english-v3.0` |
|
||||
| `HINDSIGHT_API_RERANKER_COHERE_BASE_URL` | Custom base URL for Cohere-compatible API (e.g., Azure-hosted) | - |
|
||||
| `HINDSIGHT_API_RERANKER_LITELLM_MODEL` | LiteLLM rerank model (use provider prefix, e.g., `cohere/rerank-english-v3.0`) | `cohere/rerank-english-v3.0` |
|
||||
|
||||
```bash
|
||||
# Local (default) - uses SentenceTransformers CrossEncoder
|
||||
@@ -208,8 +233,27 @@ export HINDSIGHT_API_RERANKER_TEI_URL=http://localhost:8081
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-api-key # shared with embeddings
|
||||
export HINDSIGHT_API_RERANKER_COHERE_MODEL=rerank-english-v3.0
|
||||
|
||||
# Azure-hosted Cohere - reranking via custom endpoint
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-azure-api-key
|
||||
export HINDSIGHT_API_RERANKER_COHERE_MODEL=rerank-english-v3.0
|
||||
export HINDSIGHT_API_RERANKER_COHERE_BASE_URL=https://your-azure-cohere-endpoint.com
|
||||
|
||||
# LiteLLM proxy - unified gateway for multiple reranking providers
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=litellm
|
||||
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
|
||||
export HINDSIGHT_API_LITELLM_API_KEY=your-litellm-key # optional
|
||||
export HINDSIGHT_API_RERANKER_LITELLM_MODEL=cohere/rerank-english-v3.0 # or voyage/rerank-2, together_ai/...
|
||||
```
|
||||
|
||||
LiteLLM supports multiple reranking providers via the `/rerank` endpoint:
|
||||
- Cohere (`cohere/rerank-english-v3.0`, `cohere/rerank-multilingual-v3.0`)
|
||||
- Together AI (`together_ai/...`)
|
||||
- Voyage AI (`voyage/rerank-2`)
|
||||
- Jina AI (`jina_ai/...`)
|
||||
- AWS Bedrock (`bedrock/...`)
|
||||
|
||||
### Authentication
|
||||
|
||||
By default, Hindsight runs without authentication. For production deployments, enable API key authentication using the built-in tenant extension:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -249,6 +249,72 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/default/banks/{bank_id}/memories/{memory_id}": {
|
||||
"get": {
|
||||
"tags": [
|
||||
"Memory"
|
||||
],
|
||||
"summary": "Get memory unit",
|
||||
"description": "Get a single memory unit by ID with all its metadata including entities and tags.",
|
||||
"operationId": "get_memory",
|
||||
"parameters": [
|
||||
{
|
||||
"name": "bank_id",
|
||||
"in": "path",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"title": "Bank Id"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "memory_id",
|
||||
"in": "path",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"title": "Memory Id"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "authorization",
|
||||
"in": "header",
|
||||
"required": false,
|
||||
"schema": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Authorization"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Successful Response",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
}
|
||||
},
|
||||
"422": {
|
||||
"description": "Validation Error",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/default/banks/{bank_id}/memories/recall": {
|
||||
"post": {
|
||||
"tags": [
|
||||
@@ -944,6 +1010,107 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/default/banks/{bank_id}/tags": {
|
||||
"get": {
|
||||
"tags": [
|
||||
"Memory"
|
||||
],
|
||||
"summary": "List tags",
|
||||
"description": "List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.",
|
||||
"operationId": "list_tags",
|
||||
"parameters": [
|
||||
{
|
||||
"name": "bank_id",
|
||||
"in": "path",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"type": "string",
|
||||
"title": "Bank Id"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "q",
|
||||
"in": "query",
|
||||
"required": false,
|
||||
"schema": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.",
|
||||
"title": "Q"
|
||||
},
|
||||
"description": "Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive."
|
||||
},
|
||||
{
|
||||
"name": "limit",
|
||||
"in": "query",
|
||||
"required": false,
|
||||
"schema": {
|
||||
"type": "integer",
|
||||
"description": "Maximum number of tags to return",
|
||||
"default": 100,
|
||||
"title": "Limit"
|
||||
},
|
||||
"description": "Maximum number of tags to return"
|
||||
},
|
||||
{
|
||||
"name": "offset",
|
||||
"in": "query",
|
||||
"required": false,
|
||||
"schema": {
|
||||
"type": "integer",
|
||||
"description": "Offset for pagination",
|
||||
"default": 0,
|
||||
"title": "Offset"
|
||||
},
|
||||
"description": "Offset for pagination"
|
||||
},
|
||||
{
|
||||
"name": "authorization",
|
||||
"in": "header",
|
||||
"required": false,
|
||||
"schema": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Authorization"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"description": "Successful Response",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/ListTagsResponse"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"422": {
|
||||
"description": "Validation Error",
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
"/v1/default/chunks/{chunk_id}": {
|
||||
"get": {
|
||||
"tags": [
|
||||
@@ -2219,6 +2386,14 @@
|
||||
"memory_unit_count": {
|
||||
"type": "integer",
|
||||
"title": "Memory Unit Count"
|
||||
},
|
||||
"tags": {
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Tags",
|
||||
"description": "Tags associated with this document"
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -2240,6 +2415,10 @@
|
||||
"id": "session_1",
|
||||
"memory_unit_count": 15,
|
||||
"original_text": "Full document text here...",
|
||||
"tags": [
|
||||
"user_a",
|
||||
"session_123"
|
||||
],
|
||||
"updated_at": "2024-01-15T10:30:00Z"
|
||||
}
|
||||
},
|
||||
@@ -2752,6 +2931,57 @@
|
||||
"total": 150
|
||||
}
|
||||
},
|
||||
"ListTagsResponse": {
|
||||
"properties": {
|
||||
"items": {
|
||||
"items": {
|
||||
"$ref": "#/components/schemas/TagItem"
|
||||
},
|
||||
"type": "array",
|
||||
"title": "Items"
|
||||
},
|
||||
"total": {
|
||||
"type": "integer",
|
||||
"title": "Total"
|
||||
},
|
||||
"limit": {
|
||||
"type": "integer",
|
||||
"title": "Limit"
|
||||
},
|
||||
"offset": {
|
||||
"type": "integer",
|
||||
"title": "Offset"
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
"required": [
|
||||
"items",
|
||||
"total",
|
||||
"limit",
|
||||
"offset"
|
||||
],
|
||||
"title": "ListTagsResponse",
|
||||
"description": "Response model for list tags endpoint.",
|
||||
"example": {
|
||||
"items": [
|
||||
{
|
||||
"count": 42,
|
||||
"tag": "user:alice"
|
||||
},
|
||||
{
|
||||
"count": 15,
|
||||
"tag": "user:bob"
|
||||
},
|
||||
{
|
||||
"count": 8,
|
||||
"tag": "session:abc123"
|
||||
}
|
||||
],
|
||||
"limit": 100,
|
||||
"offset": 0,
|
||||
"total": 25
|
||||
}
|
||||
},
|
||||
"MemoryItem": {
|
||||
"properties": {
|
||||
"content": {
|
||||
@@ -2821,6 +3051,21 @@
|
||||
],
|
||||
"title": "Entities",
|
||||
"description": "Optional entities to combine with auto-extracted entities."
|
||||
},
|
||||
"tags": {
|
||||
"anyOf": [
|
||||
{
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"type": "array"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Tags",
|
||||
"description": "Optional tags for visibility scoping. Memories with tags can be filtered during recall."
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -2846,6 +3091,10 @@
|
||||
"channel": "engineering",
|
||||
"source": "slack"
|
||||
},
|
||||
"tags": [
|
||||
"user_a",
|
||||
"user_b"
|
||||
],
|
||||
"timestamp": "2024-01-15T10:30:00Z"
|
||||
}
|
||||
},
|
||||
@@ -2999,6 +3248,33 @@
|
||||
"include": {
|
||||
"$ref": "#/components/schemas/IncludeOptions",
|
||||
"description": "Options for including additional data (entities are included by default)"
|
||||
},
|
||||
"tags": {
|
||||
"anyOf": [
|
||||
{
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"type": "array"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Tags",
|
||||
"description": "Filter memories by tags. If not specified, all memories are returned."
|
||||
},
|
||||
"tags_match": {
|
||||
"type": "string",
|
||||
"enum": [
|
||||
"any",
|
||||
"all",
|
||||
"any_strict",
|
||||
"all_strict"
|
||||
],
|
||||
"title": "Tags Match",
|
||||
"description": "How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
|
||||
"default": "any"
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -3017,6 +3293,10 @@
|
||||
"max_tokens": 4096,
|
||||
"query": "What did Alice say about machine learning?",
|
||||
"query_timestamp": "2023-05-30T23:40:00",
|
||||
"tags": [
|
||||
"user_a"
|
||||
],
|
||||
"tags_match": "any",
|
||||
"trace": true,
|
||||
"types": [
|
||||
"world",
|
||||
@@ -3238,6 +3518,20 @@
|
||||
}
|
||||
],
|
||||
"title": "Chunk Id"
|
||||
},
|
||||
"tags": {
|
||||
"anyOf": [
|
||||
{
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"type": "array"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Tags"
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -3262,6 +3556,10 @@
|
||||
},
|
||||
"occurred_end": "2024-01-15T10:30:00Z",
|
||||
"occurred_start": "2024-01-15T10:30:00Z",
|
||||
"tags": [
|
||||
"user_a",
|
||||
"user_b"
|
||||
],
|
||||
"text": "Alice works at Google on the AI team",
|
||||
"type": "world"
|
||||
}
|
||||
@@ -3404,6 +3702,33 @@
|
||||
],
|
||||
"title": "Response Schema",
|
||||
"description": "Optional JSON Schema for structured output. When provided, the response will include a 'structured_output' field with the LLM response parsed according to this schema."
|
||||
},
|
||||
"tags": {
|
||||
"anyOf": [
|
||||
{
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"type": "array"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Tags",
|
||||
"description": "Filter memories by tags during reflection. If not specified, all memories are considered."
|
||||
},
|
||||
"tags_match": {
|
||||
"type": "string",
|
||||
"enum": [
|
||||
"any",
|
||||
"all",
|
||||
"any_strict",
|
||||
"all_strict"
|
||||
],
|
||||
"title": "Tags Match",
|
||||
"description": "How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).",
|
||||
"default": "any"
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -3437,7 +3762,11 @@
|
||||
"key_points"
|
||||
],
|
||||
"type": "object"
|
||||
}
|
||||
},
|
||||
"tags": [
|
||||
"user_a"
|
||||
],
|
||||
"tags_match": "any"
|
||||
}
|
||||
},
|
||||
"ReflectResponse": {
|
||||
@@ -3527,6 +3856,21 @@
|
||||
"title": "Async",
|
||||
"description": "If true, process asynchronously in background. If false, wait for completion (default: false)",
|
||||
"default": false
|
||||
},
|
||||
"document_tags": {
|
||||
"anyOf": [
|
||||
{
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"type": "array"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Document Tags",
|
||||
"description": "Tags applied to all items in this request. These are merged with any item-level tags."
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
@@ -3537,6 +3881,10 @@
|
||||
"description": "Request model for retain endpoint.",
|
||||
"example": {
|
||||
"async": false,
|
||||
"document_tags": [
|
||||
"user_a",
|
||||
"user_b"
|
||||
],
|
||||
"items": [
|
||||
{
|
||||
"content": "Alice works at Google",
|
||||
@@ -3615,6 +3963,27 @@
|
||||
}
|
||||
}
|
||||
},
|
||||
"TagItem": {
|
||||
"properties": {
|
||||
"tag": {
|
||||
"type": "string",
|
||||
"title": "Tag",
|
||||
"description": "The tag value"
|
||||
},
|
||||
"count": {
|
||||
"type": "integer",
|
||||
"title": "Count",
|
||||
"description": "Number of memories with this tag"
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
"required": [
|
||||
"tag",
|
||||
"count"
|
||||
],
|
||||
"title": "TagItem",
|
||||
"description": "Single tag with usage count."
|
||||
},
|
||||
"TokenUsage": {
|
||||
"properties": {
|
||||
"input_tokens": {
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "hindsight-all"
|
||||
version = "0.2.1"
|
||||
version = "0.3.0"
|
||||
description = "Hindsight: Agent Memory That Works Like Human Memory - All-in-One Bundle"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
|
||||
@@ -54,7 +54,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(hindsight_http_requests_total)",
|
||||
"expr": "sum(hindsight_http_requests_total{tenant=~\"$tenant\"})",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -99,7 +99,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(rate(hindsight_http_requests_total[1m]))",
|
||||
"expr": "sum(rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -191,7 +191,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"5xx\"}[5m])) / sum(rate(hindsight_http_requests_total[5m]))",
|
||||
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"5xx\", tenant=~\"$tenant\"}[5m])) / sum(rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[5m]))",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -236,7 +236,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))",
|
||||
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -313,7 +313,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum by (endpoint) (rate(hindsight_http_requests_total[1m]))",
|
||||
"expr": "sum by (endpoint) (rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "{{endpoint}}",
|
||||
"refId": "A"
|
||||
}
|
||||
@@ -404,17 +404,17 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "histogram_quantile(0.50, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))",
|
||||
"expr": "histogram_quantile(0.50, sum by (le) (rate(hindsight_http_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
|
||||
"legendFormat": "p50",
|
||||
"refId": "A"
|
||||
},
|
||||
{
|
||||
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))",
|
||||
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
|
||||
"legendFormat": "p95",
|
||||
"refId": "B"
|
||||
},
|
||||
{
|
||||
"expr": "histogram_quantile(0.99, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))",
|
||||
"expr": "histogram_quantile(0.99, sum by (le) (rate(hindsight_http_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
|
||||
"legendFormat": "p99",
|
||||
"refId": "C"
|
||||
}
|
||||
@@ -505,12 +505,12 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"5xx\"}[1m])) / sum(rate(hindsight_http_requests_total[1m]))",
|
||||
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"5xx\", tenant=~\"$tenant\"}[1m])) / sum(rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "5xx Error Rate",
|
||||
"refId": "A"
|
||||
},
|
||||
{
|
||||
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"4xx\"}[1m])) / sum(rate(hindsight_http_requests_total[1m]))",
|
||||
"expr": "sum(rate(hindsight_http_requests_total{status_class=\"4xx\", tenant=~\"$tenant\"}[1m])) / sum(rate(hindsight_http_requests_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "4xx Error Rate",
|
||||
"refId": "B"
|
||||
}
|
||||
@@ -1276,7 +1276,36 @@
|
||||
"schemaVersion": 38,
|
||||
"tags": ["hindsight", "api", "service"],
|
||||
"templating": {
|
||||
"list": []
|
||||
"list": [
|
||||
{
|
||||
"allValue": ".*",
|
||||
"current": {
|
||||
"selected": true,
|
||||
"text": "All",
|
||||
"value": "$__all"
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "prometheus"
|
||||
},
|
||||
"definition": "label_values(hindsight_http_requests_total, tenant)",
|
||||
"hide": 0,
|
||||
"includeAll": true,
|
||||
"label": "Tenant",
|
||||
"multi": false,
|
||||
"name": "tenant",
|
||||
"options": [],
|
||||
"query": {
|
||||
"query": "label_values(hindsight_http_requests_total, tenant)",
|
||||
"refId": "PrometheusVariableQueryEditor-VariableQuery"
|
||||
},
|
||||
"refresh": 2,
|
||||
"regex": "",
|
||||
"skipUrlSync": false,
|
||||
"sort": 1,
|
||||
"type": "query"
|
||||
}
|
||||
]
|
||||
},
|
||||
"time": {
|
||||
"from": "now-30m",
|
||||
|
||||
@@ -46,7 +46,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(hindsight_llm_calls_total)",
|
||||
"expr": "sum(hindsight_llm_calls_total{tenant=~\"$tenant\"})",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -91,7 +91,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(hindsight_llm_tokens_input_tokens_total) + sum(hindsight_llm_tokens_output_tokens_total)",
|
||||
"expr": "sum(hindsight_llm_tokens_input_tokens_total{tenant=~\"$tenant\"}) + sum(hindsight_llm_tokens_output_tokens_total{tenant=~\"$tenant\"})",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -137,7 +137,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(hindsight_llm_tokens_input_tokens_total)",
|
||||
"expr": "sum(hindsight_llm_tokens_input_tokens_total{tenant=~\"$tenant\"})",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -183,7 +183,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(hindsight_llm_tokens_output_tokens_total)",
|
||||
"expr": "sum(hindsight_llm_tokens_output_tokens_total{tenant=~\"$tenant\"})",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -260,7 +260,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum by (scope) (rate(hindsight_llm_calls_total[1m]))",
|
||||
"expr": "sum by (scope) (rate(hindsight_llm_calls_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "{{scope}}",
|
||||
"refId": "A"
|
||||
}
|
||||
@@ -347,12 +347,12 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(rate(hindsight_llm_tokens_input_tokens_total[1m]))",
|
||||
"expr": "sum(rate(hindsight_llm_tokens_input_tokens_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "Input",
|
||||
"refId": "A"
|
||||
},
|
||||
{
|
||||
"expr": "sum(rate(hindsight_llm_tokens_output_tokens_total[1m]))",
|
||||
"expr": "sum(rate(hindsight_llm_tokens_output_tokens_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "Output",
|
||||
"refId": "B"
|
||||
}
|
||||
@@ -430,7 +430,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "histogram_quantile(0.95, sum by (scope, le) (rate(hindsight_llm_duration_seconds_bucket[5m])))",
|
||||
"expr": "histogram_quantile(0.95, sum by (scope, le) (rate(hindsight_llm_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
|
||||
"legendFormat": "{{scope}}",
|
||||
"refId": "A"
|
||||
}
|
||||
@@ -508,12 +508,12 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum by (scope) (rate(hindsight_llm_tokens_input_tokens_total[1m]))",
|
||||
"expr": "sum by (scope) (rate(hindsight_llm_tokens_input_tokens_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "{{scope}} (input)",
|
||||
"refId": "A"
|
||||
},
|
||||
{
|
||||
"expr": "sum by (scope) (rate(hindsight_llm_tokens_output_tokens_total[1m]))",
|
||||
"expr": "sum by (scope) (rate(hindsight_llm_tokens_output_tokens_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "{{scope}} (output)",
|
||||
"refId": "B"
|
||||
}
|
||||
@@ -526,7 +526,36 @@
|
||||
"schemaVersion": 38,
|
||||
"tags": ["hindsight", "llm"],
|
||||
"templating": {
|
||||
"list": []
|
||||
"list": [
|
||||
{
|
||||
"allValue": ".*",
|
||||
"current": {
|
||||
"selected": true,
|
||||
"text": "All",
|
||||
"value": "$__all"
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "prometheus"
|
||||
},
|
||||
"definition": "label_values(hindsight_llm_calls_total, tenant)",
|
||||
"hide": 0,
|
||||
"includeAll": true,
|
||||
"label": "Tenant",
|
||||
"multi": false,
|
||||
"name": "tenant",
|
||||
"options": [],
|
||||
"query": {
|
||||
"query": "label_values(hindsight_llm_calls_total, tenant)",
|
||||
"refId": "PrometheusVariableQueryEditor-VariableQuery"
|
||||
},
|
||||
"refresh": 2,
|
||||
"regex": "",
|
||||
"skipUrlSync": false,
|
||||
"sort": 1,
|
||||
"type": "query"
|
||||
}
|
||||
]
|
||||
},
|
||||
"time": {
|
||||
"from": "now-30m",
|
||||
|
||||
@@ -46,7 +46,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(hindsight_operation_operations_total)",
|
||||
"expr": "sum(hindsight_operation_operations_total{tenant=~\"$tenant\"})",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -91,7 +91,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(rate(hindsight_operation_operations_total[1m]))",
|
||||
"expr": "sum(rate(hindsight_operation_operations_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -137,7 +137,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"retain\"}[1m]))",
|
||||
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"retain\", tenant=~\"$tenant\"}[1m]))",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -183,7 +183,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"recall\"}[1m]))",
|
||||
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"recall\", tenant=~\"$tenant\"}[1m]))",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -229,7 +229,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"reflect\"}[1m]))",
|
||||
"expr": "sum(rate(hindsight_operation_operations_total{operation=\"reflect\", tenant=~\"$tenant\"}[1m]))",
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
@@ -319,7 +319,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum by (operation) (rate(hindsight_operation_operations_total[1m]))",
|
||||
"expr": "sum by (operation) (rate(hindsight_operation_operations_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "{{operation}}",
|
||||
"refId": "A"
|
||||
}
|
||||
@@ -410,17 +410,17 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "histogram_quantile(0.50, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\"}[5m])))",
|
||||
"expr": "histogram_quantile(0.50, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\", tenant=~\"$tenant\"}[5m])))",
|
||||
"legendFormat": "p50",
|
||||
"refId": "A"
|
||||
},
|
||||
{
|
||||
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\"}[5m])))",
|
||||
"expr": "histogram_quantile(0.95, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\", tenant=~\"$tenant\"}[5m])))",
|
||||
"legendFormat": "p95",
|
||||
"refId": "B"
|
||||
},
|
||||
{
|
||||
"expr": "histogram_quantile(0.99, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\"}[5m])))",
|
||||
"expr": "histogram_quantile(0.99, sum by (le) (rate(hindsight_operation_duration_seconds_bucket{operation=\"recall\", tenant=~\"$tenant\"}[5m])))",
|
||||
"legendFormat": "p99",
|
||||
"refId": "C"
|
||||
}
|
||||
@@ -498,7 +498,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "histogram_quantile(0.95, sum by (operation, le) (rate(hindsight_operation_duration_seconds_bucket[5m])))",
|
||||
"expr": "histogram_quantile(0.95, sum by (operation, le) (rate(hindsight_operation_duration_seconds_bucket{tenant=~\"$tenant\"}[5m])))",
|
||||
"legendFormat": "{{operation}}",
|
||||
"refId": "A"
|
||||
}
|
||||
@@ -576,7 +576,7 @@
|
||||
"pluginVersion": "10.0.0",
|
||||
"targets": [
|
||||
{
|
||||
"expr": "sum by (bank_id) (rate(hindsight_operation_operations_total[1m]))",
|
||||
"expr": "sum by (bank_id) (rate(hindsight_operation_operations_total{tenant=~\"$tenant\"}[1m]))",
|
||||
"legendFormat": "{{bank_id}}",
|
||||
"refId": "A"
|
||||
}
|
||||
@@ -589,7 +589,36 @@
|
||||
"schemaVersion": 38,
|
||||
"tags": ["hindsight"],
|
||||
"templating": {
|
||||
"list": []
|
||||
"list": [
|
||||
{
|
||||
"allValue": ".*",
|
||||
"current": {
|
||||
"selected": true,
|
||||
"text": "All",
|
||||
"value": "$__all"
|
||||
},
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "prometheus"
|
||||
},
|
||||
"definition": "label_values(hindsight_operation_operations_total, tenant)",
|
||||
"hide": 0,
|
||||
"includeAll": true,
|
||||
"label": "Tenant",
|
||||
"multi": false,
|
||||
"name": "tenant",
|
||||
"options": [],
|
||||
"query": {
|
||||
"query": "label_values(hindsight_operation_operations_total, tenant)",
|
||||
"refId": "PrometheusVariableQueryEditor-VariableQuery"
|
||||
},
|
||||
"refresh": 2,
|
||||
"regex": "",
|
||||
"skipUrlSync": false,
|
||||
"sort": 1,
|
||||
"type": "query"
|
||||
}
|
||||
]
|
||||
},
|
||||
"time": {
|
||||
"from": "now-30m",
|
||||
|
||||
@@ -1257,7 +1257,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "hindsight-all"
|
||||
version = "0.2.1"
|
||||
version = "0.3.0"
|
||||
source = { editable = "hindsight" }
|
||||
dependencies = [
|
||||
{ name = "hindsight-api" },
|
||||
@@ -1281,7 +1281,7 @@ provides-extras = ["test"]
|
||||
|
||||
[[package]]
|
||||
name = "hindsight-api"
|
||||
version = "0.2.1"
|
||||
version = "0.3.0"
|
||||
source = { editable = "hindsight-api" }
|
||||
dependencies = [
|
||||
{ name = "alembic" },
|
||||
@@ -1397,7 +1397,7 @@ dev = [
|
||||
|
||||
[[package]]
|
||||
name = "hindsight-client"
|
||||
version = "0.2.1"
|
||||
version = "0.3.0"
|
||||
source = { editable = "hindsight-clients/python" }
|
||||
dependencies = [
|
||||
{ name = "aiohttp" },
|
||||
@@ -1431,7 +1431,7 @@ provides-extras = ["test"]
|
||||
|
||||
[[package]]
|
||||
name = "hindsight-dev"
|
||||
version = "0.2.1"
|
||||
version = "0.3.0"
|
||||
source = { editable = "hindsight-dev" }
|
||||
dependencies = [
|
||||
{ name = "hindsight-api" },
|
||||
@@ -1466,7 +1466,7 @@ dev = [
|
||||
|
||||
[[package]]
|
||||
name = "hindsight-embed"
|
||||
version = "0.2.1"
|
||||
version = "0.3.0"
|
||||
source = { editable = "hindsight-embed" }
|
||||
dependencies = [
|
||||
{ name = "httpx" },
|
||||
|
||||
Reference in New Issue
Block a user