|
5 | 5 | import json |
6 | 6 | import warnings |
7 | 7 | from collections.abc import Sequence |
| 8 | +from typing import Any |
8 | 9 |
|
9 | 10 | import pytest |
| 11 | +from pydantic import BaseModel |
10 | 12 |
|
11 | 13 | from haystack.dataclasses.chat_message import ( |
12 | 14 | ChatMessage, |
@@ -756,6 +758,85 @@ def test_to_trace_dict_with_file_content(self, base64_pdf_string): |
756 | 758 | } |
757 | 759 |
|
758 | 760 |
|
| 761 | +class MessageEnvelope(BaseModel): |
| 762 | + message: ChatMessage |
| 763 | + |
| 764 | + |
| 765 | +class TestFromDictPydanticDump: |
| 766 | + """ |
| 767 | + `ChatMessage.from_dict` supports the format Pydantic produces when it auto-serializes ChatMessage as a plain |
| 768 | + dataclass: raw dataclass fields (`_role`, `_content`, ...) with unwrapped content parts. |
| 769 | + """ |
| 770 | + |
| 771 | + def _pydantic_dump(self, message: ChatMessage) -> dict[str, Any]: |
| 772 | + return MessageEnvelope(message=message).model_dump(mode="json")["message"] |
| 773 | + |
| 774 | + def test_text_message(self): |
| 775 | + message = ChatMessage.from_user("What is the answer?", meta={"some": "info"}, name="virginia") |
| 776 | + assert ChatMessage.from_dict(self._pydantic_dump(message)) == message |
| 777 | + |
| 778 | + def test_tool_call_message(self): |
| 779 | + message = ChatMessage.from_assistant( |
| 780 | + tool_calls=[ToolCall(tool_name="mytool", arguments={"a": 1}, id="123", extra={"call_id": "123"})] |
| 781 | + ) |
| 782 | + assert ChatMessage.from_dict(self._pydantic_dump(message)) == message |
| 783 | + |
| 784 | + def test_tool_result_message(self): |
| 785 | + message = ChatMessage.from_tool( |
| 786 | + tool_result="42", origin=ToolCall(tool_name="mytool", arguments={"a": 1}, id="123"), error=False |
| 787 | + ) |
| 788 | + assert ChatMessage.from_dict(self._pydantic_dump(message)) == message |
| 789 | + |
| 790 | + def test_reasoning_message(self): |
| 791 | + message = ChatMessage.from_assistant( |
| 792 | + "Answer", reasoning=ReasoningContent(reasoning_text="Thinking...", extra={"key": "value"}) |
| 793 | + ) |
| 794 | + assert ChatMessage.from_dict(self._pydantic_dump(message)) == message |
| 795 | + |
| 796 | + def test_image_message(self, base64_image_string): |
| 797 | + message = ChatMessage.from_user( |
| 798 | + content_parts=[ |
| 799 | + TextContent(text="What is in this image?"), |
| 800 | + ImageContent(base64_image=base64_image_string, mime_type="image/png", detail="auto"), |
| 801 | + ] |
| 802 | + ) |
| 803 | + assert ChatMessage.from_dict(self._pydantic_dump(message)) == message |
| 804 | + |
| 805 | + def test_file_message(self): |
| 806 | + message = ChatMessage.from_user( |
| 807 | + content_parts=[ |
| 808 | + TextContent(text="Summarize this file."), |
| 809 | + FileContent(base64_data="aGVsbG8=", mime_type="text/plain", filename="hello.txt"), |
| 810 | + ] |
| 811 | + ) |
| 812 | + assert ChatMessage.from_dict(self._pydantic_dump(message)) == message |
| 813 | + |
| 814 | + def test_multiple_messages(self, base64_image_string): |
| 815 | + class Response(BaseModel): |
| 816 | + messages: list[ChatMessage] |
| 817 | + |
| 818 | + tool_call = ToolCall(id="123", tool_name="mytool", arguments={"a": 1}) |
| 819 | + messages = [ |
| 820 | + ChatMessage.from_user("What is the answer?"), |
| 821 | + ChatMessage.from_assistant( |
| 822 | + "Let me check.", |
| 823 | + meta={"some": "info"}, |
| 824 | + tool_calls=[tool_call], |
| 825 | + reasoning=ReasoningContent(reasoning_text="Let me think about it..."), |
| 826 | + ), |
| 827 | + ChatMessage.from_tool(tool_result="42", origin=tool_call), |
| 828 | + ChatMessage.from_user( |
| 829 | + content_parts=[ |
| 830 | + ImageContent(base64_image=base64_image_string, mime_type="image/png"), |
| 831 | + FileContent(base64_data="aGVsbG8=", mime_type="text/plain", filename="hello.txt"), |
| 832 | + ] |
| 833 | + ), |
| 834 | + ] |
| 835 | + |
| 836 | + dumped = Response(messages=messages).model_dump(mode="json") |
| 837 | + assert [ChatMessage.from_dict(message) for message in dumped["messages"]] == messages |
| 838 | + |
| 839 | + |
759 | 840 | class TestToOpenaiDictFormat: |
760 | 841 | def test_to_openai_dict_format_system_message(self): |
761 | 842 | message = ChatMessage.from_system("You are good assistant") |
|
0 commit comments