Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 7 additions & 4 deletions src/engram/_models/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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)
Expand Down Expand Up @@ -68,20 +71,20 @@ 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


@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()
Expand Down
31 changes: 24 additions & 7 deletions src/engram/_serialization/_builders.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

from datetime import UTC, datetime
from typing import Any, TypeAlias

from .._models import (
Expand Down Expand Up @@ -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}}
Expand All @@ -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:
Expand All @@ -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}


Expand Down
66 changes: 66 additions & 0 deletions tests/test_serialization.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from datetime import UTC, datetime, timedelta, timezone

from engram._models import (
ConversationInput,
MessageInput,
Expand Down Expand Up @@ -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"),
Expand Down Expand Up @@ -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(
Expand Down
Loading