diff --git a/src/claude_agent_sdk/_internal/message_parser.py b/src/claude_agent_sdk/_internal/message_parser.py index 931cc2a63..c85627e16 100644 --- a/src/claude_agent_sdk/_internal/message_parser.py +++ b/src/claude_agent_sdk/_internal/message_parser.py @@ -99,9 +99,16 @@ def parse_message(data: dict[str, Any]) -> Message | None: tool_use_result = data.get("tool_use_result") uuid = data.get("uuid") origin = _parse_origin(data) - if isinstance(data["message"]["content"], list): + message = data["message"] + if not isinstance(message, dict): + raise MessageParseError( + f"Invalid user message (expected dict, got " + f"{type(message).__name__})", + data, + ) + if isinstance(message["content"], list): user_content_blocks: list[ContentBlock] = [] - for block in data["message"]["content"]: + for block in message["content"]: if not isinstance(block, dict): raise MessageParseError( f"Invalid content block (expected dict, got " @@ -137,7 +144,7 @@ def parse_message(data: dict[str, Any]) -> Message | None: origin=origin, ) return UserMessage( - content=data["message"]["content"], + content=message["content"], uuid=uuid, parent_tool_use_id=parent_tool_use_id, tool_use_result=tool_use_result, @@ -150,7 +157,14 @@ def parse_message(data: dict[str, Any]) -> Message | None: case "assistant": try: - raw_content = data["message"]["content"] + message = data["message"] + if not isinstance(message, dict): + raise MessageParseError( + f"Invalid assistant message (expected dict, got " + f"{type(message).__name__})", + data, + ) + raw_content = message["content"] if not isinstance(raw_content, list): raise MessageParseError( f"Invalid assistant content (expected list, got " @@ -356,6 +370,12 @@ def parse_message(data: dict[str, Any]) -> Message | None: case "rate_limit_event": try: info = data["rate_limit_info"] + if not isinstance(info, dict): + raise MessageParseError( + f"Invalid rate_limit_info (expected dict, got " + f"{type(info).__name__})", + data, + ) return RateLimitEvent( rate_limit_info=RateLimitInfo( status=info["status"], diff --git a/tests/test_message_parser.py b/tests/test_message_parser.py index e55fd1556..d1ae053d8 100644 --- a/tests/test_message_parser.py +++ b/tests/test_message_parser.py @@ -1015,6 +1015,37 @@ def test_non_dict_content_block_raises_documented_error(self, role: str) -> None with pytest.raises(MessageParseError): parse_message({"type": role, "message": message}) + @pytest.mark.parametrize("role", ["assistant", "user"]) + @pytest.mark.parametrize("bad_message", ["a string", 5, ["a", "list"], None]) + def test_non_dict_message_raises_documented_error( + self, role: str, bad_message: object + ) -> None: + """A non-dict ``message`` raises MessageParseError, never a raw TypeError. + + The block-level guard (above) only runs after ``message["content"]`` is + indexed, so a ``message`` that is not a dict crashed with a raw + TypeError before this guard existed. + """ + with pytest.raises(MessageParseError) as exc_info: + parse_message({"type": role, "message": bad_message}) + assert "expected dict" in str(exc_info.value) + + @pytest.mark.parametrize("bad_info", ["a string", 5, ["a", "list"], None]) + def test_non_dict_rate_limit_info_raises_documented_error( + self, bad_info: object + ) -> None: + """A non-dict ``rate_limit_info`` raises MessageParseError, not a TypeError.""" + with pytest.raises(MessageParseError) as exc_info: + parse_message( + { + "type": "rate_limit_event", + "uuid": "u", + "session_id": "s", + "rate_limit_info": bad_info, + } + ) + assert "expected dict" in str(exc_info.value) + def test_parse_system_message_missing_fields(self): """Test that system message with missing fields raises MessageParseError.""" with pytest.raises(MessageParseError) as exc_info: