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
18 changes: 10 additions & 8 deletions src/engram/__init__.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,16 @@
from ._models import (
CommittedOperation,
CommittedOperations,
ConversationContent,
ConversationInput,
Memory,
MessageContent,
PreExtractedContent,
MessageInput,
PreExtractedInput,
PreExtractedItem,
RetrievalConfig,
Run,
RunStatus,
SearchResults,
StringContent,
StringInput,
ToolCallCustomInput,
ToolCallFuncInput,
ToolCallInput,
Expand All @@ -33,18 +34,19 @@
"CommittedOperation",
"CommittedOperations",
"ConnectionError",
"ConversationContent",
"ConversationInput",
"EngramClient",
"EngramError",
"EngramTimeoutError",
"Memory",
"MessageContent",
"PreExtractedContent",
"MessageInput",
"PreExtractedInput",
"PreExtractedItem",
"RetrievalConfig",
"Run",
"RunStatus",
"SearchResults",
"StringContent",
"StringInput",
"ToolCallCustomInput",
"ToolCallFuncInput",
"ToolCallInput",
Expand Down
22 changes: 12 additions & 10 deletions src/engram/_models/__init__.py
Original file line number Diff line number Diff line change
@@ -1,31 +1,33 @@
from .memory import (
AddContent,
ConversationContent,
AddInput,
ConversationInput,
Memory,
MessageContent,
PreExtractedContent,
MessageInput,
PreExtractedInput,
PreExtractedItem,
RetrievalConfig,
SearchResults,
StringContent,
StringInput,
ToolCallCustomInput,
ToolCallFuncInput,
ToolCallInput,
)
from .run import CommittedOperation, CommittedOperations, Run, RunStatus

__all__ = [
"AddContent",
"AddInput",
"CommittedOperation",
"CommittedOperations",
"ConversationContent",
"ConversationInput",
"Memory",
"MessageContent",
"PreExtractedContent",
"MessageInput",
"PreExtractedInput",
"PreExtractedItem",
"RetrievalConfig",
"Run",
"RunStatus",
"SearchResults",
"StringContent",
"StringInput",
"ToolCallCustomInput",
"ToolCallFuncInput",
"ToolCallInput",
Expand Down
33 changes: 21 additions & 12 deletions src/engram/_models/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,18 +6,27 @@


@dataclass(slots=True)
class PreExtractedContent:
"""Pre-extracted content that bypasses the extraction pipeline."""
class PreExtractedInput:
"""Pre-extracted input that skips the extraction step continues through the pipeline as-is.
Each individual item represents a separate memory.
"""

items: list[PreExtractedItem]


@dataclass(slots=True)
class PreExtractedItem:
"""A single pre-extracted memory."""

content: str
topic: str


@dataclass(slots=True)
class StringContent:
"""String content that bypasses the extraction pipeline."""
class StringInput:
"""String input to extract memories from."""

content: str
content: str | list[str]


@dataclass(slots=True)
Expand Down Expand Up @@ -50,7 +59,7 @@ class ToolCallInput:


@dataclass(slots=True)
class MessageContent:
class MessageInput:
"""A message in a conversation using the OpenAI Chat Completions format.

- 'tool' role (tool results) is mapped to 'user' by the server.
Expand All @@ -66,18 +75,18 @@ class MessageContent:


@dataclass(slots=True)
class ConversationContent:
"""Conversation content that bypasses the extraction pipeline."""
class ConversationInput:
"""Conversation input that bypasses the extraction pipeline."""

messages: list[MessageContent]
messages: list[MessageInput]
metadata: dict[str, Any] | None = None
created_at: str | None = None
updated_at: str | None = None


# Type alias for the content argument to memories.add()
AddContent: TypeAlias = (
str | list[dict[str, str]] | PreExtractedContent | ConversationContent | StringContent
# Type alias for the input_data argument to memories.add()
AddInput: TypeAlias = (
str | list[dict[str, str]] | PreExtractedInput | ConversationInput | StringInput
)


Expand Down
10 changes: 5 additions & 5 deletions src/engram/_resources/memories.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from uuid import UUID

from .._http import AsyncHttpTransport, HttpTransport
from .._models import AddContent, Memory, RetrievalConfig, Run, SearchResults
from .._models import AddInput, Memory, RetrievalConfig, Run, SearchResults
from .._serialization import (
build_add_body,
build_memory_params,
Expand All @@ -29,14 +29,14 @@ def __init__(self, transport: HttpTransport) -> None:

def add(
self,
content: AddContent,
input_data: AddInput,
*,
user_id: str | None = None,
conversation_id: str | None = None,
group: str | None = None,
) -> Run:
body = build_add_body(
content,
input_data,
user_id=user_id,
conversation_id=conversation_id,
group=group,
Expand Down Expand Up @@ -101,14 +101,14 @@ def __init__(self, transport: AsyncHttpTransport) -> None:

async def add(
self,
content: AddContent,
input_data: AddInput,
*,
user_id: str | None = None,
conversation_id: str | None = None,
group: str | None = None,
) -> Run:
body = build_add_body(
content,
input_data,
user_id=user_id,
conversation_id=conversation_id,
group=group,
Expand Down
51 changes: 25 additions & 26 deletions src/engram/_serialization/_builders.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,11 @@
from typing import Any

from .._models import (
AddContent,
ConversationContent,
PreExtractedContent,
AddInput,
ConversationInput,
PreExtractedInput,
RetrievalConfig,
StringContent,
StringInput,
ToolCallInput,
)

Expand All @@ -21,29 +21,28 @@ def _serialize_tool_call(tc: ToolCallInput) -> dict[str, Any]:
return out


def _serialize_content(content: AddContent) -> dict[str, Any]:
"""Build the content envelope with the type discriminator."""
if isinstance(content, str):
return {"type": "string", "content": content}
if isinstance(content, StringContent):
return {"type": "string", "content": content.content}
if isinstance(content, PreExtractedContent):
def _serialize_input(input_data: AddInput) -> dict[str, Any]:
"""Build the input envelope with the type discriminator."""
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]}}
if isinstance(input_data, PreExtractedInput):
items = [{"content": item.content, "topic": item.topic} for item in input_data.items]
return {"pre_extracted": {"items": items}}
if isinstance(input_data, list):
return {
"type": "pre_extracted",
"content": content.content,
"topic": content.topic,
"conversation": {"messages": input_data},
}
if isinstance(content, list):
return {
"type": "conversation",
"conversation": {"messages": content},
}
if isinstance(content, ConversationContent):
return _serialize_conversation_content(content)
raise TypeError(f"Unsupported content type: {type(content)}") # pragma: no cover
if isinstance(input_data, ConversationInput):
return _serialize_conversation_content(input_data)
raise TypeError(f"Unsupported input type: {type(input_data)}") # pragma: no cover


def _serialize_conversation_content(content: ConversationContent) -> dict[str, Any]:
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}
Expand All @@ -63,17 +62,17 @@ def _serialize_conversation_content(content: ConversationContent) -> dict[str, A
conversation["created_at"] = content.created_at
if content.updated_at is not None:
conversation["updated_at"] = content.updated_at
return {"type": "conversation", "conversation": conversation}
return {"conversation": conversation}


def build_add_body(
content: AddContent,
input_data: AddInput,
*,
user_id: str | None,
conversation_id: str | None,
group: str | None,
) -> dict[str, Any]:
body: dict[str, Any] = {"content": _serialize_content(content)}
body: dict[str, Any] = {"input": _serialize_input(input_data)}
if user_id is not None:
body["user_id"] = user_id
if conversation_id is not None:
Expand Down
Loading
Loading