From de90ade2942a8a5da711245429c1e2de587a923e Mon Sep 17 00:00:00 2001 From: Scott Severance Date: Sat, 29 Aug 2026 19:04:24 +0000 Subject: [PATCH 1/2] fix: validate message structure before accessing content --- agent/core/llm.py | 34 ++++++++++++++++++++++++++++------ 1 file changed, 28 insertions(+), 6 deletions(-) diff --git a/agent/core/llm.py b/agent/core/llm.py index 9ffe11d..901a15f 100644 --- a/agent/core/llm.py +++ b/agent/core/llm.py @@ -39,6 +39,25 @@ def _validate_response(self, response: dict, required_fields: list[str]) -> None f"Got: {response}" ) + def _validate_message_structure(self, response: dict) -> None: + """Validate that response message has expected structure. + + Args: + response: The response dict to validate. + + Raises: + OllamaError: If message structure is invalid. + """ + if "message" not in response: + raise OllamaError(f"Missing 'message' field in response: {response}") + message = response["message"] + if not isinstance(message, dict): + raise OllamaError( + f"Expected 'message' to be a dict, got {type(message).__name__}: {response}" + ) + if "content" not in message: + raise OllamaError(f"Missing 'content' field in message: {response}") + def _request( self, endpoint: str, @@ -162,6 +181,7 @@ def chat( result = self._request("/api/chat", data, timeout=timeout) self._validate_response(result, ["message"]) + self._validate_message_structure(result) return result def _stream_chat( @@ -181,12 +201,14 @@ def _stream_chat( try: chunk = json.loads(line) if "message" in chunk: - content = chunk["message"].get("content", "") - full_response["message"]["content"] += content - yield content - # Check for tool calls - if "tool_calls" in chunk["message"]: - full_response["message"]["tool_calls"] = chunk["message"]["tool_calls"] + message = chunk["message"] + if isinstance(message, dict) and "content" in message: + content = message.get("content", "") + full_response["message"]["content"] += content + yield content + # Check for tool calls + if "tool_calls" in message: + full_response["message"]["tool_calls"] = message["tool_calls"] if chunk.get("done", False): break except json.JSONDecodeError: From 69a61f118628c9f835c9b2df5f6ed2d3ea4a5729 Mon Sep 17 00:00:00 2001 From: "claude[bot]" <41898282+claude[bot]@users.noreply.github.com> Date: Sat, 29 Aug 2026 19:08:15 +0000 Subject: [PATCH 2/2] fix(review): keep streamed tool calls, validate content type not presence - _stream_chat no longer gates the tool_calls branch on a 'content' key, so a chunk carrying only tool_calls is recorded again. - Validate that 'content' is a str; tolerate missing/null content when 'tool_calls' is present, matching engine.py's `.get("content", "") or ""`. - Drop the redundant _validate_response(result, ["message"]) call. - Add tests/test_llm.py covering the validator and streaming chunks. Co-Authored-By: Claude Opus 5 (1M context) --- agent/core/llm.py | 34 +++++++--- tests/test_llm.py | 170 ++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 194 insertions(+), 10 deletions(-) create mode 100644 tests/test_llm.py diff --git a/agent/core/llm.py b/agent/core/llm.py index 901a15f..4168ec1 100644 --- a/agent/core/llm.py +++ b/agent/core/llm.py @@ -41,10 +41,13 @@ def _validate_response(self, response: dict, required_fields: list[str]) -> None def _validate_message_structure(self, response: dict) -> None: """Validate that response message has expected structure. - + + A missing or null 'content' is tolerated when the message carries + 'tool_calls', since a tool-call-only turn is a valid Ollama response. + Args: response: The response dict to validate. - + Raises: OllamaError: If message structure is invalid. """ @@ -55,8 +58,14 @@ def _validate_message_structure(self, response: dict) -> None: raise OllamaError( f"Expected 'message' to be a dict, got {type(message).__name__}: {response}" ) - if "content" not in message: - raise OllamaError(f"Missing 'content' field in message: {response}") + content = message.get("content") + if content is None: + if not message.get("tool_calls"): + raise OllamaError(f"Missing 'content' field in message: {response}") + elif not isinstance(content, str): + raise OllamaError( + f"Expected 'content' to be a str, got {type(content).__name__}: {response}" + ) def _request( self, @@ -180,7 +189,6 @@ def chat( return self._stream_chat(data, timeout=timeout) result = self._request("/api/chat", data, timeout=timeout) - self._validate_response(result, ["message"]) self._validate_message_structure(result) return result @@ -202,11 +210,17 @@ def _stream_chat( chunk = json.loads(line) if "message" in chunk: message = chunk["message"] - if isinstance(message, dict) and "content" in message: - content = message.get("content", "") - full_response["message"]["content"] += content - yield content - # Check for tool calls + if isinstance(message, dict): + content = message.get("content") + if content is not None and not isinstance(content, str): + raise OllamaError( + f"Expected 'content' to be a str, got " + f"{type(content).__name__}: {chunk}" + ) + if content: + full_response["message"]["content"] += content + yield content + # Tool calls can arrive in a chunk without content. if "tool_calls" in message: full_response["message"]["tool_calls"] = message["tool_calls"] if chunk.get("done", False): diff --git a/tests/test_llm.py b/tests/test_llm.py new file mode 100644 index 0000000..e9955ef --- /dev/null +++ b/tests/test_llm.py @@ -0,0 +1,170 @@ +"""Tests for the Ollama LLM client response validation.""" + +import json +import sys +import unittest +from pathlib import Path +from unittest.mock import patch + +sys.path.insert(0, str(Path(__file__).parent.parent)) + +from agent.core.config import LLMConfig +from agent.core.llm import LLMClient, OllamaError + + +def make_client() -> LLMClient: + return LLMClient(LLMConfig()) + + +def stream_lines(chunks) -> list[bytes]: + return [json.dumps(chunk).encode("utf-8") for chunk in chunks] + + +def drain(generator): + """Run a generator to completion, returning (yielded, return_value).""" + yielded = [] + while True: + try: + yielded.append(next(generator)) + except StopIteration as stop: + return yielded, stop.value + + +class TestValidateMessageStructure(unittest.TestCase): + """Test _validate_message_structure.""" + + def test_missing_message_raises(self): + with self.assertRaises(OllamaError): + make_client()._validate_message_structure({"done": True}) + + def test_non_dict_message_raises(self): + with self.assertRaises(OllamaError): + make_client()._validate_message_structure({"message": "hello"}) + + def test_non_string_content_raises(self): + with self.assertRaises(OllamaError): + make_client()._validate_message_structure( + {"message": {"role": "assistant", "content": ["a", "b"]}} + ) + + def test_missing_content_without_tool_calls_raises(self): + with self.assertRaises(OllamaError): + make_client()._validate_message_structure( + {"message": {"role": "assistant"}} + ) + + def test_missing_content_with_tool_calls_is_allowed(self): + make_client()._validate_message_structure( + { + "message": { + "role": "assistant", + "tool_calls": [{"function": {"name": "read_file"}}], + } + } + ) + + def test_null_content_with_tool_calls_is_allowed(self): + make_client()._validate_message_structure( + { + "message": { + "role": "assistant", + "content": None, + "tool_calls": [{"function": {"name": "read_file"}}], + } + } + ) + + def test_valid_message_passes(self): + make_client()._validate_message_structure( + {"message": {"role": "assistant", "content": "hi"}} + ) + + def test_chat_accepts_tool_call_only_response(self): + client = make_client() + response = { + "message": { + "role": "assistant", + "tool_calls": [{"function": {"name": "read_file", "arguments": {}}}], + } + } + with patch.object(client, "_request", return_value=response): + self.assertEqual(client.chat([{"role": "user", "content": "hi"}]), response) + + def test_chat_rejects_non_dict_message(self): + client = make_client() + with patch.object(client, "_request", return_value={"message": None}): + with self.assertRaises(OllamaError): + client.chat([{"role": "user", "content": "hi"}]) + + +class TestStreamChat(unittest.TestCase): + """Test _stream_chat chunk handling.""" + + def test_tool_calls_without_content_are_kept(self): + client = make_client() + tool_calls = [{"function": {"name": "read_file", "arguments": {}}}] + lines = stream_lines( + [ + {"message": {"role": "assistant", "tool_calls": tool_calls}}, + {"message": {"role": "assistant", "content": ""}, "done": True}, + ] + ) + with patch.object(client, "_request", return_value=lines): + yielded, result = drain(client._stream_chat({})) + + self.assertEqual(result["message"].get("tool_calls"), tool_calls) + self.assertEqual(yielded, []) + + def test_content_chunks_are_accumulated_and_yielded(self): + client = make_client() + lines = stream_lines( + [ + {"message": {"role": "assistant", "content": "he"}}, + {"message": {"role": "assistant", "content": "llo"}}, + {"message": {"role": "assistant", "content": ""}, "done": True}, + ] + ) + with patch.object(client, "_request", return_value=lines): + yielded, result = drain(client._stream_chat({})) + + self.assertEqual(yielded, ["he", "llo"]) + self.assertEqual(result["message"]["content"], "hello") + + def test_null_content_does_not_crash(self): + client = make_client() + lines = stream_lines( + [ + {"message": {"role": "assistant", "content": None}}, + {"message": {"role": "assistant", "content": "ok"}, "done": True}, + ] + ) + with patch.object(client, "_request", return_value=lines): + yielded, result = drain(client._stream_chat({})) + + self.assertEqual(yielded, ["ok"]) + self.assertEqual(result["message"]["content"], "ok") + + def test_non_string_content_raises(self): + client = make_client() + lines = stream_lines([{"message": {"role": "assistant", "content": [1, 2]}}]) + with patch.object(client, "_request", return_value=lines): + with self.assertRaises(OllamaError): + drain(client._stream_chat({})) + + def test_non_dict_message_is_skipped(self): + client = make_client() + lines = stream_lines( + [ + {"message": "not a dict"}, + {"message": {"role": "assistant", "content": "ok"}, "done": True}, + ] + ) + with patch.object(client, "_request", return_value=lines): + yielded, result = drain(client._stream_chat({})) + + self.assertEqual(yielded, ["ok"]) + self.assertEqual(result["message"]["content"], "ok") + + +if __name__ == "__main__": + unittest.main()