You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

test_serialization.py 4.2 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114
  1. from pydantic import BaseModel
  2. from dataclasses import dataclass
  3. import pytest
  4. from autogen_core.base import Serialization
  5. from autogen_core.base import JSON_DATA_CONTENT_TYPE, MessageSerializer, try_get_known_serializers_for_type
  6. class PydanticMessage(BaseModel):
  7. message: str
  8. class NestingPydanticMessage(BaseModel):
  9. message: str
  10. nested: PydanticMessage
  11. @dataclass
  12. class DataclassMessage:
  13. message: str
  14. @dataclass
  15. class NestingDataclassMessage:
  16. message: str
  17. nested: DataclassMessage
  18. @dataclass
  19. class NestingPydanticDataclassMessage:
  20. message: str
  21. nested: PydanticMessage
  22. def test_pydantic() -> None:
  23. serde = Serialization()
  24. serde.add_serializer(try_get_known_serializers_for_type(PydanticMessage))
  25. message = PydanticMessage(message="hello")
  26. name = serde.type_name(message)
  27. json = serde.serialize(message, type_name=name, data_content_type=JSON_DATA_CONTENT_TYPE)
  28. assert name == "PydanticMessage"
  29. assert json == b'{"message":"hello"}'
  30. deserialized = serde.deserialize(json, type_name=name, data_content_type=JSON_DATA_CONTENT_TYPE)
  31. assert deserialized == message
  32. def test_nested_pydantic() -> None:
  33. serde = Serialization()
  34. serde.add_serializer(try_get_known_serializers_for_type(NestingPydanticMessage))
  35. message = NestingPydanticMessage(message="hello", nested=PydanticMessage(message="world"))
  36. name = serde.type_name(message)
  37. json = serde.serialize(message, type_name=name, data_content_type=JSON_DATA_CONTENT_TYPE)
  38. assert json == b'{"message":"hello","nested":{"message":"world"}}'
  39. deserialized = serde.deserialize(json, type_name=name, data_content_type=JSON_DATA_CONTENT_TYPE)
  40. assert deserialized == message
  41. def test_dataclass() -> None:
  42. serde = Serialization()
  43. serde.add_serializer(try_get_known_serializers_for_type(DataclassMessage))
  44. message = DataclassMessage(message="hello")
  45. name = serde.type_name(message)
  46. json = serde.serialize(message, type_name=name, data_content_type=JSON_DATA_CONTENT_TYPE)
  47. assert json == b'{"message": "hello"}'
  48. deserialized = serde.deserialize(json, type_name=name, data_content_type=JSON_DATA_CONTENT_TYPE)
  49. assert deserialized == message
  50. def test_nesting_dataclass_dataclass() -> None:
  51. serde = Serialization()
  52. serde.add_serializer(try_get_known_serializers_for_type(NestingDataclassMessage))
  53. message = NestingDataclassMessage(message="hello", nested=DataclassMessage(message="world"))
  54. name = serde.type_name(message)
  55. with pytest.raises(ValueError):
  56. _json = serde.serialize(message, type_name=name, data_content_type=JSON_DATA_CONTENT_TYPE)
  57. def test_nesting_dataclass_pydantic() -> None:
  58. serde = Serialization()
  59. serde.add_serializer(try_get_known_serializers_for_type(NestingPydanticDataclassMessage))
  60. message = NestingPydanticDataclassMessage(message="hello", nested=PydanticMessage(message="world"))
  61. name = serde.type_name(message)
  62. with pytest.raises(ValueError):
  63. _json = serde.serialize(message, type_name=name, data_content_type=JSON_DATA_CONTENT_TYPE)
  64. def test_invalid_type() -> None:
  65. serde = Serialization()
  66. try:
  67. serde.add_serializer(try_get_known_serializers_for_type(str))
  68. except ValueError as e:
  69. assert str(e) == "Unsupported type <class 'str'>"
  70. def test_custom_type() -> None:
  71. serde = Serialization()
  72. class CustomStringTypeSerializer(MessageSerializer[str]):
  73. @property
  74. def data_content_type(self) -> str:
  75. return "str"
  76. @property
  77. def type_name(self) -> str:
  78. return "custom_str"
  79. def deserialize(self, payload: bytes) -> str:
  80. message = payload.decode("utf-8")
  81. return message[1:-1]
  82. def serialize(self, message: str) -> bytes:
  83. return f'"{message}"'.encode("utf-8")
  84. serde.add_serializer(CustomStringTypeSerializer())
  85. message = "hello"
  86. json = serde.serialize(message, type_name="custom_str", data_content_type="str")
  87. assert json == b'"hello"'
  88. deserialized = serde.deserialize(json, type_name="custom_str", data_content_type="str")
  89. assert deserialized == message