diff --git a/python/packages/autogen-core/src/autogen_core/memory/_base_memory.py b/python/packages/autogen-core/src/autogen_core/memory/_base_memory.py index 5385b37d3e35..f59fa0167217 100644 --- a/python/packages/autogen-core/src/autogen_core/memory/_base_memory.py +++ b/python/packages/autogen-core/src/autogen_core/memory/_base_memory.py @@ -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 @@ -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.""" diff --git a/python/packages/autogen-core/tests/test_memory.py b/python/packages/autogen-core/tests/test_memory.py index ce98aaffb97f..52b90072a702 100644 --- a/python/packages/autogen-core/tests/test_memory.py +++ b/python/packages/autogen-core/tests/test_memory.py @@ -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."""