diff --git a/src/engram/_models/memory.py b/src/engram/_models/memory.py index b9c4753..0158bef 100644 --- a/src/engram/_models/memory.py +++ b/src/engram/_models/memory.py @@ -2,6 +2,7 @@ from collections.abc import Iterator, Sequence from dataclasses import dataclass, field +from datetime import datetime from typing import Any, Literal, TypeAlias @@ -27,6 +28,8 @@ class StringInput: """String input to extract memories from.""" content: str | list[str] + created_at: str | datetime | None = None + updated_at: str | datetime | None = None @dataclass(slots=True) @@ -68,7 +71,7 @@ class MessageInput: role: Literal["user", "assistant", "system", "tool", "developer"] content: str = "" - created_at: str | None = None + created_at: str | datetime | None = None tool_call_id: str | None = None name: str | None = None tool_calls: list[ToolCallInput] | None = None @@ -76,12 +79,12 @@ class MessageInput: @dataclass(slots=True) class ConversationInput: - """Conversation input that bypasses the extraction pipeline.""" + """Conversation input to extract memories from.""" messages: list[MessageInput] metadata: dict[str, Any] | None = None - created_at: str | None = None - updated_at: str | None = None + created_at: str | datetime | None = None + updated_at: str | datetime | None = None # Type alias for the input_data argument to memories.add() diff --git a/src/engram/_serialization/_builders.py b/src/engram/_serialization/_builders.py index b4e3289..a273c38 100644 --- a/src/engram/_serialization/_builders.py +++ b/src/engram/_serialization/_builders.py @@ -1,5 +1,6 @@ from __future__ import annotations +from datetime import UTC, datetime from typing import Any, TypeAlias from .._models import ( @@ -33,10 +34,7 @@ def _serialize_input(input_data: AddInput) -> dict[str, Any]: if isinstance(input_data, str): return {"string": {"content": [input_data]}} if isinstance(input_data, StringInput): - if isinstance(input_data.content, list): - return {"string": {"content": input_data.content}} - else: - return {"string": {"content": [input_data.content]}} + return _serialize_string_input(input_data) if isinstance(input_data, PreExtractedInput): items = [{"content": item.content, "topic": item.topic} for item in input_data.items] return {"pre_extracted": {"items": items}} @@ -49,12 +47,31 @@ def _serialize_input(input_data: AddInput) -> dict[str, Any]: raise TypeError(f"Unsupported input type: {type(input_data)}") # pragma: no cover +def _serialize_timestamp(value: str | datetime) -> str: + """Format a timestamp as RFC 3339 (e.g. "2024-01-01T00:00:00Z").""" + if isinstance(value, str): + return value + if value.tzinfo is None: + value = value.replace(tzinfo=UTC) + return value.isoformat().replace("+00:00", "Z") + + +def _serialize_string_input(input_data: StringInput) -> dict[str, Any]: + content = input_data.content if isinstance(input_data.content, list) else [input_data.content] + body: dict[str, Any] = {"content": content} + if input_data.created_at is not None: + body["created_at"] = _serialize_timestamp(input_data.created_at) + if input_data.updated_at is not None: + body["updated_at"] = _serialize_timestamp(input_data.updated_at) + return {"string": body} + + def _serialize_conversation_content(content: ConversationInput) -> dict[str, Any]: messages = [] for msg in content.messages: m: dict[str, Any] = {"role": msg.role, "content": msg.content} if msg.created_at is not None: - m["created_at"] = msg.created_at + m["created_at"] = _serialize_timestamp(msg.created_at) if msg.tool_call_id is not None: m["tool_call_id"] = msg.tool_call_id if msg.name is not None: @@ -66,9 +83,9 @@ def _serialize_conversation_content(content: ConversationInput) -> dict[str, Any if content.metadata is not None: conversation["metadata"] = content.metadata if content.created_at is not None: - conversation["created_at"] = content.created_at + conversation["created_at"] = _serialize_timestamp(content.created_at) if content.updated_at is not None: - conversation["updated_at"] = content.updated_at + conversation["updated_at"] = _serialize_timestamp(content.updated_at) return {"conversation": conversation} diff --git a/tests/test_serialization.py b/tests/test_serialization.py index f78da74..4464033 100644 --- a/tests/test_serialization.py +++ b/tests/test_serialization.py @@ -1,3 +1,5 @@ +from datetime import UTC, datetime, timedelta, timezone + from engram._models import ( ConversationInput, MessageInput, @@ -94,6 +96,51 @@ def test_build_add_body_string_content_with_options() -> None: } +def test_build_add_body_string_content_with_timestamps() -> None: + body = build_add_body( + StringInput( + content="hello world", + created_at="2024-01-01T00:00:00Z", + updated_at="2024-01-02T00:00:00Z", + ), + user_id=None, + group=None, + ) + assert body == { + "input": { + "string": { + "content": ["hello world"], + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-02T00:00:00Z", + }, + }, + } + + +def test_build_add_body_string_content_with_datetime_timestamps() -> None: + body = build_add_body( + StringInput( + content="hello world", + created_at=datetime(2024, 1, 1, tzinfo=UTC), + updated_at=datetime(2024, 1, 2, tzinfo=timezone(timedelta(hours=-5))), + ), + user_id=None, + group=None, + ) + string_body = body["input"]["string"] + assert string_body["created_at"] == "2024-01-01T00:00:00Z" + assert string_body["updated_at"] == "2024-01-02T00:00:00-05:00" + + +def test_build_add_body_string_content_naive_datetime_assumed_utc() -> None: + body = build_add_body( + StringInput(content="hello", created_at=datetime(2024, 1, 1)), + user_id=None, + group=None, + ) + assert body["input"]["string"]["created_at"] == "2024-01-01T00:00:00Z" + + def test_build_add_body_conversation_content() -> None: messages = [ MessageInput(role="user", content="hi"), @@ -147,6 +194,25 @@ def test_build_add_body_conversation_content_with_message_timestamps() -> None: assert "tool_call_metadata" not in msg +def test_build_add_body_conversation_content_with_datetime_timestamps() -> None: + messages = [ + MessageInput(role="user", content="hi", created_at=datetime(2024, 1, 1, tzinfo=UTC)) + ] + body = build_add_body( + ConversationInput( + messages=messages, + created_at=datetime(2024, 1, 1, tzinfo=UTC), + updated_at=datetime(2024, 1, 2, tzinfo=timezone(timedelta(hours=-5))), + ), + user_id=None, + group=None, + ) + conv = body["input"]["conversation"] + assert conv["messages"][0]["created_at"] == "2024-01-01T00:00:00Z" + assert conv["created_at"] == "2024-01-01T00:00:00Z" + assert conv["updated_at"] == "2024-01-02T00:00:00-05:00" + + def test_build_add_body_conversation_content_with_tool_calls() -> None: messages = [ MessageInput(