fix: retain async with timestamp might fails (#253)
This commit is contained in:
parent
c33b9b8bb2
commit
cbb8fc6723
4 changed files with 113 additions and 10 deletions
|
|
@ -759,6 +759,10 @@ async def _extract_facts_from_chunk(
|
|||
|
||||
# Build user message with metadata and chunk content in a clear format
|
||||
# Format event_date with day of week for better temporal reasoning
|
||||
# Handle both datetime objects and ISO string formats (from deserialized async tasks)
|
||||
from .orchestrator import parse_datetime_flexible
|
||||
|
||||
event_date = parse_datetime_flexible(event_date)
|
||||
event_date_formatted = event_date.strftime("%A, %B %d, %Y") # e.g., "Monday, June 10, 2024"
|
||||
user_message = f"""Extract facts from the following text chunk.
|
||||
|
||||
|
|
@ -1346,6 +1350,8 @@ def _add_temporal_offsets(facts: list[ExtractedFactType], contents: list[RetainC
|
|||
|
||||
Modifies facts in place.
|
||||
"""
|
||||
from .orchestrator import parse_datetime_flexible
|
||||
|
||||
# Group facts by content_index
|
||||
current_content_idx = 0
|
||||
content_fact_start = 0
|
||||
|
|
@ -1360,10 +1366,10 @@ def _add_temporal_offsets(facts: list[ExtractedFactType], contents: list[RetainC
|
|||
fact_position = i - content_fact_start
|
||||
offset = timedelta(seconds=fact_position * SECONDS_PER_FACT)
|
||||
|
||||
# Apply offset to all temporal fields
|
||||
# Apply offset to all temporal fields (handle both datetime objects and ISO strings)
|
||||
if fact.occurred_start:
|
||||
fact.occurred_start = fact.occurred_start + offset
|
||||
fact.occurred_start = parse_datetime_flexible(fact.occurred_start) + offset
|
||||
if fact.occurred_end:
|
||||
fact.occurred_end = fact.occurred_end + offset
|
||||
fact.occurred_end = parse_datetime_flexible(fact.occurred_end) + offset
|
||||
if fact.mentioned_at:
|
||||
fact.mentioned_at = fact.mentioned_at + offset
|
||||
fact.mentioned_at = parse_datetime_flexible(fact.mentioned_at) + offset
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import logging
|
|||
import time
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from ..db_utils import acquire_with_retry
|
||||
from . import bank_utils
|
||||
|
|
@ -18,6 +19,39 @@ def utcnow():
|
|||
return datetime.now(UTC)
|
||||
|
||||
|
||||
def parse_datetime_flexible(value: Any) -> datetime:
|
||||
"""
|
||||
Parse a datetime value that could be either a datetime object or an ISO string.
|
||||
|
||||
This handles datetime values from both direct Python calls and deserialized JSON
|
||||
(where datetime objects are serialized as ISO strings).
|
||||
|
||||
Args:
|
||||
value: Either a datetime object or an ISO format string
|
||||
|
||||
Returns:
|
||||
datetime object (timezone-aware)
|
||||
|
||||
Raises:
|
||||
TypeError: If value is neither datetime nor string
|
||||
ValueError: If string is not a valid ISO datetime
|
||||
"""
|
||||
if isinstance(value, datetime):
|
||||
# Ensure timezone-aware
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=UTC)
|
||||
return value
|
||||
elif isinstance(value, str):
|
||||
# Parse ISO format string (handles both 'Z' and '+00:00' timezone formats)
|
||||
dt = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
# Ensure timezone-aware
|
||||
if dt.tzinfo is None:
|
||||
return dt.replace(tzinfo=UTC)
|
||||
return dt
|
||||
else:
|
||||
raise TypeError(f"Expected datetime or string, got {type(value).__name__}")
|
||||
|
||||
|
||||
from ..response_models import TokenUsage
|
||||
from . import (
|
||||
chunk_storage,
|
||||
|
|
@ -89,10 +123,18 @@ async def retain_batch(
|
|||
# Merge item-level tags with document-level tags
|
||||
item_tags = item.get("tags", []) or []
|
||||
merged_tags = list(set(item_tags + (document_tags or [])))
|
||||
|
||||
# Handle event_date: parse flexibly (handles both datetime objects and ISO strings)
|
||||
event_date_value = item.get("event_date")
|
||||
if event_date_value:
|
||||
event_date_value = parse_datetime_flexible(event_date_value)
|
||||
else:
|
||||
event_date_value = utcnow()
|
||||
|
||||
content = RetainContent(
|
||||
content=item["content"],
|
||||
context=item.get("context", ""),
|
||||
event_date=item.get("event_date") or utcnow(),
|
||||
event_date=event_date_value,
|
||||
metadata=item.get("metadata", {}),
|
||||
entities=item.get("entities", []),
|
||||
tags=merged_tags,
|
||||
|
|
|
|||
|
|
@ -1174,3 +1174,58 @@ async def test_retain_with_multiple_timestamps(api_client, test_bank_id):
|
|||
data = response.json()
|
||||
assert data["success"] is True
|
||||
assert data["items_count"] == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retain_with_timestamp_async_complete_processing(api_client, test_bank_id):
|
||||
"""Test that async retain with timestamp completes full processing including fact extraction."""
|
||||
# Submit async retain with timestamp
|
||||
response = await api_client.post(
|
||||
f"/v1/default/banks/{test_bank_id}/memories",
|
||||
json={
|
||||
"items": [
|
||||
{
|
||||
"content": "The quarterly meeting was held on January 30th 2026",
|
||||
"context": "meetings",
|
||||
"timestamp": "2026-01-30T11:45:00Z"
|
||||
}
|
||||
],
|
||||
"async": True
|
||||
}
|
||||
)
|
||||
|
||||
assert response.status_code == 200, f"Expected 200, got {response.status_code}: {response.text}"
|
||||
data = response.json()
|
||||
assert data["success"] is True
|
||||
assert data["async"] is True
|
||||
operation_id = data["operation_id"]
|
||||
|
||||
# Wait for async processing to complete (poll operation status)
|
||||
max_wait_seconds = 30
|
||||
poll_interval = 0.5
|
||||
elapsed = 0
|
||||
operation_completed = False
|
||||
|
||||
while elapsed < max_wait_seconds:
|
||||
response = await api_client.get(f"/v1/default/banks/{test_bank_id}/operations/{operation_id}")
|
||||
if response.status_code == 200:
|
||||
op_status = response.json()
|
||||
if op_status.get("status") == "completed":
|
||||
operation_completed = True
|
||||
break
|
||||
elif op_status.get("status") == "failed":
|
||||
raise AssertionError(f"Operation failed: {op_status.get('error_message')}")
|
||||
|
||||
await asyncio.sleep(poll_interval)
|
||||
elapsed += poll_interval
|
||||
|
||||
assert operation_completed, f"Async operation did not complete within {max_wait_seconds} seconds"
|
||||
|
||||
# Verify memories were actually stored
|
||||
response = await api_client.get(
|
||||
f"/v1/default/banks/{test_bank_id}/memories/list",
|
||||
params={"limit": 10}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
items = response.json()["items"]
|
||||
assert len(items) > 0, "Should have stored memories after async processing"
|
||||
|
|
|
|||
10
uv.lock
10
uv.lock
|
|
@ -1295,7 +1295,7 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "hindsight-all"
|
||||
version = "0.4.3"
|
||||
version = "0.4.4"
|
||||
source = { editable = "hindsight" }
|
||||
dependencies = [
|
||||
{ name = "hindsight-api" },
|
||||
|
|
@ -1319,7 +1319,7 @@ provides-extras = ["test"]
|
|||
|
||||
[[package]]
|
||||
name = "hindsight-api"
|
||||
version = "0.4.3"
|
||||
version = "0.4.4"
|
||||
source = { editable = "hindsight-api" }
|
||||
dependencies = [
|
||||
{ name = "aiohttp" },
|
||||
|
|
@ -1449,7 +1449,7 @@ dev = [
|
|||
|
||||
[[package]]
|
||||
name = "hindsight-client"
|
||||
version = "0.4.3"
|
||||
version = "0.4.4"
|
||||
source = { editable = "hindsight-clients/python" }
|
||||
dependencies = [
|
||||
{ name = "aiohttp" },
|
||||
|
|
@ -1483,7 +1483,7 @@ provides-extras = ["test"]
|
|||
|
||||
[[package]]
|
||||
name = "hindsight-dev"
|
||||
version = "0.4.3"
|
||||
version = "0.4.4"
|
||||
source = { editable = "hindsight-dev" }
|
||||
dependencies = [
|
||||
{ name = "hindsight-api" },
|
||||
|
|
@ -1529,7 +1529,7 @@ dev = [
|
|||
|
||||
[[package]]
|
||||
name = "hindsight-embed"
|
||||
version = "0.4.3"
|
||||
version = "0.4.4"
|
||||
source = { editable = "hindsight-embed" }
|
||||
dependencies = [
|
||||
{ name = "httpx" },
|
||||
|
|
|
|||
Loading…
Reference in a new issue