Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2795e98699 | ||
|
|
4221814de7 | ||
|
|
6859b5a60e | ||
|
|
c9c949b34f | ||
|
|
70ce979fbe | ||
|
|
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)
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
[Documentation](https://hindsight.vectorize.io) • [Paper](https://arxiv.org/abs/2512.12818) • [Cookbook](https://hindsight.vectorize.io/cookbook) • [Hindsight Cloud](https://vectorize.io/hindsight/cloud)
|
||||
|
||||
[](https://github.com/vectorize-io/hindsight/actions/workflows/release.yml)
|
||||
[](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
|
||||
[](https://join.slack.com/t/hindsight-space/shared_invite/zt-3nhbm4w29-LeSJ5Ixi6j8PdiYOCPlOgg)
|
||||
[](https://opensource.org/licenses/MIT)
|
||||

|
||||

|
||||
@@ -242,7 +242,7 @@ client.reflect(bank_id="my-bank", query="What should I know about Alice?")
|
||||
- [CLI](https://hindsight.vectorize.io/sdks/cli)
|
||||
|
||||
**Community:**
|
||||
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3klo21kua-VUCC_zHP5rIcXFB1_5yw6A)
|
||||
- [Slack](https://join.slack.com/t/hindsight-space/shared_invite/zt-3nhbm4w29-LeSJ5Ixi6j8PdiYOCPlOgg)
|
||||
- [GitHub Issues](https://github.com/vectorize-io/hindsight/issues)
|
||||
|
||||
---
|
||||
|
||||
@@ -2,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(
|
||||
|
||||
@@ -41,10 +41,19 @@ ENV_EMBEDDINGS_LOCAL_MODEL = "HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL"
|
||||
ENV_EMBEDDINGS_TEI_URL = "HINDSIGHT_API_EMBEDDINGS_TEI_URL"
|
||||
ENV_EMBEDDINGS_OPENAI_API_KEY = "HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY"
|
||||
ENV_EMBEDDINGS_OPENAI_MODEL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL"
|
||||
ENV_EMBEDDINGS_OPENAI_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_OPENAI_BASE_URL"
|
||||
|
||||
ENV_COHERE_API_KEY = "HINDSIGHT_API_COHERE_API_KEY"
|
||||
ENV_EMBEDDINGS_COHERE_MODEL = "HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL"
|
||||
ENV_EMBEDDINGS_COHERE_BASE_URL = "HINDSIGHT_API_EMBEDDINGS_COHERE_BASE_URL"
|
||||
ENV_RERANKER_COHERE_MODEL = "HINDSIGHT_API_RERANKER_COHERE_MODEL"
|
||||
ENV_RERANKER_COHERE_BASE_URL = "HINDSIGHT_API_RERANKER_COHERE_BASE_URL"
|
||||
|
||||
# LiteLLM gateway configuration (for embeddings and reranker via LiteLLM proxy)
|
||||
ENV_LITELLM_API_BASE = "HINDSIGHT_API_LITELLM_API_BASE"
|
||||
ENV_LITELLM_API_KEY = "HINDSIGHT_API_LITELLM_API_KEY"
|
||||
ENV_EMBEDDINGS_LITELLM_MODEL = "HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL"
|
||||
ENV_RERANKER_LITELLM_MODEL = "HINDSIGHT_API_RERANKER_LITELLM_MODEL"
|
||||
|
||||
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
|
||||
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
|
||||
@@ -121,6 +130,11 @@ DEFAULT_RERANKER_FLASHRANK_CACHE_DIR = None # Use default cache directory
|
||||
DEFAULT_EMBEDDINGS_COHERE_MODEL = "embed-english-v3.0"
|
||||
DEFAULT_RERANKER_COHERE_MODEL = "rerank-english-v3.0"
|
||||
|
||||
# LiteLLM defaults
|
||||
DEFAULT_LITELLM_API_BASE = "http://localhost:4000"
|
||||
DEFAULT_EMBEDDINGS_LITELLM_MODEL = "text-embedding-3-small"
|
||||
DEFAULT_RERANKER_LITELLM_MODEL = "cohere/rerank-english-v3.0"
|
||||
|
||||
DEFAULT_HOST = "0.0.0.0"
|
||||
DEFAULT_PORT = 8888
|
||||
DEFAULT_LOG_LEVEL = "info"
|
||||
@@ -224,6 +238,8 @@ class HindsightConfig:
|
||||
embeddings_provider: str
|
||||
embeddings_local_model: str
|
||||
embeddings_tei_url: str | None
|
||||
embeddings_openai_base_url: str | None
|
||||
embeddings_cohere_base_url: str | None
|
||||
|
||||
# Reranker
|
||||
reranker_provider: str
|
||||
@@ -232,6 +248,7 @@ class HindsightConfig:
|
||||
reranker_tei_batch_size: int
|
||||
reranker_tei_max_concurrent: int
|
||||
reranker_max_candidates: int
|
||||
reranker_cohere_base_url: str | None
|
||||
|
||||
# Server
|
||||
host: str
|
||||
@@ -300,6 +317,8 @@ class HindsightConfig:
|
||||
embeddings_provider=os.getenv(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER),
|
||||
embeddings_local_model=os.getenv(ENV_EMBEDDINGS_LOCAL_MODEL, DEFAULT_EMBEDDINGS_LOCAL_MODEL),
|
||||
embeddings_tei_url=os.getenv(ENV_EMBEDDINGS_TEI_URL),
|
||||
embeddings_openai_base_url=os.getenv(ENV_EMBEDDINGS_OPENAI_BASE_URL) or None,
|
||||
embeddings_cohere_base_url=os.getenv(ENV_EMBEDDINGS_COHERE_BASE_URL) or None,
|
||||
# Reranker
|
||||
reranker_provider=os.getenv(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER),
|
||||
reranker_local_model=os.getenv(ENV_RERANKER_LOCAL_MODEL, DEFAULT_RERANKER_LOCAL_MODEL),
|
||||
@@ -309,6 +328,7 @@ class HindsightConfig:
|
||||
os.getenv(ENV_RERANKER_TEI_MAX_CONCURRENT, str(DEFAULT_RERANKER_TEI_MAX_CONCURRENT))
|
||||
),
|
||||
reranker_max_candidates=int(os.getenv(ENV_RERANKER_MAX_CANDIDATES, str(DEFAULT_RERANKER_MAX_CANDIDATES))),
|
||||
reranker_cohere_base_url=os.getenv(ENV_RERANKER_COHERE_BASE_URL) or None,
|
||||
# Server
|
||||
host=os.getenv(ENV_HOST, DEFAULT_HOST),
|
||||
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
|
||||
|
||||
@@ -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'"
|
||||
)
|
||||
|
||||
@@ -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'"
|
||||
)
|
||||
|
||||
@@ -151,6 +151,7 @@ from .retain import bank_utils, embedding_utils
|
||||
from .retain.types import RetainContentDict
|
||||
from .search import observation_utils, think_utils
|
||||
from .search.reranking import CrossEncoderReranker
|
||||
from .search.tags import TagsMatch
|
||||
from .task_backend import AsyncIOQueueBackend, NoopTaskBackend, TaskBackend
|
||||
|
||||
|
||||
@@ -1059,6 +1060,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
document_id: str | None = None,
|
||||
fact_type_override: str | None = None,
|
||||
confidence_score: float | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
return_usage: bool = False,
|
||||
):
|
||||
"""
|
||||
@@ -1191,6 +1193,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch=i == 1, # Only upsert on first batch
|
||||
fact_type_override=fact_type_override,
|
||||
confidence_score=confidence_score,
|
||||
document_tags=document_tags,
|
||||
)
|
||||
all_results.extend(sub_results)
|
||||
total_usage = total_usage + sub_usage
|
||||
@@ -1209,6 +1212,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch=True,
|
||||
fact_type_override=fact_type_override,
|
||||
confidence_score=confidence_score,
|
||||
document_tags=document_tags,
|
||||
)
|
||||
|
||||
# Call post-operation hook if validator is configured
|
||||
@@ -1243,6 +1247,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch: bool = True,
|
||||
fact_type_override: str | None = None,
|
||||
confidence_score: float | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
) -> tuple[list[list[str]], "TokenUsage"]:
|
||||
"""
|
||||
Internal method for batch processing without chunking logic.
|
||||
@@ -1259,6 +1264,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch: Whether this is the first batch (for chunked operations, only delete on first batch)
|
||||
fact_type_override: Override fact type for all facts
|
||||
confidence_score: Confidence score for opinions
|
||||
document_tags: Tags applied to all items in this batch
|
||||
|
||||
Returns:
|
||||
Tuple of (unit ID lists, token usage for fact extraction)
|
||||
@@ -1283,6 +1289,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
is_first_batch=is_first_batch,
|
||||
fact_type_override=fact_type_override,
|
||||
confidence_score=confidence_score,
|
||||
document_tags=document_tags,
|
||||
)
|
||||
|
||||
def recall(
|
||||
@@ -1341,6 +1348,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
include_chunks: bool = False,
|
||||
max_chunk_tokens: int = 8192,
|
||||
request_context: "RequestContext",
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> RecallResultModel:
|
||||
"""
|
||||
Recall memories using N*4-way parallel retrieval (N fact types × 4 retrieval methods).
|
||||
@@ -1366,6 +1375,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_entity_tokens: Maximum tokens for entity observations (default 500)
|
||||
include_chunks: Whether to include raw chunks in the response
|
||||
max_chunk_tokens: Maximum tokens for chunks (default 8192)
|
||||
tags: Optional list of tags for visibility filtering (OR matching - returns
|
||||
memories that have at least one matching tag)
|
||||
|
||||
Returns:
|
||||
RecallResultModel containing:
|
||||
@@ -1438,6 +1449,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_chunk_tokens,
|
||||
request_context,
|
||||
semaphore_wait=semaphore_wait,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
break # Success - exit retry loop
|
||||
except Exception as e:
|
||||
@@ -1556,6 +1569,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
max_chunk_tokens: int = 8192,
|
||||
request_context: "RequestContext" = None,
|
||||
semaphore_wait: float = 0.0,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> RecallResultModel:
|
||||
"""
|
||||
Search implementation with modular retrieval and reranking.
|
||||
@@ -1585,7 +1600,9 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
# Initialize tracer if requested
|
||||
from .search.tracer import SearchTracer
|
||||
|
||||
tracer = SearchTracer(query, thinking_budget, max_tokens) if enable_trace else None
|
||||
tracer = (
|
||||
SearchTracer(query, thinking_budget, max_tokens, tags=tags, tags_match=tags_match) if enable_trace else None
|
||||
)
|
||||
if tracer:
|
||||
tracer.start()
|
||||
|
||||
@@ -1595,8 +1612,9 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
# Buffer logs for clean output in concurrent scenarios
|
||||
recall_id = f"{bank_id[:8]}-{int(time.time() * 1000) % 100000}"
|
||||
log_buffer = []
|
||||
tags_info = f", tags={tags}, tags_match={tags_match}" if tags else ""
|
||||
log_buffer.append(
|
||||
f"[RECALL {recall_id}] Query: '{query[:50]}...' (budget={thinking_budget}, max_tokens={max_tokens})"
|
||||
f"[RECALL {recall_id}] Query: '{query[:50]}...' (budget={thinking_budget}, max_tokens={max_tokens}{tags_info})"
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1642,6 +1660,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
thinking_budget,
|
||||
question_date,
|
||||
self.query_analyzer,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
parallel_duration = time.time() - parallel_start
|
||||
|
||||
@@ -1746,6 +1766,11 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
f"edges={hd.get('edges_loaded', 0)}"
|
||||
)
|
||||
|
||||
# Record temporal constraint in tracer if detected
|
||||
if tracer and detected_temporal_constraint:
|
||||
start_dt, end_dt = detected_temporal_constraint
|
||||
tracer.record_temporal_constraint(start_dt, end_dt)
|
||||
|
||||
# Record retrieval results for tracer - per fact type
|
||||
if tracer:
|
||||
# Convert RetrievalResult to old tuple format for tracer
|
||||
@@ -1788,14 +1813,22 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
fact_type=ft_name,
|
||||
)
|
||||
|
||||
# Add temporal retrieval results for this fact type (even if empty, to show it ran)
|
||||
if rr.temporal is not None:
|
||||
# Add temporal retrieval results for this fact type
|
||||
# Show temporal even with 0 results if constraint was detected
|
||||
if rr.temporal is not None or rr.temporal_constraint is not None:
|
||||
temporal_metadata = {"budget": thinking_budget}
|
||||
if rr.temporal_constraint:
|
||||
start_dt, end_dt = rr.temporal_constraint
|
||||
temporal_metadata["constraint"] = {
|
||||
"start": start_dt.isoformat() if start_dt else None,
|
||||
"end": end_dt.isoformat() if end_dt else None,
|
||||
}
|
||||
tracer.add_retrieval_results(
|
||||
method_name="temporal",
|
||||
results=to_tuple_format(rr.temporal),
|
||||
results=to_tuple_format(rr.temporal or []),
|
||||
duration_seconds=rr.timings.get("temporal", 0.0),
|
||||
score_field="temporal_score",
|
||||
metadata={"budget": thinking_budget},
|
||||
metadata=temporal_metadata,
|
||||
fact_type=ft_name,
|
||||
)
|
||||
|
||||
@@ -2055,6 +2088,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
mentioned_at=result_dict.get("mentioned_at"),
|
||||
document_id=result_dict.get("document_id"),
|
||||
chunk_id=result_dict.get("chunk_id"),
|
||||
tags=result_dict.get("tags"),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -2270,11 +2304,11 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
doc = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT d.id, d.bank_id, d.original_text, d.content_hash,
|
||||
d.created_at, d.updated_at, COUNT(mu.id) as unit_count
|
||||
d.created_at, d.updated_at, d.tags, COUNT(mu.id) as unit_count
|
||||
FROM {fq_table("documents")} d
|
||||
LEFT JOIN {fq_table("memory_units")} mu ON mu.document_id = d.id
|
||||
WHERE d.id = $1 AND d.bank_id = $2
|
||||
GROUP BY d.id, d.bank_id, d.original_text, d.content_hash, d.created_at, d.updated_at
|
||||
GROUP BY d.id, d.bank_id, d.original_text, d.content_hash, d.created_at, d.updated_at, d.tags
|
||||
""",
|
||||
document_id,
|
||||
bank_id,
|
||||
@@ -2291,6 +2325,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
"memory_unit_count": doc["unit_count"],
|
||||
"created_at": doc["created_at"].isoformat() if doc["created_at"] else None,
|
||||
"updated_at": doc["updated_at"].isoformat() if doc["updated_at"] else None,
|
||||
"tags": list(doc["tags"]) if doc["tags"] else [],
|
||||
}
|
||||
|
||||
async def delete_document(
|
||||
@@ -2779,6 +2814,68 @@ class MemoryEngine(MemoryEngineInterface):
|
||||
|
||||
return {"items": items, "total": total, "limit": limit, "offset": offset}
|
||||
|
||||
async def get_memory_unit(
|
||||
self,
|
||||
bank_id: str,
|
||||
memory_id: str,
|
||||
request_context: "RequestContext",
|
||||
):
|
||||
"""
|
||||
Get a single memory unit by ID.
|
||||
|
||||
Args:
|
||||
bank_id: Bank ID
|
||||
memory_id: Memory unit ID
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with memory unit data or None if not found
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
pool = await self._get_pool()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Get the memory unit
|
||||
row = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, fact_type, document_id, chunk_id, tags
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = $1 AND bank_id = $2
|
||||
""",
|
||||
memory_id,
|
||||
bank_id,
|
||||
)
|
||||
|
||||
if not row:
|
||||
return None
|
||||
|
||||
# Get entity information
|
||||
entities_rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT e.canonical_name
|
||||
FROM {fq_table("unit_entities")} ue
|
||||
JOIN {fq_table("entities")} e ON ue.entity_id = e.id
|
||||
WHERE ue.unit_id = $1
|
||||
""",
|
||||
row["id"],
|
||||
)
|
||||
entities = [r["canonical_name"] for r in entities_rows]
|
||||
|
||||
return {
|
||||
"id": str(row["id"]),
|
||||
"text": row["text"],
|
||||
"context": row["context"] if row["context"] else "",
|
||||
"date": row["event_date"].isoformat() if row["event_date"] else "",
|
||||
"type": row["fact_type"],
|
||||
"mentioned_at": row["mentioned_at"].isoformat() if row["mentioned_at"] else None,
|
||||
"occurred_start": row["occurred_start"].isoformat() if row["occurred_start"] else None,
|
||||
"occurred_end": row["occurred_end"].isoformat() if row["occurred_end"] else None,
|
||||
"entities": entities,
|
||||
"document_id": row["document_id"] if row["document_id"] else None,
|
||||
"chunk_id": str(row["chunk_id"]) if row["chunk_id"] else None,
|
||||
"tags": row["tags"] if row["tags"] else [],
|
||||
}
|
||||
|
||||
async def list_documents(
|
||||
self,
|
||||
bank_id: str,
|
||||
@@ -3302,6 +3399,8 @@ Guidelines:
|
||||
max_tokens: int = 4096,
|
||||
response_schema: dict | None = None,
|
||||
request_context: "RequestContext",
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> ReflectResult:
|
||||
"""
|
||||
Reflect and formulate an answer using bank identity, world facts, and opinions.
|
||||
@@ -3369,6 +3468,8 @@ Guidelines:
|
||||
fact_type=["experience", "world", "opinion"],
|
||||
include_entities=True,
|
||||
request_context=request_context,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
recall_time = time.time() - recall_start
|
||||
|
||||
@@ -3682,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,
|
||||
@@ -3721,6 +3822,85 @@ Guidelines:
|
||||
"offset": offset,
|
||||
}
|
||||
|
||||
async def list_tags(
|
||||
self,
|
||||
bank_id: str,
|
||||
*,
|
||||
pattern: str | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
request_context: "RequestContext",
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
List all unique tags for a bank with usage counts.
|
||||
|
||||
Use this to discover available tags or expand wildcard patterns.
|
||||
Supports '*' as wildcard for flexible matching (case-insensitive):
|
||||
- 'user:*' matches user:alice, user:bob
|
||||
- '*-admin' matches role-admin, super-admin
|
||||
- 'env*-prod' matches env-prod, environment-prod
|
||||
|
||||
Args:
|
||||
bank_id: Bank identifier
|
||||
pattern: Wildcard pattern to filter tags (use '*' as wildcard, case-insensitive)
|
||||
limit: Maximum number of tags to return
|
||||
offset: Offset for pagination
|
||||
request_context: Request context for authentication.
|
||||
|
||||
Returns:
|
||||
Dict with items (list of {tag, count}), total, limit, offset
|
||||
"""
|
||||
await self._authenticate_tenant(request_context)
|
||||
pool = await self._get_pool()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
# Build pattern filter if provided (convert * to % for ILIKE)
|
||||
pattern_clause = ""
|
||||
params: list[Any] = [bank_id]
|
||||
if pattern:
|
||||
# Convert wildcard pattern: * -> % for SQL ILIKE
|
||||
sql_pattern = pattern.replace("*", "%")
|
||||
pattern_clause = "AND tag ILIKE $2"
|
||||
params.append(sql_pattern)
|
||||
|
||||
# Get total count of distinct tags matching pattern
|
||||
total_row = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT COUNT(DISTINCT tag) as total
|
||||
FROM {fq_table("memory_units")}, unnest(tags) AS tag
|
||||
WHERE bank_id = $1 AND tags IS NOT NULL AND tags != '{{}}'
|
||||
{pattern_clause}
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
total = total_row["total"] if total_row else 0
|
||||
|
||||
# Get paginated tags with counts, ordered by frequency
|
||||
limit_param = len(params) + 1
|
||||
offset_param = len(params) + 2
|
||||
params.extend([limit, offset])
|
||||
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT tag, COUNT(*) as count
|
||||
FROM {fq_table("memory_units")}, unnest(tags) AS tag
|
||||
WHERE bank_id = $1 AND tags IS NOT NULL AND tags != '{{}}'
|
||||
{pattern_clause}
|
||||
GROUP BY tag
|
||||
ORDER BY count DESC, tag ASC
|
||||
LIMIT ${limit_param} OFFSET ${offset_param}
|
||||
""",
|
||||
*params,
|
||||
)
|
||||
|
||||
items = [{"tag": row["tag"], "count": row["count"]} for row in rows]
|
||||
|
||||
return {
|
||||
"items": items,
|
||||
"total": total,
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
}
|
||||
|
||||
async def get_entity_state(
|
||||
self,
|
||||
bank_id: str,
|
||||
@@ -4365,6 +4545,7 @@ Guidelines:
|
||||
contents: list[dict[str, Any]],
|
||||
*,
|
||||
request_context: "RequestContext",
|
||||
document_tags: list[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Submit a batch retain operation to run asynchronously."""
|
||||
await self._authenticate_tenant(request_context)
|
||||
@@ -4388,14 +4569,16 @@ Guidelines:
|
||||
)
|
||||
|
||||
# Submit task to background queue
|
||||
await self._task_backend.submit_task(
|
||||
{
|
||||
"type": "batch_retain",
|
||||
"operation_id": str(operation_id),
|
||||
"bank_id": bank_id,
|
||||
"contents": contents,
|
||||
}
|
||||
)
|
||||
task_payload = {
|
||||
"type": "batch_retain",
|
||||
"operation_id": str(operation_id),
|
||||
"bank_id": bank_id,
|
||||
"contents": contents,
|
||||
}
|
||||
if document_tags:
|
||||
task_payload["document_tags"] = document_tags
|
||||
|
||||
await self._task_backend.submit_task(task_payload)
|
||||
|
||||
logger.info(f"Retain task queued for bank_id={bank_id}, {len(contents)} items, operation_id={operation_id}")
|
||||
|
||||
|
||||
@@ -85,6 +85,7 @@ class MemoryFact(BaseModel):
|
||||
"metadata": {"source": "slack"},
|
||||
"chunk_id": "bank123_session_abc123_0",
|
||||
"activation": 0.95,
|
||||
"tags": ["user_a", "session_123"],
|
||||
}
|
||||
}
|
||||
)
|
||||
@@ -102,6 +103,7 @@ class MemoryFact(BaseModel):
|
||||
chunk_id: str | None = Field(
|
||||
None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)"
|
||||
)
|
||||
tags: list[str] | None = Field(None, description="Visibility scope tags associated with this fact")
|
||||
|
||||
|
||||
class ChunkInfo(BaseModel):
|
||||
|
||||
@@ -1268,6 +1268,7 @@ async def extract_facts_from_contents(
|
||||
# mentioned_at: always the event_date (when the conversation/document occurred)
|
||||
mentioned_at=content.event_date,
|
||||
metadata=content.metadata,
|
||||
tags=content.tags,
|
||||
)
|
||||
|
||||
extracted_facts.append(extracted_fact)
|
||||
|
||||
@@ -45,6 +45,7 @@ async def insert_facts_batch(
|
||||
metadata_jsons = []
|
||||
chunk_ids = []
|
||||
document_ids = []
|
||||
tags_list = []
|
||||
|
||||
for fact in facts:
|
||||
fact_texts.append(fact.fact_text)
|
||||
@@ -65,16 +66,31 @@ async def insert_facts_batch(
|
||||
chunk_ids.append(fact.chunk_id)
|
||||
# Use per-fact document_id if available, otherwise fallback to batch-level document_id
|
||||
document_ids.append(fact.document_id if fact.document_id else document_id)
|
||||
# Convert tags to JSON string for proper batch insertion (PostgreSQL unnest doesn't handle 2D arrays well)
|
||||
tags_list.append(json.dumps(fact.tags if fact.tags else []))
|
||||
|
||||
# Batch insert all facts
|
||||
# Note: tags are passed as JSON strings and converted back to varchar[] via jsonb_array_elements_text + array_agg
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id)
|
||||
SELECT $1, * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[]
|
||||
WITH input_data AS (
|
||||
SELECT * FROM unnest(
|
||||
$2::text[], $3::vector[], $4::timestamptz[], $5::timestamptz[], $6::timestamptz[], $7::timestamptz[],
|
||||
$8::text[], $9::text[], $10::float[], $11::int[], $12::jsonb[], $13::text[], $14::text[], $15::jsonb[]
|
||||
) AS t(text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags_json)
|
||||
)
|
||||
INSERT INTO {fq_table("memory_units")} (bank_id, text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id, tags)
|
||||
SELECT
|
||||
$1,
|
||||
text, embedding, event_date, occurred_start, occurred_end, mentioned_at,
|
||||
context, fact_type, confidence_score, access_count, metadata, chunk_id, document_id,
|
||||
COALESCE(
|
||||
(SELECT array_agg(elem) FROM jsonb_array_elements_text(tags_json) AS elem),
|
||||
'{{}}'::varchar[]
|
||||
)
|
||||
FROM input_data
|
||||
RETURNING id
|
||||
""",
|
||||
bank_id,
|
||||
@@ -91,6 +107,7 @@ async def insert_facts_batch(
|
||||
metadata_jsons,
|
||||
chunk_ids,
|
||||
document_ids,
|
||||
tags_list,
|
||||
)
|
||||
|
||||
unit_ids = [str(row["id"]) for row in results]
|
||||
@@ -121,7 +138,13 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
|
||||
|
||||
|
||||
async def handle_document_tracking(
|
||||
conn, bank_id: str, document_id: str, combined_content: str, is_first_batch: bool, retain_params: dict | None = None
|
||||
conn,
|
||||
bank_id: str,
|
||||
document_id: str,
|
||||
combined_content: str,
|
||||
is_first_batch: bool,
|
||||
retain_params: dict | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handle document tracking in the database.
|
||||
@@ -133,6 +156,7 @@ async def handle_document_tracking(
|
||||
combined_content: Combined content text from all content items
|
||||
is_first_batch: Whether this is the first batch (for chunked operations)
|
||||
retain_params: Optional parameters passed during retain (context, event_date, etc.)
|
||||
document_tags: Optional list of tags to associate with the document
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
@@ -149,13 +173,14 @@ async def handle_document_tracking(
|
||||
# Insert document (or update if exists from concurrent operations)
|
||||
await conn.execute(
|
||||
f"""
|
||||
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params)
|
||||
VALUES ($1, $2, $3, $4, $5, $6)
|
||||
INSERT INTO {fq_table("documents")} (id, bank_id, original_text, content_hash, metadata, retain_params, tags)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
||||
ON CONFLICT (id, bank_id) DO UPDATE
|
||||
SET original_text = EXCLUDED.original_text,
|
||||
content_hash = EXCLUDED.content_hash,
|
||||
metadata = EXCLUDED.metadata,
|
||||
retain_params = EXCLUDED.retain_params,
|
||||
tags = EXCLUDED.tags,
|
||||
updated_at = NOW()
|
||||
""",
|
||||
document_id,
|
||||
@@ -164,4 +189,5 @@ async def handle_document_tracking(
|
||||
content_hash,
|
||||
json.dumps({}), # Empty metadata dict
|
||||
json.dumps(retain_params) if retain_params else None,
|
||||
document_tags or [],
|
||||
)
|
||||
|
||||
@@ -49,6 +49,7 @@ async def retain_batch(
|
||||
is_first_batch: bool = True,
|
||||
fact_type_override: str | None = None,
|
||||
confidence_score: float | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
) -> tuple[list[list[str]], TokenUsage]:
|
||||
"""
|
||||
Process a batch of content through the retain pipeline.
|
||||
@@ -67,6 +68,7 @@ async def retain_batch(
|
||||
is_first_batch: Whether this is the first batch
|
||||
fact_type_override: Override fact type for all facts
|
||||
confidence_score: Confidence score for opinions
|
||||
document_tags: Tags applied to all items in this batch
|
||||
|
||||
Returns:
|
||||
Tuple of (unit ID lists, token usage for fact extraction)
|
||||
@@ -88,12 +90,16 @@ async def retain_batch(
|
||||
# Convert dicts to RetainContent objects
|
||||
contents = []
|
||||
for item in contents_dicts:
|
||||
# Merge item-level tags with document-level tags
|
||||
item_tags = item.get("tags", []) or []
|
||||
merged_tags = list(set(item_tags + (document_tags or [])))
|
||||
content = RetainContent(
|
||||
content=item["content"],
|
||||
context=item.get("context", ""),
|
||||
event_date=item.get("event_date") or utcnow(),
|
||||
metadata=item.get("metadata", {}),
|
||||
entities=item.get("entities", []),
|
||||
tags=merged_tags,
|
||||
)
|
||||
contents.append(content)
|
||||
|
||||
@@ -131,7 +137,7 @@ async def retain_batch(
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
else:
|
||||
# Check for per-item document_ids
|
||||
@@ -159,7 +165,7 @@ async def retain_batch(
|
||||
if first_item.get("metadata"):
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params
|
||||
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
|
||||
total_time = time.time() - start_time
|
||||
@@ -225,7 +231,7 @@ async def retain_batch(
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params
|
||||
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, document_tags
|
||||
)
|
||||
document_ids_added.append(document_id)
|
||||
doc_id_mapping[None] = document_id # For backwards compatibility
|
||||
@@ -269,7 +275,13 @@ async def retain_batch(
|
||||
retain_params["metadata"] = first_item["metadata"]
|
||||
|
||||
await fact_storage.handle_document_tracking(
|
||||
conn, bank_id, actual_doc_id, combined_content, is_first_batch, retain_params
|
||||
conn,
|
||||
bank_id,
|
||||
actual_doc_id,
|
||||
combined_content,
|
||||
is_first_batch,
|
||||
retain_params,
|
||||
document_tags,
|
||||
)
|
||||
document_ids_added.append(actual_doc_id)
|
||||
|
||||
|
||||
@@ -21,6 +21,7 @@ class RetainContentDict(TypedDict, total=False):
|
||||
metadata: Custom key-value metadata (optional)
|
||||
document_id: Document ID for this content item (optional)
|
||||
entities: User-provided entities to merge with extracted entities (optional)
|
||||
tags: Visibility scope tags for this content item (optional)
|
||||
"""
|
||||
|
||||
content: str # Required
|
||||
@@ -29,6 +30,7 @@ class RetainContentDict(TypedDict, total=False):
|
||||
metadata: dict[str, str]
|
||||
document_id: str
|
||||
entities: list[dict[str, str]] # [{"text": "...", "type": "..."}]
|
||||
tags: list[str] # Visibility scope tags
|
||||
|
||||
|
||||
def _now_utc() -> datetime:
|
||||
@@ -49,6 +51,7 @@ class RetainContent:
|
||||
event_date: datetime = field(default_factory=_now_utc)
|
||||
metadata: dict[str, str] = field(default_factory=dict)
|
||||
entities: list[dict[str, str]] = field(default_factory=list) # User-provided entities
|
||||
tags: list[str] = field(default_factory=list) # Visibility scope tags
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -113,6 +116,7 @@ class ExtractedFact:
|
||||
context: str = ""
|
||||
mentioned_at: datetime | None = None
|
||||
metadata: dict[str, str] = field(default_factory=dict)
|
||||
tags: list[str] = field(default_factory=list) # Visibility scope tags
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -158,6 +162,9 @@ class ProcessedFact:
|
||||
# Track which content this fact came from (for user entity merging)
|
||||
content_index: int = 0
|
||||
|
||||
# Visibility scope tags
|
||||
tags: list[str] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def is_duplicate(self) -> bool:
|
||||
"""Check if this fact was marked as a duplicate."""
|
||||
@@ -201,6 +208,7 @@ class ProcessedFact:
|
||||
causal_relations=extracted_fact.causal_relations,
|
||||
chunk_id=chunk_id,
|
||||
content_index=extracted_fact.content_index,
|
||||
tags=extracted_fact.tags,
|
||||
)
|
||||
|
||||
|
||||
@@ -232,6 +240,7 @@ class RetainBatch:
|
||||
document_id: str | None = None
|
||||
fact_type_override: str | None = None
|
||||
confidence_score: float | None = None
|
||||
document_tags: list[str] = field(default_factory=list) # Tags applied to all items
|
||||
|
||||
# Extracted data (populated during processing)
|
||||
extracted_facts: list[ExtractedFact] = field(default_factory=list)
|
||||
|
||||
@@ -11,6 +11,7 @@ from abc import ABC, abstractmethod
|
||||
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .tags import TagsMatch, filter_results_by_tags
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -43,6 +44,8 @@ class GraphRetriever(ABC):
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # TypedAdjacency, optional pre-loaded graph
|
||||
tags: list[str] | None = None, # Visibility scope tags for filtering
|
||||
tags_match: TagsMatch = "any", # How to match tags: 'any' (OR) or 'all' (AND)
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve relevant facts via graph traversal.
|
||||
@@ -57,6 +60,7 @@ class GraphRetriever(ABC):
|
||||
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
|
||||
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
|
||||
adjacency: Pre-loaded typed adjacency graph (optional, for MPFP)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
Tuple of (List of RetrievalResult with activation scores, optional timing info)
|
||||
@@ -114,6 +118,8 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # Not used by BFS
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts using BFS spreading activation.
|
||||
@@ -129,7 +135,9 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
for interface compatibility but not used.
|
||||
"""
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
results = await self._retrieve_with_conn(conn, query_embedding_str, bank_id, fact_type, budget)
|
||||
results = await self._retrieve_with_conn(
|
||||
conn, query_embedding_str, bank_id, fact_type, budget, tags=tags, tags_match=tags_match
|
||||
)
|
||||
return results, None
|
||||
|
||||
async def _retrieve_with_conn(
|
||||
@@ -139,33 +147,46 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
budget: int,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[RetrievalResult]:
|
||||
"""Internal implementation with connection."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, self.entry_point_threshold, self.entry_point_limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
# Step 1: Find entry points
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
self.entry_point_threshold,
|
||||
self.entry_point_limit,
|
||||
*params,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
logger.debug(
|
||||
f"[BFS] No entry points found for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
return []
|
||||
|
||||
logger.debug(
|
||||
f"[BFS] Found {len(entry_points)} entry points for fact_type={fact_type} "
|
||||
f"(tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Step 2: BFS spreading activation
|
||||
visited = set()
|
||||
results = []
|
||||
@@ -196,7 +217,7 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.occurred_start, mu.occurred_end,
|
||||
mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type,
|
||||
mu.document_id, mu.chunk_id,
|
||||
mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight, ml.link_type, ml.from_unit_id
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||
@@ -236,4 +257,8 @@ class BFSGraphRetriever(GraphRetriever):
|
||||
neighbor_result = RetrievalResult.from_db_row(dict(n))
|
||||
queue.append((neighbor_result, new_activation))
|
||||
|
||||
# Apply tags filtering (BFS may traverse into memories that don't match tags criteria)
|
||||
if tags:
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
return results
|
||||
|
||||
@@ -18,6 +18,7 @@ import time
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .graph_retrieval import GraphRetriever
|
||||
from .tags import TagsMatch, filter_results_by_tags
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -30,26 +31,32 @@ async def _find_semantic_seeds(
|
||||
fact_type: str,
|
||||
limit: int = 20,
|
||||
threshold: float = 0.3,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[RetrievalResult]:
|
||||
"""Find semantic seeds via embedding search."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
threshold,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
|
||||
|
||||
@@ -95,6 +102,8 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts by expanding links from seeds.
|
||||
@@ -109,6 +118,7 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
semantic_seeds: Pre-computed semantic entry points
|
||||
temporal_seeds: Pre-computed temporal entry points
|
||||
adjacency: Unused, kept for interface compatibility
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
Tuple of (results, timings)
|
||||
@@ -125,15 +135,27 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
else:
|
||||
seeds_start = time.time()
|
||||
all_seeds = await _find_semantic_seeds(
|
||||
conn, query_embedding_str, bank_id, fact_type, limit=20, threshold=0.3
|
||||
conn,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit=20,
|
||||
threshold=0.3,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
timings.seeds_time = time.time() - seeds_start
|
||||
logger.debug(
|
||||
f"[LinkExpansion] Found {len(all_seeds)} semantic seeds for fact_type={fact_type} "
|
||||
f"(tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Add temporal seeds if provided
|
||||
if temporal_seeds:
|
||||
all_seeds.extend(temporal_seeds)
|
||||
|
||||
if not all_seeds:
|
||||
logger.debug("[LinkExpansion] No seeds found, returning empty results")
|
||||
return [], timings
|
||||
|
||||
seed_ids = list({s.id for s in all_seeds})
|
||||
@@ -147,7 +169,7 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
SELECT
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
COUNT(*)::float AS score
|
||||
FROM {fq_table("unit_entities")} seed_ue
|
||||
JOIN {fq_table("entities")} e ON seed_ue.entity_id = e.id
|
||||
@@ -172,7 +194,7 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
SELECT DISTINCT ON (mu.id)
|
||||
mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start,
|
||||
mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id,
|
||||
mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight + 1.0 AS score
|
||||
FROM {fq_table("memory_links")} ml
|
||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||
@@ -219,6 +241,10 @@ class LinkExpansionRetriever(GraphRetriever):
|
||||
result.activation = row["score"]
|
||||
results.append(result)
|
||||
|
||||
# Apply tags filtering (graph expansion may reach untagged memories)
|
||||
if tags:
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
timings.result_count = len(results)
|
||||
timings.traverse = time.time() - start_time
|
||||
|
||||
|
||||
@@ -23,6 +23,7 @@ from dataclasses import dataclass, field
|
||||
from ..db_utils import acquire_with_retry
|
||||
from ..memory_engine import fq_table
|
||||
from .graph_retrieval import GraphRetriever
|
||||
from .tags import TagsMatch
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -448,7 +449,7 @@ async def fetch_memory_units_by_ids(
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end,
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id
|
||||
mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE id = ANY($1::uuid[])
|
||||
AND fact_type = $2
|
||||
@@ -503,6 +504,8 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
semantic_seeds: list[RetrievalResult] | None = None,
|
||||
temporal_seeds: list[RetrievalResult] | None = None,
|
||||
adjacency=None, # Ignored - kept for interface compatibility
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||
"""
|
||||
Retrieve facts using MPFP algorithm with lazy edge loading.
|
||||
@@ -517,6 +520,7 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
semantic_seeds: Pre-computed semantic entry points
|
||||
temporal_seeds: Pre-computed temporal entry points
|
||||
adjacency: Ignored (kept for interface compatibility)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
Tuple of (List of RetrievalResult with activation scores, MPFPTimings)
|
||||
@@ -532,8 +536,13 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
# If no semantic seeds provided, fall back to finding our own
|
||||
if not semantic_seed_nodes:
|
||||
seeds_start = time.time()
|
||||
semantic_seed_nodes = await self._find_semantic_seeds(pool, query_embedding_str, bank_id, fact_type)
|
||||
semantic_seed_nodes = await self._find_semantic_seeds(
|
||||
pool, query_embedding_str, bank_id, fact_type, tags=tags, tags_match=tags_match
|
||||
)
|
||||
timings.seeds_time = time.time() - seeds_start
|
||||
logger.debug(
|
||||
f"[MPFP] Found {len(semantic_seed_nodes)} semantic seeds for fact_type={fact_type} (tags={tags}, tags_match={tags_match})"
|
||||
)
|
||||
|
||||
# Collect all pattern jobs
|
||||
pattern_jobs = []
|
||||
@@ -549,6 +558,9 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
pattern_jobs.append((temporal_seed_nodes, pattern))
|
||||
|
||||
if not pattern_jobs:
|
||||
logger.debug(
|
||||
f"[MPFP] No pattern jobs (semantic_seeds={len(semantic_seed_nodes)}, temporal_seeds={len(temporal_seed_nodes)})"
|
||||
)
|
||||
return [], timings
|
||||
|
||||
timings.pattern_count = len(pattern_jobs)
|
||||
@@ -587,6 +599,7 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
timings.fusion = time.time() - step_start
|
||||
|
||||
if not fused:
|
||||
logger.debug(f"[MPFP] No fused results after RRF fusion (pattern_count={len(pattern_results)})")
|
||||
return [], timings
|
||||
|
||||
# Get top result IDs
|
||||
@@ -596,6 +609,13 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
step_start = time.time()
|
||||
results = await fetch_memory_units_by_ids(pool, result_ids, fact_type)
|
||||
timings.fetch = time.time() - step_start
|
||||
|
||||
# Filter results by tags (graph traversal may have picked up unfiltered memories)
|
||||
if tags:
|
||||
from .tags import filter_results_by_tags
|
||||
|
||||
results = filter_results_by_tags(results, tags, match=tags_match)
|
||||
|
||||
timings.result_count = len(results)
|
||||
|
||||
# Add activation scores from fusion
|
||||
@@ -634,8 +654,17 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
fact_type: str,
|
||||
limit: int = 20,
|
||||
threshold: float = 0.3,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> list[SeedNode]:
|
||||
"""Fallback: find semantic seeds via embedding search."""
|
||||
from .tags import build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_embedding_str, bank_id, fact_type, threshold, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
@@ -645,14 +674,11 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= $4
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $5
|
||||
""",
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
threshold,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
|
||||
return [SeedNode(node_id=str(r["id"]), score=r["similarity"]) for r in rows]
|
||||
|
||||
@@ -20,6 +20,7 @@ from ..memory_engine import fq_table
|
||||
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
|
||||
from .link_expansion_retrieval import LinkExpansionRetriever
|
||||
from .mpfp_retrieval import MPFPGraphRetriever
|
||||
from .tags import TagsMatch, build_tags_where_clause_simple
|
||||
from .types import MPFPTimings, RetrievalResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -85,7 +86,12 @@ def set_default_graph_retriever(retriever: GraphRetriever) -> None:
|
||||
|
||||
|
||||
async def retrieve_semantic(
|
||||
conn, query_emb_str: str, bank_id: str, fact_type: str, limit: int
|
||||
conn,
|
||||
query_emb_str: str,
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
limit: int,
|
||||
tags: list[str] | None = None,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Semantic retrieval via vector similarity.
|
||||
@@ -96,31 +102,44 @@ async def retrieve_semantic(
|
||||
agent_id: bank ID
|
||||
fact_type: Fact type to filter
|
||||
limit: Maximum results to return
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
List of RetrievalResult objects
|
||||
"""
|
||||
from .tags import TagsMatch, build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 5)
|
||||
params = [query_emb_str, bank_id, fact_type, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = $3
|
||||
AND (1 - (embedding <=> $1::vector)) >= 0.3
|
||||
{tags_clause}
|
||||
ORDER BY embedding <=> $1::vector
|
||||
LIMIT $4
|
||||
""",
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in results]
|
||||
|
||||
|
||||
async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, limit: int) -> list[RetrievalResult]:
|
||||
async def retrieve_bm25(
|
||||
conn,
|
||||
query_text: str,
|
||||
bank_id: str,
|
||||
fact_type: str,
|
||||
limit: int,
|
||||
tags: list[str] | None = None,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
BM25 keyword retrieval via full-text search.
|
||||
|
||||
@@ -130,12 +149,15 @@ async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, lim
|
||||
agent_id: bank ID
|
||||
fact_type: Fact type to filter
|
||||
limit: Maximum results to return
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
List of RetrievalResult objects
|
||||
"""
|
||||
import re
|
||||
|
||||
from .tags import TagsMatch, build_tags_where_clause_simple
|
||||
|
||||
# Sanitize query text: remove special characters that have meaning in tsquery
|
||||
# Keep only alphanumeric characters and spaces
|
||||
sanitized_text = re.sub(r"[^\w\s]", " ", query_text.lower())
|
||||
@@ -151,21 +173,24 @@ async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, lim
|
||||
# This prevents empty results when some terms are missing
|
||||
query_tsquery = " | ".join(tokens)
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 5)
|
||||
params = [query_tsquery, bank_id, fact_type, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
ts_rank_cd(search_vector, to_tsquery('english', $1)) AS bm25_score
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
AND fact_type = $3
|
||||
AND search_vector @@ to_tsquery('english', $1)
|
||||
{tags_clause}
|
||||
ORDER BY bm25_score DESC
|
||||
LIMIT $4
|
||||
""",
|
||||
query_tsquery,
|
||||
bank_id,
|
||||
fact_type,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
return [RetrievalResult.from_db_row(dict(r)) for r in results]
|
||||
|
||||
@@ -177,6 +202,8 @@ async def retrieve_semantic_bm25_combined(
|
||||
bank_id: str,
|
||||
fact_types: list[str],
|
||||
limit: int,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]]:
|
||||
"""
|
||||
Combined semantic + BM25 retrieval for multiple fact types in a single query.
|
||||
@@ -203,10 +230,14 @@ async def retrieve_semantic_bm25_combined(
|
||||
|
||||
# If no valid tokens for BM25, just run semantic
|
||||
if not tokens:
|
||||
tags_clause = build_tags_where_clause_simple(tags, 5, match=tags_match)
|
||||
params = [query_emb_str, bank_id, fact_types, limit]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
WITH semantic_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
NULL::float AS bm25_score,
|
||||
'semantic' AS source,
|
||||
@@ -216,16 +247,14 @@ async def retrieve_semantic_bm25_combined(
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = ANY($3)
|
||||
AND (1 - (embedding <=> $1::vector)) >= 0.3
|
||||
{tags_clause}
|
||||
)
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM semantic_ranked
|
||||
WHERE rn <= $4
|
||||
""",
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_types,
|
||||
limit,
|
||||
*params,
|
||||
)
|
||||
# Group by fact_type
|
||||
result_dict: dict[str, tuple[list[RetrievalResult], list[RetrievalResult]]] = {
|
||||
@@ -241,12 +270,18 @@ async def retrieve_semantic_bm25_combined(
|
||||
|
||||
query_tsquery = " | ".join(tokens)
|
||||
|
||||
# Build tags clause - param 6 if tags provided
|
||||
tags_clause = build_tags_where_clause_simple(tags, 6, match=tags_match)
|
||||
params = [query_emb_str, bank_id, fact_types, limit, query_tsquery]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
# Combined CTE query for both semantic and BM25 across all fact types
|
||||
# Uses window functions to limit per fact_type per method
|
||||
results = await conn.fetch(
|
||||
f"""
|
||||
WITH semantic_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
NULL::float AS bm25_score,
|
||||
'semantic' AS source,
|
||||
@@ -256,9 +291,10 @@ async def retrieve_semantic_bm25_combined(
|
||||
AND embedding IS NOT NULL
|
||||
AND fact_type = ANY($3)
|
||||
AND (1 - (embedding <=> $1::vector)) >= 0.3
|
||||
{tags_clause}
|
||||
),
|
||||
bm25_ranked AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
NULL::float AS similarity,
|
||||
ts_rank_cd(search_vector, to_tsquery('english', $5)) AS bm25_score,
|
||||
'bm25' AS source,
|
||||
@@ -267,14 +303,15 @@ async def retrieve_semantic_bm25_combined(
|
||||
WHERE bank_id = $2
|
||||
AND fact_type = ANY($3)
|
||||
AND search_vector @@ to_tsquery('english', $5)
|
||||
{tags_clause}
|
||||
),
|
||||
semantic AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM semantic_ranked WHERE rn <= $4
|
||||
),
|
||||
bm25 AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
similarity, bm25_score, source
|
||||
FROM bm25_ranked WHERE rn <= $4
|
||||
)
|
||||
@@ -282,11 +319,7 @@ async def retrieve_semantic_bm25_combined(
|
||||
UNION ALL
|
||||
SELECT * FROM bm25
|
||||
""",
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_types,
|
||||
limit,
|
||||
query_tsquery,
|
||||
*params,
|
||||
)
|
||||
|
||||
# Group results by fact_type and source
|
||||
@@ -313,6 +346,8 @@ async def retrieve_temporal_combined(
|
||||
end_date: datetime,
|
||||
budget: int,
|
||||
semantic_threshold: float = 0.1,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> dict[str, list[RetrievalResult]]:
|
||||
"""
|
||||
Temporal retrieval for multiple fact types in a single query.
|
||||
@@ -341,11 +376,17 @@ async def retrieve_temporal_combined(
|
||||
if end_date.tzinfo is None:
|
||||
end_date = end_date.replace(tzinfo=UTC)
|
||||
|
||||
# Build tags clause
|
||||
tags_clause = build_tags_where_clause_simple(tags, 7, match=tags_match)
|
||||
params = [query_emb_str, bank_id, fact_types, start_date, end_date, semantic_threshold]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
# Batch query: Get entry points for ALL fact types at once with window function
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
WITH ranked_entries AS (
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity,
|
||||
ROW_NUMBER() OVER (PARTITION BY fact_type ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, embedding <=> $1::vector) AS rn
|
||||
FROM {fq_table("memory_units")}
|
||||
@@ -363,17 +404,13 @@ async def retrieve_temporal_combined(
|
||||
(occurred_end IS NOT NULL AND occurred_end BETWEEN $4 AND $5)
|
||||
)
|
||||
AND (1 - (embedding <=> $1::vector)) >= $6
|
||||
{tags_clause}
|
||||
)
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, similarity
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags, similarity
|
||||
FROM ranked_entries
|
||||
WHERE rn <= 10
|
||||
""",
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_types,
|
||||
start_date,
|
||||
end_date,
|
||||
semantic_threshold,
|
||||
*params,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
@@ -436,13 +473,20 @@ async def retrieve_temporal_combined(
|
||||
budget_remaining = budget - len(ft_entry_points)
|
||||
batch_size = 20
|
||||
|
||||
# Build tags clause for spreading (use param 6 since 1-5 are used)
|
||||
spreading_tags_clause = build_tags_where_clause_simple(tags, 6, table_alias="mu.", match=tags_match)
|
||||
|
||||
while frontier and budget_remaining > 0:
|
||||
batch_ids = frontier[:batch_size]
|
||||
frontier = frontier[batch_size:]
|
||||
|
||||
spreading_params = [query_emb_str, batch_ids, ft, semantic_threshold, batch_size * 10]
|
||||
if tags:
|
||||
spreading_params.append(tags)
|
||||
|
||||
neighbors = await conn.fetch(
|
||||
f"""
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id,
|
||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id, mu.tags,
|
||||
ml.weight, ml.link_type, ml.from_unit_id,
|
||||
1 - (mu.embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_links")} ml
|
||||
@@ -453,14 +497,11 @@ async def retrieve_temporal_combined(
|
||||
AND mu.fact_type = $3
|
||||
AND mu.embedding IS NOT NULL
|
||||
AND (1 - (mu.embedding <=> $1::vector)) >= $4
|
||||
{spreading_tags_clause}
|
||||
ORDER BY ml.weight DESC
|
||||
LIMIT $5
|
||||
""",
|
||||
query_emb_str,
|
||||
batch_ids,
|
||||
ft,
|
||||
semantic_threshold,
|
||||
batch_size * 10,
|
||||
*spreading_params,
|
||||
)
|
||||
|
||||
for n in neighbors:
|
||||
@@ -529,6 +570,7 @@ async def retrieve_temporal(
|
||||
end_date: datetime,
|
||||
budget: int,
|
||||
semantic_threshold: float = 0.1,
|
||||
tags: list[str] | None = None,
|
||||
) -> list[RetrievalResult]:
|
||||
"""
|
||||
Temporal retrieval with spreading activation.
|
||||
@@ -547,6 +589,7 @@ async def retrieve_temporal(
|
||||
end_date: End of time range
|
||||
budget: Node budget for spreading
|
||||
semantic_threshold: Minimum semantic similarity to include
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
List of RetrievalResult objects with temporal scores
|
||||
@@ -558,9 +601,16 @@ async def retrieve_temporal(
|
||||
if end_date.tzinfo is None:
|
||||
end_date = end_date.replace(tzinfo=UTC)
|
||||
|
||||
from .tags import TagsMatch, build_tags_where_clause_simple
|
||||
|
||||
tags_clause = build_tags_where_clause_simple(tags, 7)
|
||||
params = [query_emb_str, bank_id, fact_type, start_date, end_date, semantic_threshold]
|
||||
if tags:
|
||||
params.append(tags)
|
||||
|
||||
entry_points = await conn.fetch(
|
||||
f"""
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id,
|
||||
SELECT id, text, context, event_date, occurred_start, occurred_end, mentioned_at, access_count, embedding, fact_type, document_id, chunk_id, tags,
|
||||
1 - (embedding <=> $1::vector) AS similarity
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $2
|
||||
@@ -580,15 +630,11 @@ async def retrieve_temporal(
|
||||
(occurred_end IS NOT NULL AND occurred_end BETWEEN $4 AND $5)
|
||||
)
|
||||
AND (1 - (embedding <=> $1::vector)) >= $6
|
||||
{tags_clause}
|
||||
ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, (embedding <=> $1::vector) ASC
|
||||
LIMIT 10
|
||||
""",
|
||||
query_emb_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
start_date,
|
||||
end_date,
|
||||
semantic_threshold,
|
||||
*params,
|
||||
)
|
||||
|
||||
if not entry_points:
|
||||
@@ -740,6 +786,7 @@ async def retrieve_parallel(
|
||||
query_analyzer: Optional["QueryAnalyzer"] = None,
|
||||
graph_retriever: GraphRetriever | None = None,
|
||||
temporal_constraint: tuple | None = None, # Pre-extracted temporal constraint
|
||||
tags: list[str] | None = None, # Visibility scope tags for filtering
|
||||
) -> ParallelRetrievalResult:
|
||||
"""
|
||||
Run 3-way or 4-way parallel retrieval (adds temporal if detected).
|
||||
@@ -755,6 +802,7 @@ async def retrieve_parallel(
|
||||
query_analyzer: Query analyzer to use (defaults to TransformerQueryAnalyzer)
|
||||
graph_retriever: Graph retrieval strategy (defaults to configured retriever)
|
||||
temporal_constraint: Pre-extracted temporal constraint (optional)
|
||||
tags: Optional list of tags for visibility filtering (OR matching)
|
||||
|
||||
Returns:
|
||||
ParallelRetrievalResult with semantic, bm25, graph, temporal results and timings
|
||||
@@ -775,6 +823,7 @@ async def retrieve_parallel(
|
||||
retriever,
|
||||
question_date,
|
||||
query_analyzer,
|
||||
tags=tags,
|
||||
)
|
||||
else:
|
||||
# For BFS, extract temporal constraint upfront (legacy path)
|
||||
@@ -785,7 +834,15 @@ async def retrieve_parallel(
|
||||
query_text, reference_date=question_date, analyzer=query_analyzer
|
||||
)
|
||||
return await _retrieve_parallel_bfs(
|
||||
pool, query_text, query_embedding_str, bank_id, fact_type, thinking_budget, temporal_constraint, retriever
|
||||
pool,
|
||||
query_text,
|
||||
query_embedding_str,
|
||||
bank_id,
|
||||
fact_type,
|
||||
thinking_budget,
|
||||
temporal_constraint,
|
||||
retriever,
|
||||
tags=tags,
|
||||
)
|
||||
|
||||
|
||||
@@ -809,6 +866,7 @@ async def _retrieve_parallel_mpfp(
|
||||
retriever: GraphRetriever,
|
||||
question_date: datetime | None = None,
|
||||
query_analyzer=None,
|
||||
tags: list[str] | None = None,
|
||||
) -> ParallelRetrievalResult:
|
||||
"""
|
||||
MPFP retrieval with true parallelization.
|
||||
@@ -830,7 +888,9 @@ async def _retrieve_parallel_mpfp(
|
||||
acquire_start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
conn_wait = time.time() - acquire_start
|
||||
results = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
|
||||
results = await retrieve_semantic(
|
||||
conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget, tags=tags
|
||||
)
|
||||
return _TimedResult(results, time.time() - start, conn_wait)
|
||||
|
||||
async def run_bm25() -> _TimedResult:
|
||||
@@ -839,7 +899,7 @@ async def _retrieve_parallel_mpfp(
|
||||
acquire_start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
conn_wait = time.time() - acquire_start
|
||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget)
|
||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget, tags=tags)
|
||||
return _TimedResult(results, time.time() - start, conn_wait)
|
||||
|
||||
async def run_graph() -> tuple[list[RetrievalResult], float, MPFPTimings | None]:
|
||||
@@ -857,6 +917,7 @@ async def _retrieve_parallel_mpfp(
|
||||
query_text=query_text,
|
||||
semantic_seeds=None, # Let MPFP find its own seeds
|
||||
temporal_seeds=None, # Don't wait for temporal extraction
|
||||
tags=tags,
|
||||
)
|
||||
return results, time.time() - start, mpfp_timing
|
||||
|
||||
@@ -1028,6 +1089,7 @@ async def _retrieve_parallel_bfs(
|
||||
thinking_budget: int,
|
||||
temporal_constraint: tuple | None,
|
||||
retriever: GraphRetriever,
|
||||
tags: list[str] | None = None,
|
||||
) -> ParallelRetrievalResult:
|
||||
"""BFS retrieval: all methods run in parallel (original behavior)."""
|
||||
import time
|
||||
@@ -1035,13 +1097,15 @@ async def _retrieve_parallel_bfs(
|
||||
async def run_semantic() -> _TimedResult:
|
||||
start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
results = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
|
||||
results = await retrieve_semantic(
|
||||
conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget, tags=tags
|
||||
)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
async def run_bm25() -> _TimedResult:
|
||||
start = time.time()
|
||||
async with acquire_with_retry(pool) as conn:
|
||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget)
|
||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget, tags=tags)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
async def run_graph() -> _TimedResult:
|
||||
@@ -1053,6 +1117,7 @@ async def _retrieve_parallel_bfs(
|
||||
fact_type=fact_type,
|
||||
budget=thinking_budget,
|
||||
query_text=query_text,
|
||||
tags=tags,
|
||||
)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
@@ -1068,6 +1133,7 @@ async def _retrieve_parallel_bfs(
|
||||
tc_end,
|
||||
budget=thinking_budget,
|
||||
semantic_threshold=0.1,
|
||||
tags=tags,
|
||||
)
|
||||
return _TimedResult(results, time.time() - start)
|
||||
|
||||
@@ -1122,6 +1188,8 @@ async def retrieve_all_fact_types_parallel(
|
||||
question_date: datetime | None = None,
|
||||
query_analyzer: Optional["QueryAnalyzer"] = None,
|
||||
graph_retriever: GraphRetriever | None = None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: TagsMatch = "any",
|
||||
) -> MultiFactTypeRetrievalResult:
|
||||
"""
|
||||
Optimized retrieval for multiple fact types using batched queries.
|
||||
@@ -1171,7 +1239,14 @@ async def retrieve_all_fact_types_parallel(
|
||||
|
||||
# Semantic + BM25 combined
|
||||
semantic_bm25_results = await retrieve_semantic_bm25_combined(
|
||||
conn, query_embedding_str, query_text, bank_id, fact_types, thinking_budget
|
||||
conn,
|
||||
query_embedding_str,
|
||||
query_text,
|
||||
bank_id,
|
||||
fact_types,
|
||||
thinking_budget,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
semantic_bm25_time = time.time() - semantic_bm25_start
|
||||
|
||||
@@ -1188,6 +1263,8 @@ async def retrieve_all_fact_types_parallel(
|
||||
tc_end,
|
||||
budget=thinking_budget,
|
||||
semantic_threshold=0.1,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
temporal_time = time.time() - temporal_start
|
||||
|
||||
@@ -1206,6 +1283,8 @@ async def retrieve_all_fact_types_parallel(
|
||||
query_text=query_text,
|
||||
semantic_seeds=None,
|
||||
temporal_seeds=None,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
return ft, results, time.time() - graph_start, mpfp_timing
|
||||
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
"""
|
||||
Tags filtering utilities for retrieval.
|
||||
|
||||
Provides SQL building functions for filtering memories by tags.
|
||||
Supports four matching modes via TagsMatch enum:
|
||||
- "any": OR matching, includes untagged memories (default, backward compatible)
|
||||
- "all": AND matching, includes untagged memories
|
||||
- "any_strict": OR matching, excludes untagged memories
|
||||
- "all_strict": AND matching, excludes untagged memories
|
||||
|
||||
OR matching (any/any_strict): Memory matches if ANY of its tags overlap with request tags
|
||||
AND matching (all/all_strict): Memory matches if ALL request tags are present in its tags
|
||||
"""
|
||||
|
||||
from typing import Literal
|
||||
|
||||
TagsMatch = Literal["any", "all", "any_strict", "all_strict"]
|
||||
|
||||
|
||||
def _parse_tags_match(match: TagsMatch) -> tuple[str, bool]:
|
||||
"""
|
||||
Parse TagsMatch into operator and include_untagged flag.
|
||||
|
||||
Returns:
|
||||
Tuple of (operator, include_untagged)
|
||||
- operator: "&&" for any/any_strict, "@>" for all/all_strict
|
||||
- include_untagged: True for any/all, False for any_strict/all_strict
|
||||
"""
|
||||
if match == "any":
|
||||
return "&&", True
|
||||
elif match == "all":
|
||||
return "@>", True
|
||||
elif match == "any_strict":
|
||||
return "&&", False
|
||||
elif match == "all_strict":
|
||||
return "@>", False
|
||||
else:
|
||||
# Default to "any" behavior
|
||||
return "&&", True
|
||||
|
||||
|
||||
def build_tags_where_clause(
|
||||
tags: list[str] | None,
|
||||
param_offset: int = 1,
|
||||
table_alias: str = "",
|
||||
match: TagsMatch = "any",
|
||||
) -> tuple[str, list, int]:
|
||||
"""
|
||||
Build a SQL WHERE clause for filtering by tags.
|
||||
|
||||
Supports four matching modes:
|
||||
- "any" (default): OR matching, includes untagged memories
|
||||
- "all": AND matching, includes untagged memories
|
||||
- "any_strict": OR matching, excludes untagged memories
|
||||
- "all_strict": AND matching, excludes untagged memories
|
||||
|
||||
Args:
|
||||
tags: List of tags to filter by. If None or empty, returns empty clause (no filtering).
|
||||
param_offset: Starting parameter number for SQL placeholders (default 1).
|
||||
table_alias: Optional table alias prefix (e.g., "mu." for "memory_units mu").
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
Tuple of (sql_clause, params, next_param_offset):
|
||||
- sql_clause: SQL WHERE clause string
|
||||
- params: List of parameter values to bind
|
||||
- next_param_offset: Next available parameter number
|
||||
|
||||
Example:
|
||||
>>> clause, params, next_offset = build_tags_where_clause(['user_a'], 3, 'mu.', 'any_strict')
|
||||
>>> print(clause) # "AND mu.tags IS NOT NULL AND mu.tags != '{}' AND mu.tags && $3"
|
||||
"""
|
||||
if not tags:
|
||||
return "", [], param_offset
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(match)
|
||||
|
||||
if include_untagged:
|
||||
# Include untagged memories (NULL or empty array) OR matching tags
|
||||
clause = f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_offset})"
|
||||
else:
|
||||
# Strict: only memories with matching tags (exclude NULL and empty)
|
||||
clause = f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_offset}"
|
||||
|
||||
return clause, [tags], param_offset + 1
|
||||
|
||||
|
||||
def build_tags_where_clause_simple(
|
||||
tags: list[str] | None,
|
||||
param_num: int,
|
||||
table_alias: str = "",
|
||||
match: TagsMatch = "any",
|
||||
) -> str:
|
||||
"""
|
||||
Build a simple SQL WHERE clause for tags filtering.
|
||||
|
||||
This is a convenience version that returns just the clause string,
|
||||
assuming the caller will add the tags array to their params list.
|
||||
|
||||
Args:
|
||||
tags: List of tags to filter by. If None or empty, returns empty string.
|
||||
param_num: Parameter number to use in the clause.
|
||||
table_alias: Optional table alias prefix.
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
SQL clause string or empty string.
|
||||
"""
|
||||
if not tags:
|
||||
return ""
|
||||
|
||||
column = f"{table_alias}tags" if table_alias else "tags"
|
||||
operator, include_untagged = _parse_tags_match(match)
|
||||
|
||||
if include_untagged:
|
||||
# Include untagged memories (NULL or empty array) OR matching tags
|
||||
return f"AND ({column} IS NULL OR {column} = '{{}}' OR {column} {operator} ${param_num})"
|
||||
else:
|
||||
# Strict: only memories with matching tags (exclude NULL and empty)
|
||||
return f"AND {column} IS NOT NULL AND {column} != '{{}}' AND {column} {operator} ${param_num}"
|
||||
|
||||
|
||||
def filter_results_by_tags(
|
||||
results: list,
|
||||
tags: list[str] | None,
|
||||
match: TagsMatch = "any",
|
||||
) -> list:
|
||||
"""
|
||||
Filter retrieval results by tags in Python (for post-processing).
|
||||
|
||||
Used when SQL filtering isn't possible (e.g., graph traversal results).
|
||||
|
||||
Args:
|
||||
results: List of RetrievalResult objects with a 'tags' attribute.
|
||||
tags: List of tags to filter by. If None or empty, returns all results.
|
||||
match: Matching mode. Defaults to "any".
|
||||
|
||||
Returns:
|
||||
Filtered list of results.
|
||||
"""
|
||||
if not tags:
|
||||
return results
|
||||
|
||||
_, include_untagged = _parse_tags_match(match)
|
||||
is_any_match = match in ("any", "any_strict")
|
||||
|
||||
tags_set = set(tags)
|
||||
filtered = []
|
||||
|
||||
for result in results:
|
||||
result_tags = getattr(result, "tags", None)
|
||||
|
||||
# Check if untagged
|
||||
is_untagged = result_tags is None or len(result_tags) == 0
|
||||
|
||||
if is_untagged:
|
||||
if include_untagged:
|
||||
filtered.append(result)
|
||||
# else: skip untagged
|
||||
else:
|
||||
result_tags_set = set(result_tags)
|
||||
if is_any_match:
|
||||
# Any overlap
|
||||
if result_tags_set & tags_set:
|
||||
filtered.append(result)
|
||||
else:
|
||||
# All tags must be present
|
||||
if tags_set <= result_tags_set:
|
||||
filtered.append(result)
|
||||
|
||||
return filtered
|
||||
@@ -11,6 +11,13 @@ from typing import Any, Literal
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class TemporalConstraint(BaseModel):
|
||||
"""Detected temporal constraint from query analysis."""
|
||||
|
||||
start: datetime | None = Field(default=None, description="Start of temporal range")
|
||||
end: datetime | None = Field(default=None, description="End of temporal range")
|
||||
|
||||
|
||||
class QueryInfo(BaseModel):
|
||||
"""Information about the search query."""
|
||||
|
||||
@@ -19,6 +26,11 @@ class QueryInfo(BaseModel):
|
||||
timestamp: datetime = Field(description="When the query was executed")
|
||||
budget: int = Field(description="Maximum nodes to explore")
|
||||
max_tokens: int = Field(description="Maximum tokens to return in results")
|
||||
tags: list[str] | None = Field(default=None, description="Tags filter applied to recall")
|
||||
tags_match: str | None = Field(default=None, description="Tags matching mode: any, all, any_strict, all_strict")
|
||||
temporal_constraint: TemporalConstraint | None = Field(
|
||||
default=None, description="Detected temporal range from query"
|
||||
)
|
||||
|
||||
|
||||
class EntryPoint(BaseModel):
|
||||
|
||||
@@ -22,6 +22,7 @@ from .trace import (
|
||||
SearchPhaseMetrics,
|
||||
SearchSummary,
|
||||
SearchTrace,
|
||||
TemporalConstraint,
|
||||
WeightComponents,
|
||||
)
|
||||
|
||||
@@ -45,7 +46,14 @@ class SearchTracer:
|
||||
json_output = trace.to_json()
|
||||
"""
|
||||
|
||||
def __init__(self, query: str, budget: int, max_tokens: int):
|
||||
def __init__(
|
||||
self,
|
||||
query: str,
|
||||
budget: int,
|
||||
max_tokens: int,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: str | None = None,
|
||||
):
|
||||
"""
|
||||
Initialize tracer.
|
||||
|
||||
@@ -53,10 +61,14 @@ class SearchTracer:
|
||||
query: Search query text
|
||||
budget: Maximum nodes to explore
|
||||
max_tokens: Maximum tokens to return in results
|
||||
tags: Tags filter applied to recall
|
||||
tags_match: Tags matching mode (any, all, any_strict, all_strict)
|
||||
"""
|
||||
self.query_text = query
|
||||
self.budget = budget
|
||||
self.max_tokens = max_tokens
|
||||
self.tags = tags
|
||||
self.tags_match = tags_match
|
||||
|
||||
# Trace data
|
||||
self.query_embedding: list[float] | None = None
|
||||
@@ -66,6 +78,9 @@ class SearchTracer:
|
||||
self.pruned: list[PruningDecision] = []
|
||||
self.phase_metrics: list[SearchPhaseMetrics] = []
|
||||
|
||||
# Temporal constraint detected from query
|
||||
self.temporal_constraint: TemporalConstraint | None = None
|
||||
|
||||
# New 4-way retrieval tracking
|
||||
self.retrieval_results: list[RetrievalMethodResults] = []
|
||||
self.rrf_merged: list[RRFMergeResult] = []
|
||||
@@ -88,6 +103,11 @@ class SearchTracer:
|
||||
"""Record the query embedding."""
|
||||
self.query_embedding = embedding
|
||||
|
||||
def record_temporal_constraint(self, start: datetime | None, end: datetime | None):
|
||||
"""Record the detected temporal constraint from query analysis."""
|
||||
if start is not None or end is not None:
|
||||
self.temporal_constraint = TemporalConstraint(start=start, end=end)
|
||||
|
||||
def add_entry_point(self, node_id: str, text: str, similarity: float, rank: int):
|
||||
"""
|
||||
Record an entry point.
|
||||
@@ -428,6 +448,9 @@ class SearchTracer:
|
||||
timestamp=datetime.now(UTC),
|
||||
budget=self.budget,
|
||||
max_tokens=self.max_tokens,
|
||||
tags=self.tags,
|
||||
tags_match=self.tags_match,
|
||||
temporal_constraint=self.temporal_constraint,
|
||||
)
|
||||
|
||||
# Create summary
|
||||
|
||||
@@ -48,6 +48,7 @@ class RetrievalResult:
|
||||
chunk_id: str | None = None
|
||||
access_count: int = 0
|
||||
embedding: list[float] | None = None
|
||||
tags: list[str] | None = None # Visibility scope tags
|
||||
|
||||
# Retrieval-specific scores (only one will be set depending on retrieval method)
|
||||
similarity: float | None = None # Semantic retrieval
|
||||
@@ -72,6 +73,7 @@ class RetrievalResult:
|
||||
chunk_id=row.get("chunk_id"),
|
||||
access_count=row.get("access_count", 0),
|
||||
embedding=row.get("embedding"),
|
||||
tags=row.get("tags"),
|
||||
similarity=row.get("similarity"),
|
||||
bm25_score=row.get("bm25_score"),
|
||||
activation=row.get("activation"),
|
||||
@@ -156,6 +158,7 @@ class ScoredResult:
|
||||
"chunk_id": self.retrieval.chunk_id,
|
||||
"access_count": self.retrieval.access_count,
|
||||
"embedding": self.retrieval.embedding,
|
||||
"tags": self.retrieval.tags,
|
||||
"semantic_similarity": self.retrieval.similarity,
|
||||
"bm25_score": self.retrieval.bm25_score,
|
||||
}
|
||||
|
||||
@@ -187,12 +187,15 @@ def main():
|
||||
embeddings_provider=config.embeddings_provider,
|
||||
embeddings_local_model=config.embeddings_local_model,
|
||||
embeddings_tei_url=config.embeddings_tei_url,
|
||||
embeddings_openai_base_url=config.embeddings_openai_base_url,
|
||||
embeddings_cohere_base_url=config.embeddings_cohere_base_url,
|
||||
reranker_provider=config.reranker_provider,
|
||||
reranker_local_model=config.reranker_local_model,
|
||||
reranker_tei_url=config.reranker_tei_url,
|
||||
reranker_tei_batch_size=config.reranker_tei_batch_size,
|
||||
reranker_tei_max_concurrent=config.reranker_tei_max_concurrent,
|
||||
reranker_max_candidates=config.reranker_max_candidates,
|
||||
reranker_cohere_base_url=config.reranker_cohere_base_url,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
log_level=args.log_level,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -115,6 +115,7 @@ class Hindsight:
|
||||
document_id: Optional[str] = None,
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
entities: Optional[List[Dict[str, str]]] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
) -> RetainResponse:
|
||||
"""
|
||||
Store a single memory (simplified interface).
|
||||
@@ -127,13 +128,14 @@ class Hindsight:
|
||||
document_id: Optional document ID for grouping
|
||||
metadata: Optional user-defined metadata
|
||||
entities: Optional list of entities [{"text": "...", "type": "..."}]
|
||||
tags: Optional list of tags for this memory
|
||||
|
||||
Returns:
|
||||
RetainResponse with success status
|
||||
"""
|
||||
return self.retain_batch(
|
||||
bank_id=bank_id,
|
||||
items=[{"content": content, "timestamp": timestamp, "context": context, "metadata": metadata, "entities": entities}],
|
||||
items=[{"content": content, "timestamp": timestamp, "context": context, "metadata": metadata, "entities": entities, "tags": tags}],
|
||||
document_id=document_id,
|
||||
)
|
||||
|
||||
@@ -143,15 +145,17 @@ class Hindsight:
|
||||
items: List[Dict[str, Any]],
|
||||
document_id: Optional[str] = None,
|
||||
retain_async: bool = False,
|
||||
document_tags: Optional[List[str]] = None,
|
||||
) -> RetainResponse:
|
||||
"""
|
||||
Store multiple memories in batch.
|
||||
|
||||
Args:
|
||||
bank_id: The memory bank ID
|
||||
items: List of memory items with 'content' and optional 'timestamp', 'context', 'metadata', 'document_id', 'entities'
|
||||
items: List of memory items with 'content' and optional 'timestamp', 'context', 'metadata', 'document_id', 'entities', 'tags'
|
||||
document_id: Optional document ID for grouping memories (applied to items that don't have their own)
|
||||
retain_async: If True, process asynchronously in background (default: False)
|
||||
document_tags: Optional list of tags to apply to all memories in this batch
|
||||
|
||||
Returns:
|
||||
RetainResponse with success status and item count
|
||||
@@ -175,12 +179,14 @@ class Hindsight:
|
||||
# Use item's document_id if provided, otherwise fall back to batch-level document_id
|
||||
document_id=item.get("document_id") or document_id,
|
||||
entities=entities,
|
||||
tags=item.get("tags"),
|
||||
)
|
||||
)
|
||||
|
||||
request_obj = retain_request.RetainRequest(
|
||||
items=memory_items,
|
||||
async_=retain_async,
|
||||
document_tags=document_tags,
|
||||
)
|
||||
|
||||
return _run_async(self._memory_api.retain_memories(bank_id, request_obj))
|
||||
@@ -198,6 +204,8 @@ class Hindsight:
|
||||
max_entity_tokens: int = 500,
|
||||
include_chunks: bool = False,
|
||||
max_chunk_tokens: int = 8192,
|
||||
tags: Optional[List[str]] = None,
|
||||
tags_match: str = "any",
|
||||
) -> RecallResponse:
|
||||
"""
|
||||
Recall memories using semantic similarity.
|
||||
@@ -214,6 +222,9 @@ class Hindsight:
|
||||
max_entity_tokens: Maximum tokens for entity observations (default: 500)
|
||||
include_chunks: Include raw text chunks in results (default: False)
|
||||
max_chunk_tokens: Maximum tokens for chunks (default: 8192)
|
||||
tags: Optional list of tags to filter memories by
|
||||
tags_match: How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged),
|
||||
'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged). Default: 'any'
|
||||
|
||||
Returns:
|
||||
RecallResponse with results, optional entities, optional chunks, and optional trace
|
||||
@@ -233,6 +244,8 @@ class Hindsight:
|
||||
trace=trace,
|
||||
query_timestamp=query_timestamp,
|
||||
include=include_opts,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
|
||||
return _run_async(self._memory_api.recall_memories(bank_id, request_obj))
|
||||
@@ -245,6 +258,8 @@ class Hindsight:
|
||||
context: Optional[str] = None,
|
||||
max_tokens: Optional[int] = None,
|
||||
response_schema: Optional[Dict[str, Any]] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
tags_match: str = "any",
|
||||
) -> ReflectResponse:
|
||||
"""
|
||||
Generate a contextual answer based on bank identity and memories.
|
||||
@@ -258,6 +273,9 @@ class Hindsight:
|
||||
response_schema: Optional JSON Schema for structured output. When provided,
|
||||
the response will include a 'structured_output' field with the LLM
|
||||
response parsed according to this schema.
|
||||
tags: Optional list of tags to filter memories by
|
||||
tags_match: How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged),
|
||||
'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged). Default: 'any'
|
||||
|
||||
Returns:
|
||||
ReflectResponse with answer text, optionally facts used, and optionally
|
||||
@@ -269,6 +287,8 @@ class Hindsight:
|
||||
context=context,
|
||||
max_tokens=max_tokens,
|
||||
response_schema=response_schema,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
|
||||
return _run_async(self._memory_api.reflect(bank_id, request_obj))
|
||||
|
||||
@@ -64,6 +64,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
|
||||
from hindsight_client_api.models.include_options import IncludeOptions
|
||||
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.list_tags_response import ListTagsResponse
|
||||
from hindsight_client_api.models.memory_item import MemoryItem
|
||||
from hindsight_client_api.models.operation_response import OperationResponse
|
||||
from hindsight_client_api.models.operations_list_response import OperationsListResponse
|
||||
@@ -76,6 +77,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
|
||||
from hindsight_client_api.models.reflect_response import ReflectResponse
|
||||
from hindsight_client_api.models.retain_request import RetainRequest
|
||||
from hindsight_client_api.models.retain_response import RetainResponse
|
||||
from hindsight_client_api.models.tag_item import TagItem
|
||||
from hindsight_client_api.models.token_usage import TokenUsage
|
||||
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
|
||||
from hindsight_client_api.models.validation_error import ValidationError
|
||||
|
||||
@@ -17,11 +17,12 @@ from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from typing_extensions import Annotated
|
||||
|
||||
from pydantic import Field, StrictInt, StrictStr
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
from typing_extensions import Annotated
|
||||
from hindsight_client_api.models.delete_response import DeleteResponse
|
||||
from hindsight_client_api.models.graph_data_response import GraphDataResponse
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.list_tags_response import ListTagsResponse
|
||||
from hindsight_client_api.models.recall_request import RecallRequest
|
||||
from hindsight_client_api.models.recall_response import RecallResponse
|
||||
from hindsight_client_api.models.reflect_request import ReflectRequest
|
||||
@@ -654,6 +655,299 @@ class MemoryApi:
|
||||
|
||||
|
||||
|
||||
@validate_call
|
||||
async def get_memory(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
memory_id: StrictStr,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> object:
|
||||
"""Get memory unit
|
||||
|
||||
Get a single memory unit by ID with all its metadata including entities and tags.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param memory_id: (required)
|
||||
:type memory_id: str
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._get_memory_serialize(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "object",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
await response_data.read()
|
||||
return self.api_client.response_deserialize(
|
||||
response_data=response_data,
|
||||
response_types_map=_response_types_map,
|
||||
).data
|
||||
|
||||
|
||||
@validate_call
|
||||
async def get_memory_with_http_info(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
memory_id: StrictStr,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> ApiResponse[object]:
|
||||
"""Get memory unit
|
||||
|
||||
Get a single memory unit by ID with all its metadata including entities and tags.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param memory_id: (required)
|
||||
:type memory_id: str
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._get_memory_serialize(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "object",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
await response_data.read()
|
||||
return self.api_client.response_deserialize(
|
||||
response_data=response_data,
|
||||
response_types_map=_response_types_map,
|
||||
)
|
||||
|
||||
|
||||
@validate_call
|
||||
async def get_memory_without_preload_content(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
memory_id: StrictStr,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> RESTResponseType:
|
||||
"""Get memory unit
|
||||
|
||||
Get a single memory unit by ID with all its metadata including entities and tags.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param memory_id: (required)
|
||||
:type memory_id: str
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._get_memory_serialize(
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "object",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
return response_data.response
|
||||
|
||||
|
||||
def _get_memory_serialize(
|
||||
self,
|
||||
bank_id,
|
||||
memory_id,
|
||||
authorization,
|
||||
_request_auth,
|
||||
_content_type,
|
||||
_headers,
|
||||
_host_index,
|
||||
) -> RequestSerialized:
|
||||
|
||||
_host = None
|
||||
|
||||
_collection_formats: Dict[str, str] = {
|
||||
}
|
||||
|
||||
_path_params: Dict[str, str] = {}
|
||||
_query_params: List[Tuple[str, str]] = []
|
||||
_header_params: Dict[str, Optional[str]] = _headers or {}
|
||||
_form_params: List[Tuple[str, str]] = []
|
||||
_files: Dict[
|
||||
str, Union[str, bytes, List[str], List[bytes], List[Tuple[str, bytes]]]
|
||||
] = {}
|
||||
_body_params: Optional[bytes] = None
|
||||
|
||||
# process the path parameters
|
||||
if bank_id is not None:
|
||||
_path_params['bank_id'] = bank_id
|
||||
if memory_id is not None:
|
||||
_path_params['memory_id'] = memory_id
|
||||
# process the query parameters
|
||||
# process the header parameters
|
||||
if authorization is not None:
|
||||
_header_params['authorization'] = authorization
|
||||
# process the form parameters
|
||||
# process the body parameter
|
||||
|
||||
|
||||
# set the HTTP header `Accept`
|
||||
if 'Accept' not in _header_params:
|
||||
_header_params['Accept'] = self.api_client.select_header_accept(
|
||||
[
|
||||
'application/json'
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# authentication setting
|
||||
_auth_settings: List[str] = [
|
||||
]
|
||||
|
||||
return self.api_client.param_serialize(
|
||||
method='GET',
|
||||
resource_path='/v1/default/banks/{bank_id}/memories/{memory_id}',
|
||||
path_params=_path_params,
|
||||
query_params=_query_params,
|
||||
header_params=_header_params,
|
||||
body=_body_params,
|
||||
post_params=_form_params,
|
||||
files=_files,
|
||||
auth_settings=_auth_settings,
|
||||
collection_formats=_collection_formats,
|
||||
_host=_host,
|
||||
_request_auth=_request_auth
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
@validate_call
|
||||
async def list_memories(
|
||||
self,
|
||||
@@ -1000,6 +1294,335 @@ class MemoryApi:
|
||||
|
||||
|
||||
|
||||
@validate_call
|
||||
async def list_tags(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
|
||||
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
|
||||
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> ListTagsResponse:
|
||||
"""List tags
|
||||
|
||||
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
|
||||
:type q: str
|
||||
:param limit: Maximum number of tags to return
|
||||
:type limit: int
|
||||
:param offset: Offset for pagination
|
||||
:type offset: int
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._list_tags_serialize(
|
||||
bank_id=bank_id,
|
||||
q=q,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "ListTagsResponse",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
await response_data.read()
|
||||
return self.api_client.response_deserialize(
|
||||
response_data=response_data,
|
||||
response_types_map=_response_types_map,
|
||||
).data
|
||||
|
||||
|
||||
@validate_call
|
||||
async def list_tags_with_http_info(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
|
||||
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
|
||||
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> ApiResponse[ListTagsResponse]:
|
||||
"""List tags
|
||||
|
||||
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
|
||||
:type q: str
|
||||
:param limit: Maximum number of tags to return
|
||||
:type limit: int
|
||||
:param offset: Offset for pagination
|
||||
:type offset: int
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._list_tags_serialize(
|
||||
bank_id=bank_id,
|
||||
q=q,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "ListTagsResponse",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
await response_data.read()
|
||||
return self.api_client.response_deserialize(
|
||||
response_data=response_data,
|
||||
response_types_map=_response_types_map,
|
||||
)
|
||||
|
||||
|
||||
@validate_call
|
||||
async def list_tags_without_preload_content(
|
||||
self,
|
||||
bank_id: StrictStr,
|
||||
q: Annotated[Optional[StrictStr], Field(description="Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.")] = None,
|
||||
limit: Annotated[Optional[StrictInt], Field(description="Maximum number of tags to return")] = None,
|
||||
offset: Annotated[Optional[StrictInt], Field(description="Offset for pagination")] = None,
|
||||
authorization: Optional[StrictStr] = None,
|
||||
_request_timeout: Union[
|
||||
None,
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Tuple[
|
||||
Annotated[StrictFloat, Field(gt=0)],
|
||||
Annotated[StrictFloat, Field(gt=0)]
|
||||
]
|
||||
] = None,
|
||||
_request_auth: Optional[Dict[StrictStr, Any]] = None,
|
||||
_content_type: Optional[StrictStr] = None,
|
||||
_headers: Optional[Dict[StrictStr, Any]] = None,
|
||||
_host_index: Annotated[StrictInt, Field(ge=0, le=0)] = 0,
|
||||
) -> RESTResponseType:
|
||||
"""List tags
|
||||
|
||||
List all unique tags in a memory bank with usage counts. Supports wildcard search using '*' (e.g., 'user:*', '*-fred', 'tag*-2'). Case-insensitive.
|
||||
|
||||
:param bank_id: (required)
|
||||
:type bank_id: str
|
||||
:param q: Wildcard pattern to filter tags (e.g., 'user:*' for user:alice, '*-admin' for role-admin). Use '*' as wildcard. Case-insensitive.
|
||||
:type q: str
|
||||
:param limit: Maximum number of tags to return
|
||||
:type limit: int
|
||||
:param offset: Offset for pagination
|
||||
:type offset: int
|
||||
:param authorization:
|
||||
:type authorization: str
|
||||
:param _request_timeout: timeout setting for this request. If one
|
||||
number provided, it will be total request
|
||||
timeout. It can also be a pair (tuple) of
|
||||
(connection, read) timeouts.
|
||||
:type _request_timeout: int, tuple(int, int), optional
|
||||
:param _request_auth: set to override the auth_settings for an a single
|
||||
request; this effectively ignores the
|
||||
authentication in the spec for a single request.
|
||||
:type _request_auth: dict, optional
|
||||
:param _content_type: force content-type for the request.
|
||||
:type _content_type: str, Optional
|
||||
:param _headers: set to override the headers for a single
|
||||
request; this effectively ignores the headers
|
||||
in the spec for a single request.
|
||||
:type _headers: dict, optional
|
||||
:param _host_index: set to override the host_index for a single
|
||||
request; this effectively ignores the host_index
|
||||
in the spec for a single request.
|
||||
:type _host_index: int, optional
|
||||
:return: Returns the result object.
|
||||
""" # noqa: E501
|
||||
|
||||
_param = self._list_tags_serialize(
|
||||
bank_id=bank_id,
|
||||
q=q,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
authorization=authorization,
|
||||
_request_auth=_request_auth,
|
||||
_content_type=_content_type,
|
||||
_headers=_headers,
|
||||
_host_index=_host_index
|
||||
)
|
||||
|
||||
_response_types_map: Dict[str, Optional[str]] = {
|
||||
'200': "ListTagsResponse",
|
||||
'422': "HTTPValidationError",
|
||||
}
|
||||
response_data = await self.api_client.call_api(
|
||||
*_param,
|
||||
_request_timeout=_request_timeout
|
||||
)
|
||||
return response_data.response
|
||||
|
||||
|
||||
def _list_tags_serialize(
|
||||
self,
|
||||
bank_id,
|
||||
q,
|
||||
limit,
|
||||
offset,
|
||||
authorization,
|
||||
_request_auth,
|
||||
_content_type,
|
||||
_headers,
|
||||
_host_index,
|
||||
) -> RequestSerialized:
|
||||
|
||||
_host = None
|
||||
|
||||
_collection_formats: Dict[str, str] = {
|
||||
}
|
||||
|
||||
_path_params: Dict[str, str] = {}
|
||||
_query_params: List[Tuple[str, str]] = []
|
||||
_header_params: Dict[str, Optional[str]] = _headers or {}
|
||||
_form_params: List[Tuple[str, str]] = []
|
||||
_files: Dict[
|
||||
str, Union[str, bytes, List[str], List[bytes], List[Tuple[str, bytes]]]
|
||||
] = {}
|
||||
_body_params: Optional[bytes] = None
|
||||
|
||||
# process the path parameters
|
||||
if bank_id is not None:
|
||||
_path_params['bank_id'] = bank_id
|
||||
# process the query parameters
|
||||
if q is not None:
|
||||
|
||||
_query_params.append(('q', q))
|
||||
|
||||
if limit is not None:
|
||||
|
||||
_query_params.append(('limit', limit))
|
||||
|
||||
if offset is not None:
|
||||
|
||||
_query_params.append(('offset', offset))
|
||||
|
||||
# process the header parameters
|
||||
if authorization is not None:
|
||||
_header_params['authorization'] = authorization
|
||||
# process the form parameters
|
||||
# process the body parameter
|
||||
|
||||
|
||||
# set the HTTP header `Accept`
|
||||
if 'Accept' not in _header_params:
|
||||
_header_params['Accept'] = self.api_client.select_header_accept(
|
||||
[
|
||||
'application/json'
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# authentication setting
|
||||
_auth_settings: List[str] = [
|
||||
]
|
||||
|
||||
return self.api_client.param_serialize(
|
||||
method='GET',
|
||||
resource_path='/v1/default/banks/{bank_id}/tags',
|
||||
path_params=_path_params,
|
||||
query_params=_query_params,
|
||||
header_params=_header_params,
|
||||
body=_body_params,
|
||||
post_params=_form_params,
|
||||
files=_files,
|
||||
auth_settings=_auth_settings,
|
||||
collection_formats=_collection_formats,
|
||||
_host=_host,
|
||||
_request_auth=_request_auth
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
@validate_call
|
||||
async def recall_memories(
|
||||
self,
|
||||
|
||||
@@ -42,6 +42,7 @@ from hindsight_client_api.models.http_validation_error import HTTPValidationErro
|
||||
from hindsight_client_api.models.include_options import IncludeOptions
|
||||
from hindsight_client_api.models.list_documents_response import ListDocumentsResponse
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.list_tags_response import ListTagsResponse
|
||||
from hindsight_client_api.models.memory_item import MemoryItem
|
||||
from hindsight_client_api.models.operation_response import OperationResponse
|
||||
from hindsight_client_api.models.operations_list_response import OperationsListResponse
|
||||
@@ -54,6 +55,7 @@ from hindsight_client_api.models.reflect_request import ReflectRequest
|
||||
from hindsight_client_api.models.reflect_response import ReflectResponse
|
||||
from hindsight_client_api.models.retain_request import RetainRequest
|
||||
from hindsight_client_api.models.retain_response import RetainResponse
|
||||
from hindsight_client_api.models.tag_item import TagItem
|
||||
from hindsight_client_api.models.token_usage import TokenUsage
|
||||
from hindsight_client_api.models.update_disposition_request import UpdateDispositionRequest
|
||||
from hindsight_client_api.models.validation_error import ValidationError
|
||||
|
||||
@@ -17,7 +17,7 @@ import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, StrictInt, StrictStr
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
|
||||
from typing import Any, ClassVar, Dict, List, Optional
|
||||
from typing import Optional, Set
|
||||
from typing_extensions import Self
|
||||
@@ -33,7 +33,8 @@ class DocumentResponse(BaseModel):
|
||||
created_at: StrictStr
|
||||
updated_at: StrictStr
|
||||
memory_unit_count: StrictInt
|
||||
__properties: ClassVar[List[str]] = ["id", "bank_id", "original_text", "content_hash", "created_at", "updated_at", "memory_unit_count"]
|
||||
tags: Optional[List[StrictStr]] = Field(default=None, description="Tags associated with this document")
|
||||
__properties: ClassVar[List[str]] = ["id", "bank_id", "original_text", "content_hash", "created_at", "updated_at", "memory_unit_count", "tags"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -97,7 +98,8 @@ class DocumentResponse(BaseModel):
|
||||
"content_hash": obj.get("content_hash"),
|
||||
"created_at": obj.get("created_at"),
|
||||
"updated_at": obj.get("updated_at"),
|
||||
"memory_unit_count": obj.get("memory_unit_count")
|
||||
"memory_unit_count": obj.get("memory_unit_count"),
|
||||
"tags": obj.get("tags")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
Hindsight HTTP API
|
||||
|
||||
HTTP API for Hindsight
|
||||
|
||||
The version of the OpenAPI document: 0.1.0
|
||||
Generated by OpenAPI Generator (https://openapi-generator.tech)
|
||||
|
||||
Do not edit the class manually.
|
||||
""" # noqa: E501
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, StrictInt
|
||||
from typing import Any, ClassVar, Dict, List
|
||||
from hindsight_client_api.models.tag_item import TagItem
|
||||
from typing import Optional, Set
|
||||
from typing_extensions import Self
|
||||
|
||||
class ListTagsResponse(BaseModel):
|
||||
"""
|
||||
Response model for list tags endpoint.
|
||||
""" # noqa: E501
|
||||
items: List[TagItem]
|
||||
total: StrictInt
|
||||
limit: StrictInt
|
||||
offset: StrictInt
|
||||
__properties: ClassVar[List[str]] = ["items", "total", "limit", "offset"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
validate_assignment=True,
|
||||
protected_namespaces=(),
|
||||
)
|
||||
|
||||
|
||||
def to_str(self) -> str:
|
||||
"""Returns the string representation of the model using alias"""
|
||||
return pprint.pformat(self.model_dump(by_alias=True))
|
||||
|
||||
def to_json(self) -> str:
|
||||
"""Returns the JSON representation of the model using alias"""
|
||||
# TODO: pydantic v2: use .model_dump_json(by_alias=True, exclude_unset=True) instead
|
||||
return json.dumps(self.to_dict())
|
||||
|
||||
@classmethod
|
||||
def from_json(cls, json_str: str) -> Optional[Self]:
|
||||
"""Create an instance of ListTagsResponse from a JSON string"""
|
||||
return cls.from_dict(json.loads(json_str))
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Return the dictionary representation of the model using alias.
|
||||
|
||||
This has the following differences from calling pydantic's
|
||||
`self.model_dump(by_alias=True)`:
|
||||
|
||||
* `None` is only added to the output dict for nullable fields that
|
||||
were set at model initialization. Other fields with value `None`
|
||||
are ignored.
|
||||
"""
|
||||
excluded_fields: Set[str] = set([
|
||||
])
|
||||
|
||||
_dict = self.model_dump(
|
||||
by_alias=True,
|
||||
exclude=excluded_fields,
|
||||
exclude_none=True,
|
||||
)
|
||||
# override the default output from pydantic by calling `to_dict()` of each item in items (list)
|
||||
_items = []
|
||||
if self.items:
|
||||
for _item_items in self.items:
|
||||
if _item_items:
|
||||
_items.append(_item_items.to_dict())
|
||||
_dict['items'] = _items
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional[Self]:
|
||||
"""Create an instance of ListTagsResponse from a dict"""
|
||||
if obj is None:
|
||||
return None
|
||||
|
||||
if not isinstance(obj, dict):
|
||||
return cls.model_validate(obj)
|
||||
|
||||
_obj = cls.model_validate({
|
||||
"items": [TagItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None,
|
||||
"total": obj.get("total"),
|
||||
"limit": obj.get("limit"),
|
||||
"offset": obj.get("offset")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -34,7 +34,8 @@ class MemoryItem(BaseModel):
|
||||
metadata: Optional[Dict[str, StrictStr]] = None
|
||||
document_id: Optional[StrictStr] = None
|
||||
entities: Optional[List[EntityInput]] = None
|
||||
__properties: ClassVar[List[str]] = ["content", "timestamp", "context", "metadata", "document_id", "entities"]
|
||||
tags: Optional[List[StrictStr]] = None
|
||||
__properties: ClassVar[List[str]] = ["content", "timestamp", "context", "metadata", "document_id", "entities", "tags"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -107,6 +108,11 @@ class MemoryItem(BaseModel):
|
||||
if self.entities is None and "entities" in self.model_fields_set:
|
||||
_dict['entities'] = None
|
||||
|
||||
# set to None if tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.tags is None and "tags" in self.model_fields_set:
|
||||
_dict['tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -124,7 +130,8 @@ class MemoryItem(BaseModel):
|
||||
"context": obj.get("context"),
|
||||
"metadata": obj.get("metadata"),
|
||||
"document_id": obj.get("document_id"),
|
||||
"entities": [EntityInput.from_dict(_item) for _item in obj["entities"]] if obj.get("entities") is not None else None
|
||||
"entities": [EntityInput.from_dict(_item) for _item in obj["entities"]] if obj.get("entities") is not None else None,
|
||||
"tags": obj.get("tags")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictInt, StrictStr
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictInt, StrictStr, field_validator
|
||||
from typing import Any, ClassVar, Dict, List, Optional
|
||||
from hindsight_client_api.models.budget import Budget
|
||||
from hindsight_client_api.models.include_options import IncludeOptions
|
||||
@@ -35,7 +35,19 @@ class RecallRequest(BaseModel):
|
||||
trace: Optional[StrictBool] = False
|
||||
query_timestamp: Optional[StrictStr] = None
|
||||
include: Optional[IncludeOptions] = Field(default=None, description="Options for including additional data (entities are included by default)")
|
||||
__properties: ClassVar[List[str]] = ["query", "types", "budget", "max_tokens", "trace", "query_timestamp", "include"]
|
||||
tags: Optional[List[StrictStr]] = None
|
||||
tags_match: Optional[StrictStr] = Field(default='any', description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).")
|
||||
__properties: ClassVar[List[str]] = ["query", "types", "budget", "max_tokens", "trace", "query_timestamp", "include", "tags", "tags_match"]
|
||||
|
||||
@field_validator('tags_match')
|
||||
def tags_match_validate_enum(cls, value):
|
||||
"""Validates the enum"""
|
||||
if value is None:
|
||||
return value
|
||||
|
||||
if value not in set(['any', 'all', 'any_strict', 'all_strict']):
|
||||
raise ValueError("must be one of enum values ('any', 'all', 'any_strict', 'all_strict')")
|
||||
return value
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -89,6 +101,11 @@ class RecallRequest(BaseModel):
|
||||
if self.query_timestamp is None and "query_timestamp" in self.model_fields_set:
|
||||
_dict['query_timestamp'] = None
|
||||
|
||||
# set to None if tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.tags is None and "tags" in self.model_fields_set:
|
||||
_dict['tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -107,7 +124,9 @@ class RecallRequest(BaseModel):
|
||||
"max_tokens": obj.get("max_tokens") if obj.get("max_tokens") is not None else 4096,
|
||||
"trace": obj.get("trace") if obj.get("trace") is not None else False,
|
||||
"query_timestamp": obj.get("query_timestamp"),
|
||||
"include": IncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None
|
||||
"include": IncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None,
|
||||
"tags": obj.get("tags"),
|
||||
"tags_match": obj.get("tags_match") if obj.get("tags_match") is not None else 'any'
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -37,7 +37,8 @@ class RecallResult(BaseModel):
|
||||
document_id: Optional[StrictStr] = None
|
||||
metadata: Optional[Dict[str, StrictStr]] = None
|
||||
chunk_id: Optional[StrictStr] = None
|
||||
__properties: ClassVar[List[str]] = ["id", "text", "type", "entities", "context", "occurred_start", "occurred_end", "mentioned_at", "document_id", "metadata", "chunk_id"]
|
||||
tags: Optional[List[StrictStr]] = None
|
||||
__properties: ClassVar[List[str]] = ["id", "text", "type", "entities", "context", "occurred_start", "occurred_end", "mentioned_at", "document_id", "metadata", "chunk_id", "tags"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -123,6 +124,11 @@ class RecallResult(BaseModel):
|
||||
if self.chunk_id is None and "chunk_id" in self.model_fields_set:
|
||||
_dict['chunk_id'] = None
|
||||
|
||||
# set to None if tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.tags is None and "tags" in self.model_fields_set:
|
||||
_dict['tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -145,7 +151,8 @@ class RecallResult(BaseModel):
|
||||
"mentioned_at": obj.get("mentioned_at"),
|
||||
"document_id": obj.get("document_id"),
|
||||
"metadata": obj.get("metadata"),
|
||||
"chunk_id": obj.get("chunk_id")
|
||||
"chunk_id": obj.get("chunk_id"),
|
||||
"tags": obj.get("tags")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr, field_validator
|
||||
from typing import Any, ClassVar, Dict, List, Optional
|
||||
from hindsight_client_api.models.budget import Budget
|
||||
from hindsight_client_api.models.reflect_include_options import ReflectIncludeOptions
|
||||
@@ -34,7 +34,19 @@ class ReflectRequest(BaseModel):
|
||||
max_tokens: Optional[StrictInt] = Field(default=4096, description="Maximum tokens for the response")
|
||||
include: Optional[ReflectIncludeOptions] = Field(default=None, description="Options for including additional data (disabled by default)")
|
||||
response_schema: Optional[Dict[str, Any]] = None
|
||||
__properties: ClassVar[List[str]] = ["query", "budget", "context", "max_tokens", "include", "response_schema"]
|
||||
tags: Optional[List[StrictStr]] = None
|
||||
tags_match: Optional[StrictStr] = Field(default='any', description="How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged).")
|
||||
__properties: ClassVar[List[str]] = ["query", "budget", "context", "max_tokens", "include", "response_schema", "tags", "tags_match"]
|
||||
|
||||
@field_validator('tags_match')
|
||||
def tags_match_validate_enum(cls, value):
|
||||
"""Validates the enum"""
|
||||
if value is None:
|
||||
return value
|
||||
|
||||
if value not in set(['any', 'all', 'any_strict', 'all_strict']):
|
||||
raise ValueError("must be one of enum values ('any', 'all', 'any_strict', 'all_strict')")
|
||||
return value
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -88,6 +100,11 @@ class ReflectRequest(BaseModel):
|
||||
if self.response_schema is None and "response_schema" in self.model_fields_set:
|
||||
_dict['response_schema'] = None
|
||||
|
||||
# set to None if tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.tags is None and "tags" in self.model_fields_set:
|
||||
_dict['tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -105,7 +122,9 @@ class ReflectRequest(BaseModel):
|
||||
"context": obj.get("context"),
|
||||
"max_tokens": obj.get("max_tokens") if obj.get("max_tokens") is not None else 4096,
|
||||
"include": ReflectIncludeOptions.from_dict(obj["include"]) if obj.get("include") is not None else None,
|
||||
"response_schema": obj.get("response_schema")
|
||||
"response_schema": obj.get("response_schema"),
|
||||
"tags": obj.get("tags"),
|
||||
"tags_match": obj.get("tags_match") if obj.get("tags_match") is not None else 'any'
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictBool
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictBool, StrictStr
|
||||
from typing import Any, ClassVar, Dict, List, Optional
|
||||
from hindsight_client_api.models.memory_item import MemoryItem
|
||||
from typing import Optional, Set
|
||||
@@ -29,7 +29,8 @@ class RetainRequest(BaseModel):
|
||||
""" # noqa: E501
|
||||
items: List[MemoryItem]
|
||||
var_async: Optional[StrictBool] = Field(default=False, description="If true, process asynchronously in background. If false, wait for completion (default: false)", alias="async")
|
||||
__properties: ClassVar[List[str]] = ["items", "async"]
|
||||
document_tags: Optional[List[StrictStr]] = None
|
||||
__properties: ClassVar[List[str]] = ["items", "async", "document_tags"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
@@ -77,6 +78,11 @@ class RetainRequest(BaseModel):
|
||||
if _item_items:
|
||||
_items.append(_item_items.to_dict())
|
||||
_dict['items'] = _items
|
||||
# set to None if document_tags (nullable) is None
|
||||
# and model_fields_set contains the field
|
||||
if self.document_tags is None and "document_tags" in self.model_fields_set:
|
||||
_dict['document_tags'] = None
|
||||
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
@@ -90,7 +96,8 @@ class RetainRequest(BaseModel):
|
||||
|
||||
_obj = cls.model_validate({
|
||||
"items": [MemoryItem.from_dict(_item) for _item in obj["items"]] if obj.get("items") is not None else None,
|
||||
"async": obj.get("async") if obj.get("async") is not None else False
|
||||
"async": obj.get("async") if obj.get("async") is not None else False,
|
||||
"document_tags": obj.get("document_tags")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
Hindsight HTTP API
|
||||
|
||||
HTTP API for Hindsight
|
||||
|
||||
The version of the OpenAPI document: 0.1.0
|
||||
Generated by OpenAPI Generator (https://openapi-generator.tech)
|
||||
|
||||
Do not edit the class manually.
|
||||
""" # noqa: E501
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
import pprint
|
||||
import re # noqa: F401
|
||||
import json
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, StrictInt, StrictStr
|
||||
from typing import Any, ClassVar, Dict, List
|
||||
from typing import Optional, Set
|
||||
from typing_extensions import Self
|
||||
|
||||
class TagItem(BaseModel):
|
||||
"""
|
||||
Single tag with usage count.
|
||||
""" # noqa: E501
|
||||
tag: StrictStr = Field(description="The tag value")
|
||||
count: StrictInt = Field(description="Number of memories with this tag")
|
||||
__properties: ClassVar[List[str]] = ["tag", "count"]
|
||||
|
||||
model_config = ConfigDict(
|
||||
populate_by_name=True,
|
||||
validate_assignment=True,
|
||||
protected_namespaces=(),
|
||||
)
|
||||
|
||||
|
||||
def to_str(self) -> str:
|
||||
"""Returns the string representation of the model using alias"""
|
||||
return pprint.pformat(self.model_dump(by_alias=True))
|
||||
|
||||
def to_json(self) -> str:
|
||||
"""Returns the JSON representation of the model using alias"""
|
||||
# TODO: pydantic v2: use .model_dump_json(by_alias=True, exclude_unset=True) instead
|
||||
return json.dumps(self.to_dict())
|
||||
|
||||
@classmethod
|
||||
def from_json(cls, json_str: str) -> Optional[Self]:
|
||||
"""Create an instance of TagItem from a JSON string"""
|
||||
return cls.from_dict(json.loads(json_str))
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""Return the dictionary representation of the model using alias.
|
||||
|
||||
This has the following differences from calling pydantic's
|
||||
`self.model_dump(by_alias=True)`:
|
||||
|
||||
* `None` is only added to the output dict for nullable fields that
|
||||
were set at model initialization. Other fields with value `None`
|
||||
are ignored.
|
||||
"""
|
||||
excluded_fields: Set[str] = set([
|
||||
])
|
||||
|
||||
_dict = self.model_dump(
|
||||
by_alias=True,
|
||||
exclude=excluded_fields,
|
||||
exclude_none=True,
|
||||
)
|
||||
return _dict
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, obj: Optional[Dict[str, Any]]) -> Optional[Self]:
|
||||
"""Create an instance of TagItem from a dict"""
|
||||
if obj is None:
|
||||
return None
|
||||
|
||||
if not isinstance(obj, dict):
|
||||
return cls.model_validate(obj)
|
||||
|
||||
_obj = cls.model_validate({
|
||||
"tag": obj.get("tag"),
|
||||
"count": obj.get("count")
|
||||
})
|
||||
return _obj
|
||||
|
||||
|
||||
@@ -1,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 {
|
||||
@@ -101,6 +102,8 @@ export class HindsightClient {
|
||||
documentId?: string;
|
||||
async?: boolean;
|
||||
entities?: EntityInput[];
|
||||
/** Optional list of tags for this memory */
|
||||
tags?: string[];
|
||||
}
|
||||
): Promise<RetainResponse> {
|
||||
const item: {
|
||||
@@ -110,6 +113,7 @@ export class HindsightClient {
|
||||
metadata?: Record<string, string>;
|
||||
document_id?: string;
|
||||
entities?: EntityInput[];
|
||||
tags?: string[];
|
||||
} = { content };
|
||||
if (options?.timestamp) {
|
||||
item.timestamp =
|
||||
@@ -129,6 +133,9 @@ export class HindsightClient {
|
||||
if (options?.entities) {
|
||||
item.entities = options.entities;
|
||||
}
|
||||
if (options?.tags) {
|
||||
item.tags = options.tags;
|
||||
}
|
||||
|
||||
const response = await sdk.retainMemories({
|
||||
client: this.client,
|
||||
@@ -142,13 +149,14 @@ export class HindsightClient {
|
||||
/**
|
||||
* Retain multiple memories in batch.
|
||||
*/
|
||||
async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; async?: boolean }): Promise<RetainResponse> {
|
||||
async retainBatch(bankId: string, items: MemoryItemInput[], options?: { documentId?: string; documentTags?: string[]; async?: boolean }): Promise<RetainResponse> {
|
||||
const processedItems = items.map((item) => ({
|
||||
content: item.content,
|
||||
context: item.context,
|
||||
metadata: item.metadata,
|
||||
document_id: item.document_id,
|
||||
entities: item.entities,
|
||||
tags: item.tags,
|
||||
timestamp:
|
||||
item.timestamp instanceof Date
|
||||
? item.timestamp.toISOString()
|
||||
@@ -166,6 +174,7 @@ export class HindsightClient {
|
||||
path: { bank_id: bankId },
|
||||
body: {
|
||||
items: itemsWithDocId,
|
||||
document_tags: options?.documentTags,
|
||||
async: options?.async,
|
||||
},
|
||||
});
|
||||
@@ -189,6 +198,10 @@ export class HindsightClient {
|
||||
maxEntityTokens?: number;
|
||||
includeChunks?: boolean;
|
||||
maxChunkTokens?: number;
|
||||
/** Optional list of tags to filter memories by */
|
||||
tags?: string[];
|
||||
/** How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged). Default: 'any' */
|
||||
tagsMatch?: 'any' | 'all' | 'any_strict' | 'all_strict';
|
||||
}
|
||||
): Promise<RecallResponse> {
|
||||
const response = await sdk.recallMemories({
|
||||
@@ -205,6 +218,8 @@ export class HindsightClient {
|
||||
entities: options?.includeEntities ? { max_tokens: options?.maxEntityTokens ?? 500 } : undefined,
|
||||
chunks: options?.includeChunks ? { max_tokens: options?.maxChunkTokens ?? 8192 } : undefined,
|
||||
},
|
||||
tags: options?.tags,
|
||||
tags_match: options?.tagsMatch,
|
||||
},
|
||||
});
|
||||
|
||||
@@ -217,7 +232,14 @@ export class HindsightClient {
|
||||
async reflect(
|
||||
bankId: string,
|
||||
query: string,
|
||||
options?: { context?: string; budget?: Budget }
|
||||
options?: {
|
||||
context?: string;
|
||||
budget?: Budget;
|
||||
/** Optional list of tags to filter memories by */
|
||||
tags?: string[];
|
||||
/** How to match tags: 'any' (OR, includes untagged), 'all' (AND, includes untagged), 'any_strict' (OR, excludes untagged), 'all_strict' (AND, excludes untagged). Default: 'any' */
|
||||
tagsMatch?: 'any' | 'all' | 'any_strict' | 'all_strict';
|
||||
}
|
||||
): Promise<ReflectResponse> {
|
||||
const response = await sdk.reflect({
|
||||
client: this.client,
|
||||
@@ -226,6 +248,8 @@ export class HindsightClient {
|
||||
query,
|
||||
context: options?.context,
|
||||
budget: options?.budget || 'low',
|
||||
tags: options?.tags,
|
||||
tags_match: options?.tagsMatch,
|
||||
},
|
||||
});
|
||||
|
||||
|
||||
@@ -1,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)
|
||||
|
||||
|
||||
@@ -43,11 +43,13 @@ Make sure you've completed the [Quick Start](./quickstart) to install the client
|
||||
|-----------|------|---------|-------------|
|
||||
| `query` | string | required | Natural language query |
|
||||
| `types` | list | all | Filter: `world`, `experience`, `opinion` |
|
||||
| `budget` | string | "mid" | Budget level: "low", "mid", "high" |
|
||||
| `budget` | string | "mid" | Budget level: `low`, `mid`, `high` |
|
||||
| `max_tokens` | int | 4096 | Token budget for results |
|
||||
| `trace` | bool | false | Enable trace output for debugging |
|
||||
| `include_entities` | bool | false | Include entity observations |
|
||||
| `max_entity_tokens` | int | 500 | Token budget for entity observations |
|
||||
| `tags` | list | None | Filter memories by tags (see [Tag Filtering](#filter-by-tags)) |
|
||||
| `tags_match` | string | "any" | How to match tags: `any`, `all`, `any_strict`, `all_strict` |
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python">
|
||||
@@ -127,3 +129,51 @@ The `budget` parameter controls graph traversal depth:
|
||||
<CodeSnippet code={recallMjs} section="recall-budget-levels" language="javascript" />
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Filter by Tags
|
||||
|
||||
Tags enable **visibility scoping**—filter memories based on tags assigned during [retain](./retain#tagging-memories). This is essential for multi-user agents where each user should only see their own memories.
|
||||
|
||||
### Basic Tag Filtering
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python">
|
||||
<CodeSnippet code={recallPy} section="recall-with-tags" language="python" />
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Tag Match Modes
|
||||
|
||||
The `tags_match` parameter controls how tags are matched:
|
||||
|
||||
| Mode | Behavior | Untagged Memories |
|
||||
|------|----------|-------------------|
|
||||
| `any` | OR: memory has ANY of the specified tags | **Included** |
|
||||
| `all` | AND: memory has ALL of the specified tags | **Included** |
|
||||
| `any_strict` | OR: memory has ANY of the specified tags | **Excluded** |
|
||||
| `all_strict` | AND: memory has ALL of the specified tags | **Excluded** |
|
||||
|
||||
**Strict modes** are useful when you want to ensure only tagged memories are returned:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python">
|
||||
<CodeSnippet code={recallPy} section="recall-tags-strict" language="python" />
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
**AND matching** requires all specified tags to be present:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python">
|
||||
<CodeSnippet code={recallPy} section="recall-tags-all" language="python" />
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Use Cases
|
||||
|
||||
| Scenario | Tags | Mode | Result |
|
||||
|----------|------|------|--------|
|
||||
| User A's memories only | `["user:alice"]` | `any_strict` | Only memories tagged `user:alice` |
|
||||
| Support + feedback | `["support", "feedback"]` | `any` | Memories with either tag + untagged |
|
||||
| Multi-user room | `["user:alice", "room:general"]` | `all_strict` | Only memories with both tags |
|
||||
| Global + user-specific | `["user:alice"]` | `any` | Alice's memories + shared (untagged) |
|
||||
|
||||
@@ -50,10 +50,12 @@ Make sure you've completed the [Quick Start](./quickstart) to install the client
|
||||
| Parameter | Type | Default | Description |
|
||||
|-----------|------|---------|-------------|
|
||||
| `query` | string | required | Question or prompt |
|
||||
| `budget` | string | "low" | Budget level: "low", "mid", "high" |
|
||||
| `budget` | string | "low" | Budget level: `low`, `mid`, `high` |
|
||||
| `context` | string | None | Additional context for the query |
|
||||
| `max_tokens` | int | 4096 | Maximum tokens for the response |
|
||||
| `response_schema` | object | None | JSON Schema for [structured output](#structured-output) |
|
||||
| `tags` | list | None | Filter memories by tags during reflection |
|
||||
| `tags_match` | string | "any" | How to match tags: `any`, `all`, `any_strict`, `all_strict` |
|
||||
|
||||
### Response Fields
|
||||
|
||||
@@ -247,3 +249,24 @@ hindsight memory reflect hiring-team \
|
||||
- Use `model_validate()` to parse the response back into your Pydantic model
|
||||
- Keep schemas focused — extract only what you need
|
||||
- Use `Optional` fields for data that may not always be available
|
||||
|
||||
## Filter by Tags
|
||||
|
||||
Like [recall](./recall#filter-by-tags), reflect supports tag filtering to scope which memories are considered during reasoning. This is essential for multi-user scenarios where reflection should only consider memories relevant to a specific user.
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python">
|
||||
<CodeSnippet code={reflectPy} section="reflect-with-tags" language="python" />
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
The `tags_match` parameter works the same as in recall:
|
||||
|
||||
| Mode | Behavior |
|
||||
|------|----------|
|
||||
| `any` | OR matching, includes untagged memories |
|
||||
| `all` | AND matching, includes untagged memories |
|
||||
| `any_strict` | OR matching, excludes untagged memories |
|
||||
| `all_strict` | AND matching, excludes untagged memories |
|
||||
|
||||
See [Retain API](./retain#tagging-memories) for how to tag memories and [Recall API](./recall#filter-by-tags) for more details on tag matching modes.
|
||||
|
||||
@@ -129,3 +129,55 @@ For large batches, use async ingestion to avoid blocking:
|
||||
<CodeSnippet code={retainMjs} section="retain-async" language="javascript" />
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Tagging Memories
|
||||
|
||||
Tags enable **visibility scoping**—useful when one memory bank serves multiple users but each should only see relevant memories. For example, an agent that chats with multiple users can tag memories by user ID and filter during recall.
|
||||
|
||||
### Tag Individual Items
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python">
|
||||
<CodeSnippet code={retainPy} section="retain-with-tags" language="python" />
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Apply Tags to All Items in a Batch
|
||||
|
||||
Use `document_tags` to apply the same tags to all items in a request:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python">
|
||||
<CodeSnippet code={retainPy} section="retain-with-document-tags" language="python" />
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
When both `document_tags` and item-level `tags` are provided, they are merged together.
|
||||
|
||||
### Tag Naming Conventions
|
||||
|
||||
Use consistent naming patterns for tags:
|
||||
|
||||
| Pattern | Example | Use Case |
|
||||
|---------|---------|----------|
|
||||
| `user:<id>` | `user:alice` | Multi-user agent filtering |
|
||||
| `session:<id>` | `session:123` | Session-based scoping |
|
||||
| `room:<id>` | `room:general` | Chat room isolation |
|
||||
| `topic:<name>` | `topic:feedback` | Topic categorization |
|
||||
|
||||
### Listing Tags
|
||||
|
||||
Use the list tags API to discover existing tags, useful for UI autocomplete or wildcard expansion:
|
||||
|
||||
```python
|
||||
# List all tags in a bank
|
||||
tags = client.list_tags(bank_id="my-bank")
|
||||
for tag in tags.items:
|
||||
print(f"{tag.tag}: {tag.count} memories")
|
||||
|
||||
# Search with wildcards (* matches any characters)
|
||||
user_tags = client.list_tags(bank_id="my-bank", q="user:*")
|
||||
admin_tags = client.list_tags(bank_id="my-bank", q="*-admin")
|
||||
```
|
||||
|
||||
See [Recall API](./recall#filter-by-tags) for filtering memories by tags during retrieval.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -95,29 +95,71 @@ Converts text into dense vector representations for semantic similarity search.
|
||||
|
||||
**Default:** `BAAI/bge-small-en-v1.5` (384 dimensions, ~130MB)
|
||||
|
||||
**Alternatives:**
|
||||
### Supported Providers
|
||||
|
||||
| Model | Use Case |
|
||||
|-------|----------|
|
||||
| `BAAI/bge-small-en-v1.5` | Default, fast, good quality |
|
||||
| `sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2` | Multilingual (50+ languages) |
|
||||
| Provider | Description | Best For |
|
||||
|----------|-------------|----------|
|
||||
| `local` | SentenceTransformers (default) | Development, low latency |
|
||||
| `openai` | OpenAI embeddings API | Production, high quality |
|
||||
| `cohere` | Cohere embeddings API | Production, multilingual |
|
||||
| `tei` | HuggingFace Text Embeddings Inference | Production, self-hosted |
|
||||
| `litellm` | LiteLLM proxy (unified gateway) | Multi-provider setups |
|
||||
|
||||
:::warning
|
||||
All embedding models must produce **384-dimensional vectors** to match the database schema.
|
||||
### Local Models
|
||||
|
||||
| Model | Dimensions | Use Case |
|
||||
|-------|------------|----------|
|
||||
| `BAAI/bge-small-en-v1.5` | 384 | Default, fast, good quality |
|
||||
| `sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2` | 384 | Multilingual (50+ languages) |
|
||||
|
||||
### OpenAI Models
|
||||
|
||||
| Model | Dimensions | Use Case |
|
||||
|-------|------------|----------|
|
||||
| `text-embedding-3-small` | 1536 | Default OpenAI, cost-effective |
|
||||
| `text-embedding-3-large` | 3072 | Higher quality, more expensive |
|
||||
| `text-embedding-ada-002` | 1536 | Legacy model |
|
||||
|
||||
### Cohere Models
|
||||
|
||||
| Model | Dimensions | Use Case |
|
||||
|-------|------------|----------|
|
||||
| `embed-english-v3.0` | 1024 | English text |
|
||||
| `embed-multilingual-v3.0` | 1024 | 100+ languages |
|
||||
|
||||
:::warning Embedding Dimensions
|
||||
Hindsight automatically detects the embedding dimension at startup and adjusts the database schema. Once memories are stored, you cannot change dimensions without losing data.
|
||||
:::
|
||||
|
||||
**Configuration:**
|
||||
**Configuration Examples:**
|
||||
|
||||
```bash
|
||||
# Local provider (default)
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=local
|
||||
export HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL=BAAI/bge-small-en-v1.5
|
||||
|
||||
# TEI provider (remote)
|
||||
# OpenAI
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=openai
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_API_KEY=sk-xxxxxxxxxxxx
|
||||
export HINDSIGHT_API_EMBEDDINGS_OPENAI_MODEL=text-embedding-3-small
|
||||
|
||||
# Cohere
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-api-key
|
||||
export HINDSIGHT_API_EMBEDDINGS_COHERE_MODEL=embed-english-v3.0
|
||||
|
||||
# TEI (self-hosted)
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=tei
|
||||
export HINDSIGHT_API_EMBEDDINGS_TEI_URL=http://localhost:8080
|
||||
|
||||
# LiteLLM proxy
|
||||
export HINDSIGHT_API_EMBEDDINGS_PROVIDER=litellm
|
||||
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
|
||||
export HINDSIGHT_API_EMBEDDINGS_LITELLM_MODEL=text-embedding-3-small
|
||||
```
|
||||
|
||||
See [Configuration](./configuration#embeddings) for all options including Azure OpenAI and custom endpoints.
|
||||
|
||||
---
|
||||
|
||||
## Cross-Encoder (Reranker)
|
||||
@@ -126,7 +168,18 @@ Reranks initial search results to improve precision.
|
||||
|
||||
**Default:** `cross-encoder/ms-marco-MiniLM-L-6-v2` (~85MB)
|
||||
|
||||
**Alternatives:**
|
||||
### Supported Providers
|
||||
|
||||
| Provider | Description | Best For |
|
||||
|----------|-------------|----------|
|
||||
| `local` | SentenceTransformers CrossEncoder (default) | Development, low latency |
|
||||
| `cohere` | Cohere rerank API | Production, high quality |
|
||||
| `tei` | HuggingFace Text Embeddings Inference | Production, self-hosted |
|
||||
| `flashrank` | FlashRank (lightweight, fast) | Resource-constrained environments |
|
||||
| `litellm` | LiteLLM proxy (unified gateway) | Multi-provider setups |
|
||||
| `rrf` | RRF-only (no neural reranking) | Testing, minimal resources |
|
||||
|
||||
### Local Models
|
||||
|
||||
| Model | Use Case |
|
||||
|-------|----------|
|
||||
@@ -134,14 +187,51 @@ Reranks initial search results to improve precision.
|
||||
| `cross-encoder/ms-marco-MiniLM-L-12-v2` | Higher accuracy |
|
||||
| `cross-encoder/mmarco-mMiniLMv2-L12-H384-v1` | Multilingual |
|
||||
|
||||
**Configuration:**
|
||||
### Cohere Models
|
||||
|
||||
| Model | Use Case |
|
||||
|-------|----------|
|
||||
| `rerank-english-v3.0` | English text |
|
||||
| `rerank-multilingual-v3.0` | 100+ languages |
|
||||
|
||||
### LiteLLM Supported Providers
|
||||
|
||||
LiteLLM supports multiple reranking providers via the `/rerank` endpoint:
|
||||
|
||||
| Provider | Model Example |
|
||||
|----------|---------------|
|
||||
| Cohere | `cohere/rerank-english-v3.0` |
|
||||
| Together AI | `together_ai/...` |
|
||||
| Voyage AI | `voyage/rerank-2` |
|
||||
| Jina AI | `jina_ai/...` |
|
||||
| AWS Bedrock | `bedrock/...` |
|
||||
|
||||
**Configuration Examples:**
|
||||
|
||||
```bash
|
||||
# Local provider (default)
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=local
|
||||
export HINDSIGHT_API_RERANKER_LOCAL_MODEL=cross-encoder/ms-marco-MiniLM-L-6-v2
|
||||
|
||||
# TEI provider (remote)
|
||||
# Cohere
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=cohere
|
||||
export HINDSIGHT_API_COHERE_API_KEY=your-api-key
|
||||
export HINDSIGHT_API_RERANKER_COHERE_MODEL=rerank-english-v3.0
|
||||
|
||||
# TEI (self-hosted)
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=tei
|
||||
export HINDSIGHT_API_RERANKER_TEI_URL=http://localhost:8081
|
||||
|
||||
# FlashRank (lightweight)
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=flashrank
|
||||
|
||||
# LiteLLM proxy
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=litellm
|
||||
export HINDSIGHT_API_LITELLM_API_BASE=http://localhost:4000
|
||||
export HINDSIGHT_API_RERANKER_LITELLM_MODEL=cohere/rerank-english-v3.0
|
||||
|
||||
# RRF-only (no neural reranking)
|
||||
export HINDSIGHT_API_RERANKER_PROVIDER=rrf
|
||||
```
|
||||
|
||||
See [Configuration](./configuration#reranker) for all options including Azure-hosted endpoints and batch settings.
|
||||
|
||||
@@ -183,4 +183,4 @@ Disposition creates **consistent character** across conversations while allowing
|
||||
|
||||
- [**Retain**](./retain) — How rich facts are stored
|
||||
- [**Recall**](./retrieval) — How multi-strategy search works
|
||||
- [API Reference: Reflect](./api/reflect) — Code examples and usage
|
||||
- [**Reflect API**](./api/reflect) — Code examples, parameters, and tag filtering
|
||||
|
||||
@@ -167,6 +167,18 @@ As facts accumulate about an entity, Hindsight synthesizes **observations** —
|
||||
|
||||
---
|
||||
|
||||
## Tagging Memories
|
||||
|
||||
Tags enable visibility scoping—useful when one memory bank serves multiple users but each should only see relevant memories.
|
||||
|
||||
- **Item tags**: Tag individual memories with specific scopes
|
||||
- **Document tags**: Apply tags to all items in a batch
|
||||
- **Tag filtering**: Filter during recall/reflect by tags
|
||||
|
||||
See [Retain API](./api/retain) for code examples and [Recall API](./api/recall) for filtering options.
|
||||
|
||||
---
|
||||
|
||||
## What You Get
|
||||
|
||||
After `retain()` completes:
|
||||
@@ -176,6 +188,7 @@ After `retain()` completes:
|
||||
- **Knowledge graph** with entity, temporal, semantic, and causal links
|
||||
- **Temporal grounding** for both historical and recency-based queries
|
||||
- **Background processing** that generates entity summaries
|
||||
- **Optional tags** for filtering during recall
|
||||
|
||||
All stored in your isolated **memory bank**, ready for `recall()` and `reflect()`.
|
||||
|
||||
@@ -185,4 +198,4 @@ All stored in your isolated **memory bank**, ready for `recall()` and `reflect()
|
||||
|
||||
- [**Recall**](./retrieval) — How multi-strategy search retrieves relevant memories
|
||||
- [**Reflect**](./reflect) — How disposition influences reasoning and opinion formation
|
||||
- [API Reference](./api/retain) — Code examples for retaining memories
|
||||
- [**Retain API**](./api/retain) — Code examples and parameters
|
||||
|
||||
@@ -133,7 +133,9 @@ Hindsight is built for AI agents, not humans. Traditional search systems return
|
||||
**Parameters you control:**
|
||||
- `max_tokens`: How much memory content to return (default: 4096 tokens)
|
||||
- `budget`: Search depth level (low, mid, high)
|
||||
- `fact_type`: Filter by world, experience, opinion, or all
|
||||
- `types`: Filter by world, experience, opinion, or all
|
||||
- `tags`: Filter memories by visibility tags
|
||||
- `tags_match`: How to match tags (see [Recall API](./api/recall) for all options)
|
||||
|
||||
### Expanding Context: Chunks and Entity Observations
|
||||
|
||||
@@ -241,3 +243,4 @@ See [Configuration → Retrieval](./configuration#retrieval) for available algor
|
||||
|
||||
- [**Retain**](./retain) — How memories are stored with rich context
|
||||
- [**Reflect**](./reflect) — How disposition influences reasoning
|
||||
- [**Recall API**](./api/recall) — Code examples, parameters, and tag filtering
|
||||
|
||||
@@ -116,6 +116,39 @@ results = client.recall(bank_id="my-bank", query="How are Alice and Bob connecte
|
||||
# [/docs:recall-budget-levels]
|
||||
|
||||
|
||||
# [docs:recall-with-tags]
|
||||
# Filter recall to only memories tagged for a specific user
|
||||
response = client.recall(
|
||||
bank_id="my-bank",
|
||||
query="What feedback did the user give?",
|
||||
tags=["user:alice"],
|
||||
tags_match="any" # OR matching, includes untagged (default)
|
||||
)
|
||||
# [/docs:recall-with-tags]
|
||||
|
||||
|
||||
# [docs:recall-tags-strict]
|
||||
# Strict mode: only return memories that have matching tags (exclude untagged)
|
||||
response = client.recall(
|
||||
bank_id="my-bank",
|
||||
query="What did the user say?",
|
||||
tags=["user:alice"],
|
||||
tags_match="any_strict" # OR matching, excludes untagged memories
|
||||
)
|
||||
# [/docs:recall-tags-strict]
|
||||
|
||||
|
||||
# [docs:recall-tags-all]
|
||||
# AND matching: require ALL specified tags to be present
|
||||
response = client.recall(
|
||||
bank_id="my-bank",
|
||||
query="What bugs were reported?",
|
||||
tags=["user:alice", "bug-report"],
|
||||
tags_match="all_strict" # Memory must have BOTH tags
|
||||
)
|
||||
# [/docs:recall-tags-all]
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Cleanup (not shown in docs)
|
||||
# =============================================================================
|
||||
|
||||
@@ -81,6 +81,17 @@ for fact in response.based_on or []:
|
||||
# [/docs:reflect-sources]
|
||||
|
||||
|
||||
# [docs:reflect-with-tags]
|
||||
# Filter reflection to only consider memories for a specific user
|
||||
response = client.reflect(
|
||||
bank_id="my-bank",
|
||||
query="What does this user think about our product?",
|
||||
tags=["user:alice"],
|
||||
tags_match="any_strict" # Only use memories tagged for this user
|
||||
)
|
||||
# [/docs:reflect-with-tags]
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Cleanup (not shown in docs)
|
||||
# =============================================================================
|
||||
|
||||
@@ -67,6 +67,39 @@ print(result.var_async) # True
|
||||
# [/docs:retain-async]
|
||||
|
||||
|
||||
# [docs:retain-with-tags]
|
||||
# Tag individual items for visibility scoping
|
||||
client.retain_batch(
|
||||
bank_id="my-bank",
|
||||
items=[
|
||||
{
|
||||
"content": "User Alice said she loves the new dashboard",
|
||||
"tags": ["user:alice", "feedback"]
|
||||
},
|
||||
{
|
||||
"content": "User Bob reported a bug in the search feature",
|
||||
"tags": ["user:bob", "bug-report"]
|
||||
}
|
||||
],
|
||||
document_id="user_feedback_001"
|
||||
)
|
||||
# [/docs:retain-with-tags]
|
||||
|
||||
|
||||
# [docs:retain-with-document-tags]
|
||||
# Apply tags to all items in a batch
|
||||
client.retain_batch(
|
||||
bank_id="my-bank",
|
||||
items=[
|
||||
{"content": "Alice mentioned she prefers dark mode"},
|
||||
{"content": "Bob asked about keyboard shortcuts"}
|
||||
],
|
||||
document_id="support_session_123",
|
||||
document_tags=["session:123", "support"] # Applied to all items
|
||||
)
|
||||
# [/docs:retain-with-document-tags]
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Cleanup (not shown in docs)
|
||||
# =============================================================================
|
||||
|
||||
@@ -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