Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
from enum import Enum
from typing import Any, Dict, List, Union

from pydantic import BaseModel, ConfigDict, field_serializer
from pydantic import BaseModel, ConfigDict, field_serializer, field_validator

from .._cancellation_token import CancellationToken
from .._component_config import ComponentBase
Expand Down Expand Up @@ -44,6 +44,26 @@ def serialize_mime_type(self, mime_type: MemoryMimeType | str) -> str:
return mime_type.value
return mime_type

@field_validator("mime_type", mode="before")
@classmethod
def validate_mime_type(cls, value: Any) -> Any:
"""Restore a :class:`MemoryMimeType` enum from its string value on load.

This is the inverse of :meth:`serialize_mime_type`, so that a
dump/load roundtrip preserves the enum type instead of degenerating
it to a plain string.

.. versionchanged:: 0.7.6
Added validation to restore the enum type after serialization.
"""
if isinstance(value, str):
try:
return MemoryMimeType(value)
except ValueError:
# Not a known MemoryMimeType value: keep it as a plain string.
return value
return value


class MemoryQueryResult(BaseModel):
"""Result of a memory :meth:`~autogen_core.memory.Memory.query` operation."""
Expand Down
35 changes: 35 additions & 0 deletions python/packages/autogen-core/tests/test_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,41 @@ def test_memory_component_dump_config_to_base_model() -> None:
assert len(config.config["memory_contents"]) == 1


def test_memory_content_mime_type_roundtrip_preserves_enum() -> None:
"""Test that a known MIME type stays a MemoryMimeType enum after model validation (#8293)."""
content = MemoryContent(content="test", mime_type=MemoryMimeType.TEXT)
assert content.mime_type is MemoryMimeType.TEXT

restored = MemoryContent.model_validate(content.model_dump())
assert restored.mime_type is MemoryMimeType.TEXT

restored_json = MemoryContent.model_validate(
MemoryContent(content={"key": "value"}, mime_type=MemoryMimeType.JSON).model_dump()
)
assert restored_json.mime_type is MemoryMimeType.JSON


def test_memory_content_mime_type_custom_string_preserved() -> None:
"""Test that custom MIME type strings are not converted to the enum."""
content = MemoryContent(content="test", mime_type="application/custom+json")
assert content.mime_type == "application/custom+json"

restored = MemoryContent.model_validate(content.model_dump())
assert restored.mime_type == "application/custom+json"
assert not isinstance(restored.mime_type, MemoryMimeType)


def test_memory_component_roundtrip_preserves_mime_type_enum() -> None:
"""Test that dump_component/load_component keeps mime_type as a MemoryMimeType enum (#8293)."""
memory = ListMemory(
name="test_memory", memory_contents=[MemoryContent(content="test", mime_type=MemoryMimeType.JSON)]
)
restored = ListMemory.load_component(memory.dump_component())
assert isinstance(restored, ListMemory)
assert len(restored.content) == 1
assert restored.content[0].mime_type is MemoryMimeType.JSON


def test_memory_abc_implementation() -> None:
"""Test that Memory ABC is properly implemented."""

Expand Down