Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
50 changes: 43 additions & 7 deletions agent/core/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,34 @@ 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.

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.
"""
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}"
)
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,
endpoint: str,
Expand Down Expand Up @@ -161,7 +189,7 @@ 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

def _stream_chat(
Expand All @@ -181,12 +209,20 @@ 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):
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):
break
except json.JSONDecodeError:
Expand Down
170 changes: 170 additions & 0 deletions tests/test_llm.py
Original file line number Diff line number Diff line change
@@ -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()