feat(python-sdk): add tags filtering support to high-level client (#186)
Add tags and tags_match parameters to recall/reflect methods for filtering memories by visibility scope. Also add tags support to retain methods. Changes: - recall()/arecall(): add tags, tags_match parameters - reflect()/areflect(): add tags, tags_match parameters - retain()/aretain(): add tags parameter - retain_batch()/aretain_batch(): add document_tags parameter - Add TestTags test class with 7 tests
This commit is contained in:
parent
66abad61b8
commit
aebef9408b
3 changed files with 294 additions and 83 deletions
|
|
@ -25,17 +25,18 @@ Example:
|
|||
```
|
||||
"""
|
||||
|
||||
from .hindsight_client import Hindsight
|
||||
from hindsight_client_api.models.bank_profile_response import BankProfileResponse
|
||||
from hindsight_client_api.models.disposition_traits import DispositionTraits
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.recall_response import RecallResponse as _RecallResponse
|
||||
from hindsight_client_api.models.recall_result import RecallResult as _RecallResult
|
||||
from hindsight_client_api.models.reflect_fact import ReflectFact
|
||||
from hindsight_client_api.models.reflect_response import ReflectResponse
|
||||
|
||||
# Re-export response types for convenient access
|
||||
from hindsight_client_api.models.retain_response import RetainResponse
|
||||
from hindsight_client_api.models.recall_response import RecallResponse as _RecallResponse
|
||||
from hindsight_client_api.models.recall_result import RecallResult as _RecallResult
|
||||
from hindsight_client_api.models.reflect_response import ReflectResponse
|
||||
from hindsight_client_api.models.reflect_fact import ReflectFact
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.bank_profile_response import BankProfileResponse
|
||||
from hindsight_client_api.models.disposition_traits import DispositionTraits
|
||||
|
||||
from .hindsight_client import Hindsight
|
||||
|
||||
|
||||
# Add cleaner __repr__ and __iter__ for REPL usability
|
||||
|
|
|
|||
|
|
@ -6,23 +6,23 @@ easy-to-use interface on top of the auto-generated OpenAPI client.
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Optional, List, Dict, Any, Literal
|
||||
from datetime import datetime
|
||||
from typing import Any, Literal
|
||||
|
||||
import hindsight_client_api
|
||||
from hindsight_client_api.api import memory_api, banks_api
|
||||
from hindsight_client_api.api import banks_api, memory_api
|
||||
from hindsight_client_api.models import (
|
||||
recall_request,
|
||||
retain_request,
|
||||
memory_item,
|
||||
recall_request,
|
||||
reflect_request,
|
||||
retain_request,
|
||||
)
|
||||
from hindsight_client_api.models.retain_response import RetainResponse
|
||||
from hindsight_client_api.models.bank_profile_response import BankProfileResponse
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.recall_response import RecallResponse
|
||||
from hindsight_client_api.models.recall_result import RecallResult
|
||||
from hindsight_client_api.models.reflect_response import ReflectResponse
|
||||
from hindsight_client_api.models.list_memory_units_response import ListMemoryUnitsResponse
|
||||
from hindsight_client_api.models.bank_profile_response import BankProfileResponse
|
||||
from hindsight_client_api.models.retain_response import RetainResponse
|
||||
|
||||
|
||||
def _run_async(coro):
|
||||
|
|
@ -63,7 +63,7 @@ class Hindsight:
|
|||
```
|
||||
"""
|
||||
|
||||
def __init__(self, base_url: str, api_key: Optional[str] = None, timeout: float = 30.0):
|
||||
def __init__(self, base_url: str, api_key: str | None = None, timeout: float = 30.0):
|
||||
"""
|
||||
Initialize the Hindsight client.
|
||||
|
||||
|
|
@ -110,12 +110,12 @@ class Hindsight:
|
|||
self,
|
||||
bank_id: str,
|
||||
content: str,
|
||||
timestamp: Optional[datetime] = None,
|
||||
context: Optional[str] = None,
|
||||
document_id: Optional[str] = None,
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
entities: Optional[List[Dict[str, str]]] = None,
|
||||
tags: Optional[List[str]] = None,
|
||||
timestamp: datetime | None = None,
|
||||
context: str | None = None,
|
||||
document_id: str | None = None,
|
||||
metadata: dict[str, str] | None = None,
|
||||
entities: list[dict[str, str]] | None = None,
|
||||
tags: list[str] | None = None,
|
||||
) -> RetainResponse:
|
||||
"""
|
||||
Store a single memory (simplified interface).
|
||||
|
|
@ -128,24 +128,33 @@ 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
|
||||
tags: Optional list of tags for filtering memories during recall/reflect
|
||||
|
||||
Returns:
|
||||
RetainResponse with success status
|
||||
"""
|
||||
return self.retain_batch(
|
||||
bank_id=bank_id,
|
||||
items=[{"content": content, "timestamp": timestamp, "context": context, "metadata": metadata, "entities": entities, "tags": tags}],
|
||||
items=[
|
||||
{
|
||||
"content": content,
|
||||
"timestamp": timestamp,
|
||||
"context": context,
|
||||
"metadata": metadata,
|
||||
"entities": entities,
|
||||
"tags": tags,
|
||||
}
|
||||
],
|
||||
document_id=document_id,
|
||||
)
|
||||
|
||||
def retain_batch(
|
||||
self,
|
||||
bank_id: str,
|
||||
items: List[Dict[str, Any]],
|
||||
document_id: Optional[str] = None,
|
||||
items: list[dict[str, Any]],
|
||||
document_id: str | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
retain_async: bool = False,
|
||||
document_tags: Optional[List[str]] = None,
|
||||
) -> RetainResponse:
|
||||
"""
|
||||
Store multiple memories in batch.
|
||||
|
|
@ -154,8 +163,8 @@ class Hindsight:
|
|||
bank_id: The memory bank ID
|
||||
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)
|
||||
document_tags: Optional list of tags applied to all items in this batch (merged with per-item tags)
|
||||
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
|
||||
|
|
@ -166,10 +175,7 @@ class Hindsight:
|
|||
for item in items:
|
||||
entities = None
|
||||
if item.get("entities"):
|
||||
entities = [
|
||||
EntityInput(text=e["text"], type=e.get("type"))
|
||||
for e in item["entities"]
|
||||
]
|
||||
entities = [EntityInput(text=e["text"], type=e.get("type")) for e in item["entities"]]
|
||||
memory_items.append(
|
||||
memory_item.MemoryItem(
|
||||
content=item["content"],
|
||||
|
|
@ -195,17 +201,17 @@ class Hindsight:
|
|||
self,
|
||||
bank_id: str,
|
||||
query: str,
|
||||
types: Optional[List[str]] = None,
|
||||
types: list[str] | None = None,
|
||||
max_tokens: int = 4096,
|
||||
budget: str = "mid",
|
||||
trace: bool = False,
|
||||
query_timestamp: Optional[str] = None,
|
||||
query_timestamp: str | None = None,
|
||||
include_entities: bool = False,
|
||||
max_entity_tokens: int = 500,
|
||||
include_chunks: bool = False,
|
||||
max_chunk_tokens: int = 8192,
|
||||
tags: Optional[List[str]] = None,
|
||||
tags_match: str = "any",
|
||||
tags: list[str] | None = None,
|
||||
tags_match: Literal["any", "all", "any_strict", "all_strict"] = "any",
|
||||
) -> RecallResponse:
|
||||
"""
|
||||
Recall memories using semantic similarity.
|
||||
|
|
@ -223,16 +229,18 @@ class Hindsight:
|
|||
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'
|
||||
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
|
||||
"""
|
||||
from hindsight_client_api.models import include_options, entity_include_options, chunk_include_options
|
||||
from hindsight_client_api.models import chunk_include_options, entity_include_options, include_options
|
||||
|
||||
include_opts = include_options.IncludeOptions(
|
||||
entities=entity_include_options.EntityIncludeOptions(max_tokens=max_entity_tokens) if include_entities else None,
|
||||
entities=entity_include_options.EntityIncludeOptions(max_tokens=max_entity_tokens)
|
||||
if include_entities
|
||||
else None,
|
||||
chunks=chunk_include_options.ChunkIncludeOptions(max_tokens=max_chunk_tokens) if include_chunks else None,
|
||||
)
|
||||
|
||||
|
|
@ -255,11 +263,11 @@ class Hindsight:
|
|||
bank_id: str,
|
||||
query: str,
|
||||
budget: str = "low",
|
||||
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",
|
||||
context: str | None = None,
|
||||
max_tokens: int | None = None,
|
||||
response_schema: dict[str, Any] | None = None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: Literal["any", "all", "any_strict", "all_strict"] = "any",
|
||||
) -> ReflectResponse:
|
||||
"""
|
||||
Generate a contextual answer based on bank identity and memories.
|
||||
|
|
@ -274,8 +282,8 @@ class Hindsight:
|
|||
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'
|
||||
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
|
||||
|
|
@ -296,26 +304,28 @@ class Hindsight:
|
|||
def list_memories(
|
||||
self,
|
||||
bank_id: str,
|
||||
type: Optional[str] = None,
|
||||
search_query: Optional[str] = None,
|
||||
type: str | None = None,
|
||||
search_query: str | None = None,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
) -> ListMemoryUnitsResponse:
|
||||
"""List memory units with pagination."""
|
||||
return _run_async(self._memory_api.list_memories(
|
||||
bank_id=bank_id,
|
||||
type=type,
|
||||
q=search_query,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
))
|
||||
return _run_async(
|
||||
self._memory_api.list_memories(
|
||||
bank_id=bank_id,
|
||||
type=type,
|
||||
q=search_query,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
)
|
||||
|
||||
def create_bank(
|
||||
self,
|
||||
bank_id: str,
|
||||
name: Optional[str] = None,
|
||||
background: Optional[str] = None,
|
||||
disposition: Optional[Dict[str, float]] = None,
|
||||
name: str | None = None,
|
||||
background: str | None = None,
|
||||
disposition: dict[str, float] | None = None,
|
||||
) -> BankProfileResponse:
|
||||
"""Create or update a memory bank."""
|
||||
from hindsight_client_api.models import create_bank_request, disposition_traits
|
||||
|
|
@ -357,8 +367,9 @@ class Hindsight:
|
|||
async def aretain_batch(
|
||||
self,
|
||||
bank_id: str,
|
||||
items: List[Dict[str, Any]],
|
||||
document_id: Optional[str] = None,
|
||||
items: list[dict[str, Any]],
|
||||
document_id: str | None = None,
|
||||
document_tags: list[str] | None = None,
|
||||
retain_async: bool = False,
|
||||
) -> RetainResponse:
|
||||
"""
|
||||
|
|
@ -366,8 +377,9 @@ class Hindsight:
|
|||
|
||||
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)
|
||||
document_tags: Optional list of tags applied to all items in this batch (merged with per-item tags)
|
||||
retain_async: If True, process asynchronously in background (default: False)
|
||||
|
||||
Returns:
|
||||
|
|
@ -379,10 +391,7 @@ class Hindsight:
|
|||
for item in items:
|
||||
entities = None
|
||||
if item.get("entities"):
|
||||
entities = [
|
||||
EntityInput(text=e["text"], type=e.get("type"))
|
||||
for e in item["entities"]
|
||||
]
|
||||
entities = [EntityInput(text=e["text"], type=e.get("type")) for e in item["entities"]]
|
||||
memory_items.append(
|
||||
memory_item.MemoryItem(
|
||||
content=item["content"],
|
||||
|
|
@ -392,12 +401,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 await self._memory_api.retain_memories(bank_id, request_obj)
|
||||
|
|
@ -406,11 +417,12 @@ class Hindsight:
|
|||
self,
|
||||
bank_id: str,
|
||||
content: str,
|
||||
timestamp: Optional[datetime] = None,
|
||||
context: Optional[str] = None,
|
||||
document_id: Optional[str] = None,
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
entities: Optional[List[Dict[str, str]]] = None,
|
||||
timestamp: datetime | None = None,
|
||||
context: str | None = None,
|
||||
document_id: str | None = None,
|
||||
metadata: dict[str, str] | None = None,
|
||||
entities: list[dict[str, str]] | None = None,
|
||||
tags: list[str] | None = None,
|
||||
) -> RetainResponse:
|
||||
"""
|
||||
Store a single memory (async).
|
||||
|
|
@ -423,13 +435,23 @@ 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 filtering memories during recall/reflect
|
||||
|
||||
Returns:
|
||||
RetainResponse with success status
|
||||
"""
|
||||
return await self.aretain_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,
|
||||
)
|
||||
|
||||
|
|
@ -437,10 +459,12 @@ class Hindsight:
|
|||
self,
|
||||
bank_id: str,
|
||||
query: str,
|
||||
types: Optional[List[str]] = None,
|
||||
types: list[str] | None = None,
|
||||
max_tokens: int = 4096,
|
||||
budget: str = "mid",
|
||||
) -> List[RecallResult]:
|
||||
tags: list[str] | None = None,
|
||||
tags_match: Literal["any", "all", "any_strict", "all_strict"] = "any",
|
||||
) -> list[RecallResult]:
|
||||
"""
|
||||
Recall memories using semantic similarity (async).
|
||||
|
||||
|
|
@ -450,6 +474,9 @@ class Hindsight:
|
|||
types: Optional list of fact types to filter (world, experience, opinion, observation)
|
||||
max_tokens: Maximum tokens in results (default: 4096)
|
||||
budget: Budget level for recall - "low", "mid", or "high" (default: "mid")
|
||||
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:
|
||||
List of RecallResult objects
|
||||
|
|
@ -460,17 +487,21 @@ class Hindsight:
|
|||
budget=budget,
|
||||
max_tokens=max_tokens,
|
||||
trace=False,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
|
||||
response = await self._memory_api.recall_memories(bank_id, request_obj)
|
||||
return response.results if hasattr(response, 'results') else []
|
||||
return response.results if hasattr(response, "results") else []
|
||||
|
||||
async def areflect(
|
||||
self,
|
||||
bank_id: str,
|
||||
query: str,
|
||||
budget: str = "low",
|
||||
context: Optional[str] = None,
|
||||
context: str | None = None,
|
||||
tags: list[str] | None = None,
|
||||
tags_match: Literal["any", "all", "any_strict", "all_strict"] = "any",
|
||||
) -> ReflectResponse:
|
||||
"""
|
||||
Generate a contextual answer based on bank identity and memories (async).
|
||||
|
|
@ -480,6 +511,9 @@ class Hindsight:
|
|||
query: The question or prompt
|
||||
budget: Budget level for reflection - "low", "mid", or "high" (default: "low")
|
||||
context: Optional additional context
|
||||
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 and optionally facts used
|
||||
|
|
@ -488,6 +522,8 @@ class Hindsight:
|
|||
query=query,
|
||||
budget=budget,
|
||||
context=context,
|
||||
tags=tags,
|
||||
tags_match=tags_match,
|
||||
)
|
||||
|
||||
return await self._memory_api.reflect(bank_id, request_obj)
|
||||
|
|
|
|||
|
|
@ -6,10 +6,11 @@ These tests require a running Hindsight API server.
|
|||
|
||||
import os
|
||||
import uuid
|
||||
import pytest
|
||||
from datetime import datetime
|
||||
from hindsight_client import Hindsight
|
||||
|
||||
import pytest
|
||||
|
||||
from hindsight_client import Hindsight
|
||||
|
||||
# Test configuration
|
||||
HINDSIGHT_API_URL = os.getenv("HINDSIGHT_API_URL", "http://localhost:8888")
|
||||
|
|
@ -191,14 +192,14 @@ class TestReflect:
|
|||
When response_schema is provided, the response returns structured_output
|
||||
field parsed according to the provided JSON schema.
|
||||
"""
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
# Define schema using Pydantic model
|
||||
class RecommendationResponse(BaseModel):
|
||||
recommendation: str
|
||||
reasons: list[str]
|
||||
confidence: Optional[str] = None # Optional for LLM flexibility
|
||||
confidence: str | None = None # Optional for LLM flexibility
|
||||
|
||||
response = client.reflect(
|
||||
bank_id=bank_id,
|
||||
|
|
@ -224,9 +225,7 @@ class TestListMemories:
|
|||
"""Setup: Store some test memories synchronously."""
|
||||
client.retain_batch(
|
||||
bank_id=bank_id,
|
||||
items=[
|
||||
{"content": f"Alice likes topic number {i}"} for i in range(5)
|
||||
],
|
||||
items=[{"content": f"Alice likes topic number {i}"} for i in range(5)],
|
||||
retain_async=False, # Wait for fact extraction to complete
|
||||
)
|
||||
|
||||
|
|
@ -359,6 +358,7 @@ class TestDocuments:
|
|||
def test_delete_document(self, client, bank_id):
|
||||
"""Test deleting a document."""
|
||||
import asyncio
|
||||
|
||||
from hindsight_client_api import ApiClient, Configuration
|
||||
from hindsight_client_api.api import DocumentsApi
|
||||
|
||||
|
|
@ -390,6 +390,7 @@ class TestDocuments:
|
|||
def test_get_document(self, client, bank_id):
|
||||
"""Test getting a document."""
|
||||
import asyncio
|
||||
|
||||
from hindsight_client_api import ApiClient, Configuration
|
||||
from hindsight_client_api.api import DocumentsApi
|
||||
|
||||
|
|
@ -432,6 +433,7 @@ class TestEntities:
|
|||
def test_list_entities(self, client, bank_id):
|
||||
"""Test listing entities."""
|
||||
import asyncio
|
||||
|
||||
from hindsight_client_api import ApiClient, Configuration
|
||||
from hindsight_client_api.api import EntitiesApi
|
||||
|
||||
|
|
@ -456,6 +458,7 @@ class TestEntities:
|
|||
def test_list_entities_with_pagination(self, client, bank_id):
|
||||
"""Test listing entities with pagination parameters."""
|
||||
import asyncio
|
||||
|
||||
from hindsight_client_api import ApiClient, Configuration
|
||||
from hindsight_client_api.api import EntitiesApi
|
||||
|
||||
|
|
@ -482,6 +485,7 @@ class TestEntities:
|
|||
def test_get_entity(self, client, bank_id):
|
||||
"""Test getting a specific entity."""
|
||||
import asyncio
|
||||
|
||||
from hindsight_client_api import ApiClient, Configuration
|
||||
from hindsight_client_api.api import EntitiesApi
|
||||
|
||||
|
|
@ -507,6 +511,175 @@ class TestEntities:
|
|||
assert entity is not None
|
||||
assert entity.id == entity_id
|
||||
|
||||
def test_regenerate_entity_observations(self, client, bank_id):
|
||||
"""Test regenerating observations for an entity."""
|
||||
import asyncio
|
||||
|
||||
from hindsight_client_api import ApiClient, Configuration
|
||||
from hindsight_client_api.api import EntitiesApi
|
||||
|
||||
async def do_test():
|
||||
config = Configuration(host=HINDSIGHT_API_URL)
|
||||
api_client = ApiClient(config)
|
||||
api = EntitiesApi(api_client)
|
||||
|
||||
# First list entities to get an ID
|
||||
list_response = await api.list_entities(bank_id=bank_id)
|
||||
|
||||
if list_response.items and len(list_response.items) > 0:
|
||||
entity_id = list_response.items[0].id
|
||||
|
||||
# Regenerate observations
|
||||
result = await api.regenerate_entity_observations(
|
||||
bank_id=bank_id,
|
||||
entity_id=entity_id,
|
||||
)
|
||||
return entity_id, result
|
||||
return None, None
|
||||
|
||||
entity_id, result = asyncio.get_event_loop().run_until_complete(do_test())
|
||||
|
||||
if entity_id:
|
||||
assert result is not None
|
||||
assert result.id == entity_id
|
||||
|
||||
|
||||
class TestTags:
|
||||
"""Tests for tags filtering functionality."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup_memories(self, client, bank_id):
|
||||
"""Setup: Store memories with different tags."""
|
||||
client.retain_batch(
|
||||
bank_id=bank_id,
|
||||
items=[
|
||||
{"content": "Project X meeting notes from Monday", "tags": ["project_x", "meetings"]},
|
||||
{"content": "Project X design document", "tags": ["project_x", "docs"]},
|
||||
{"content": "Project Y sprint planning", "tags": ["project_y", "meetings"]},
|
||||
{"content": "General company announcement", "tags": ["company"]},
|
||||
{"content": "Untagged memory about random things"}, # no tags
|
||||
],
|
||||
retain_async=False,
|
||||
)
|
||||
|
||||
def test_recall_with_tags_any(self, client, bank_id):
|
||||
"""Test recall with tags using 'any' match (includes untagged)."""
|
||||
response = client.recall(
|
||||
bank_id=bank_id,
|
||||
query="What are the documents?",
|
||||
tags=["project_x"],
|
||||
tags_match="any",
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.results is not None
|
||||
# Should include project_x tagged items and potentially untagged items
|
||||
result_texts = [r.text.lower() for r in response.results]
|
||||
assert any("project x" in text for text in result_texts)
|
||||
|
||||
def test_recall_with_tags_any_strict(self, client, bank_id):
|
||||
"""Test recall with tags using 'any_strict' match (excludes untagged)."""
|
||||
response = client.recall(
|
||||
bank_id=bank_id,
|
||||
query="meetings",
|
||||
tags=["project_x"],
|
||||
tags_match="any_strict",
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.results is not None
|
||||
# All results should have project_x tag - no untagged items
|
||||
result_texts = [r.text.lower() for r in response.results]
|
||||
# Should find project_x items only
|
||||
for text in result_texts:
|
||||
assert "project x" in text or "untagged" not in text
|
||||
|
||||
def test_recall_with_tags_all_strict(self, client, bank_id):
|
||||
"""Test recall with tags using 'all_strict' match (AND matching)."""
|
||||
response = client.recall(
|
||||
bank_id=bank_id,
|
||||
query="meeting notes",
|
||||
tags=["project_x", "meetings"],
|
||||
tags_match="all_strict",
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.results is not None
|
||||
# Should only return items tagged with BOTH project_x AND meetings
|
||||
if len(response.results) > 0:
|
||||
result_texts = [r.text.lower() for r in response.results]
|
||||
# The "Project X meeting notes" should be found
|
||||
assert any("project x" in text and "meeting" in text for text in result_texts)
|
||||
|
||||
def test_recall_with_multiple_tags_any(self, client, bank_id):
|
||||
"""Test recall with multiple tags using 'any' match (OR)."""
|
||||
response = client.recall(
|
||||
bank_id=bank_id,
|
||||
query="What's happening?",
|
||||
tags=["project_x", "project_y"],
|
||||
tags_match="any_strict",
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.results is not None
|
||||
# Should include items from both project_x and project_y
|
||||
result_texts = [r.text.lower() for r in response.results]
|
||||
has_project_x = any("project x" in text for text in result_texts)
|
||||
has_project_y = any("project y" in text for text in result_texts)
|
||||
# At least one of them should be present
|
||||
assert has_project_x or has_project_y
|
||||
|
||||
def test_reflect_with_tags(self, client, bank_id):
|
||||
"""Test reflect with tags filtering."""
|
||||
response = client.reflect(
|
||||
bank_id=bank_id,
|
||||
query="Summarize project X activities",
|
||||
tags=["project_x"],
|
||||
tags_match="any_strict",
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.text is not None
|
||||
assert len(response.text) > 0
|
||||
|
||||
def test_retain_with_tags(self, client, bank_id):
|
||||
"""Test storing a memory with tags."""
|
||||
response = client.retain(
|
||||
bank_id=bank_id,
|
||||
content="New feature implementation for project Z",
|
||||
tags=["project_z", "features"],
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.success is True
|
||||
|
||||
# Verify we can recall it with the tag
|
||||
recall_response = client.recall(
|
||||
bank_id=bank_id,
|
||||
query="project Z features",
|
||||
tags=["project_z"],
|
||||
tags_match="any_strict",
|
||||
)
|
||||
assert recall_response is not None
|
||||
result_texts = [r.text.lower() for r in recall_response.results]
|
||||
assert any("project z" in text for text in result_texts)
|
||||
|
||||
def test_retain_batch_with_document_tags(self, client, bank_id):
|
||||
"""Test batch retain with document-level tags."""
|
||||
response = client.retain_batch(
|
||||
bank_id=bank_id,
|
||||
items=[
|
||||
{"content": "First item in batch"},
|
||||
{"content": "Second item in batch"},
|
||||
],
|
||||
document_tags=["batch_import", "test_data"],
|
||||
retain_async=False,
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.success is True
|
||||
assert response.items_count == 2
|
||||
|
||||
|
||||
class TestDeleteBank:
|
||||
"""Tests for bank deletion."""
|
||||
|
|
@ -514,6 +687,7 @@ class TestDeleteBank:
|
|||
def test_delete_bank(self, client):
|
||||
"""Test deleting a bank."""
|
||||
import asyncio
|
||||
|
||||
from hindsight_client_api import ApiClient, Configuration
|
||||
from hindsight_client_api.api import BanksApi
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue