Build reusable bot framework
This commit is contained in:
275
tests/unit/test_parser.py
Normal file
275
tests/unit/test_parser.py
Normal file
@@ -0,0 +1,275 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from ai import parser
|
||||
from core.registry import FrameworkRegistry
|
||||
|
||||
|
||||
def _response(content=None, reasoning=None, choices=True):
|
||||
message = SimpleNamespace(content=content, reasoning=reasoning)
|
||||
return SimpleNamespace(choices=[SimpleNamespace(message=message)] if choices else [])
|
||||
|
||||
|
||||
def test_template_rendering_preserves_literal_json_braces():
|
||||
template = (
|
||||
'Schema: {"action":"create","payload":{"items":[1, 2]}}\n'
|
||||
"Input: {user_input}\nUnknown: {leave_this_alone}"
|
||||
)
|
||||
|
||||
rendered = parser._render_template(template, {"user_input": "remind me"})
|
||||
|
||||
assert '{"action":"create","payload":{"items":[1, 2]}}' in rendered
|
||||
assert "Input: remind me" in rendered
|
||||
assert "Unknown: {leave_this_alone}" in rendered
|
||||
|
||||
|
||||
def test_json_extraction_handles_fenced_text_and_nested_values():
|
||||
response = """```json
|
||||
Here is the structured result:
|
||||
{"action":"create","payload":{"items":[{"id":1},{"id":2}]}}
|
||||
```"""
|
||||
|
||||
assert parser._extract_json_from_text(response) == {
|
||||
"action": "create",
|
||||
"payload": {"items": [{"id": 1}, {"id": 2}]},
|
||||
}
|
||||
|
||||
|
||||
def test_parse_retries_invalid_json_and_validator_failures(monkeypatch):
|
||||
responses = iter(
|
||||
[
|
||||
"not JSON",
|
||||
'{"valid":false}',
|
||||
'```json\n{"valid":true,"nested":{"value":7}}\n```',
|
||||
]
|
||||
)
|
||||
prompts = []
|
||||
|
||||
def fake_llm(_system_prompt, user_prompt):
|
||||
prompts.append(user_prompt)
|
||||
return next(responses)
|
||||
|
||||
prompt = {
|
||||
"system": "Return an object",
|
||||
"user_template": (
|
||||
'Literal schema: {"valid":true}\nUser message: {user_input}'
|
||||
),
|
||||
}
|
||||
validator = lambda value: [] if value.get("valid") else ["valid must be true"]
|
||||
monkeypatch.setitem(parser.AI_CONFIG["validation"], "max_retries", 3)
|
||||
monkeypatch.setattr(parser, "_call_llm", fake_llm)
|
||||
|
||||
result = parser.parse(
|
||||
"test input",
|
||||
"focused",
|
||||
prompt_override=prompt,
|
||||
validator=validator,
|
||||
)
|
||||
|
||||
assert result == {"valid": True, "nested": {"value": 7}}
|
||||
assert len(prompts) == 3
|
||||
assert 'Literal schema: {"valid":true}' in prompts[0]
|
||||
assert "Response was not valid JSON" in prompts[1]
|
||||
assert "valid must be true" in prompts[2]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_command_parser_routes_then_uses_focused_module_prompt(monkeypatch):
|
||||
target = FrameworkRegistry()
|
||||
target.begin_module("reminders", "modules.reminders")
|
||||
target.register_command(
|
||||
"reminder",
|
||||
lambda _context, _parsed: None,
|
||||
{
|
||||
"system": "Focused reminder parser",
|
||||
"user_template": (
|
||||
'Timezone: {timezone}\nMessage: {user_input}\n'
|
||||
'Schema: {"action":"create"}'
|
||||
),
|
||||
},
|
||||
validator=lambda value: (
|
||||
[] if value.get("action") == "create" else ["invalid action"]
|
||||
),
|
||||
description="Create and manage reminders",
|
||||
)
|
||||
target.finish_module()
|
||||
|
||||
llm = AsyncMock(
|
||||
side_effect=[
|
||||
'{"interaction_type":"reminder","confidence":0.99}',
|
||||
(
|
||||
'```json\n{"action":"create","message":"call home",'
|
||||
'"payload":{"source":"dm"}}\n```'
|
||||
),
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(parser, "_call_llm_async", llm)
|
||||
|
||||
result = await parser.parse_command_async(
|
||||
"remind me to call home",
|
||||
target,
|
||||
history=[("hello", {"interaction_type": "reminder"})],
|
||||
timezone_name="America/Chicago",
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"action": "create",
|
||||
"message": "call home",
|
||||
"payload": {"source": "dm"},
|
||||
"interaction_type": "reminder",
|
||||
}
|
||||
assert llm.await_count == 2
|
||||
route_system, route_user = llm.await_args_list[0].args
|
||||
focused_system, focused_user = llm.await_args_list[1].args
|
||||
assert route_system == parser.AI_CONFIG["prompts"]["command_parser"]["system"]
|
||||
assert "- reminder: Create and manage reminders" in route_user
|
||||
assert focused_system == "Focused reminder parser"
|
||||
assert "Timezone: America/Chicago" in focused_user
|
||||
assert 'Schema: {"action":"create"}' in focused_user
|
||||
|
||||
|
||||
def test_clients_are_created_lazily_and_cached(monkeypatch):
|
||||
syncClient = SimpleNamespace()
|
||||
asyncClient = SimpleNamespace()
|
||||
syncFactory = MagicMock(return_value=syncClient)
|
||||
asyncFactory = MagicMock(return_value=asyncClient)
|
||||
monkeypatch.setattr(parser, "OpenAI", syncFactory)
|
||||
monkeypatch.setattr(parser, "AsyncOpenAI", asyncFactory)
|
||||
monkeypatch.setattr(parser, "_sync_client", None)
|
||||
monkeypatch.setattr(parser, "_async_client", None)
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "provider-key")
|
||||
|
||||
assert parser._get_client() is syncClient
|
||||
assert parser._get_client() is syncClient
|
||||
assert parser._get_client(async_client=True) is asyncClient
|
||||
assert parser._get_client(async_client=True) is asyncClient
|
||||
assert syncFactory.call_count == 1
|
||||
assert asyncFactory.call_count == 1
|
||||
|
||||
|
||||
def test_response_text_and_request_arguments(monkeypatch):
|
||||
assert parser._response_text(_response(" result ")) == "result"
|
||||
assert parser._response_text(_response(reasoning=" reason ")) == "reason"
|
||||
assert parser._response_text(_response(choices=False)) is None
|
||||
|
||||
monkeypatch.setitem(parser.AI_CONFIG, "json_mode", True)
|
||||
arguments = parser._request_args("system", "user")
|
||||
assert arguments["messages"] == [
|
||||
{"role": "system", "content": "system"},
|
||||
{"role": "user", "content": "user"},
|
||||
]
|
||||
assert arguments["response_format"] == {"type": "json_object"}
|
||||
|
||||
|
||||
def test_sync_llm_call_returns_content_and_absorbs_provider_failure(monkeypatch):
|
||||
completions = MagicMock()
|
||||
completions.create.side_effect = [
|
||||
_response('{"ok":true}'),
|
||||
RuntimeError("provider unavailable"),
|
||||
]
|
||||
client = SimpleNamespace(chat=SimpleNamespace(completions=completions))
|
||||
monkeypatch.setattr(parser, "_get_client", lambda **_kwargs: client)
|
||||
|
||||
assert parser._call_llm("system", "user") == '{"ok":true}'
|
||||
assert parser._call_llm("system", "user") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_llm_call_returns_content_and_absorbs_failure(monkeypatch):
|
||||
create = AsyncMock(
|
||||
side_effect=[_response('{"ok":true}'), RuntimeError("offline")]
|
||||
)
|
||||
client = SimpleNamespace(
|
||||
chat=SimpleNamespace(completions=SimpleNamespace(create=create))
|
||||
)
|
||||
monkeypatch.setattr(parser, "_get_client", lambda **_kwargs: client)
|
||||
|
||||
assert await parser._call_llm_async("system", "user") == '{"ok":true}'
|
||||
assert await parser._call_llm_async("system", "user") is None
|
||||
|
||||
|
||||
def test_parse_reports_unknown_unavailable_and_exhausted_results(monkeypatch):
|
||||
assert parser.parse("hello", "missing") == {
|
||||
"error": "Unknown interaction type: missing"
|
||||
}
|
||||
|
||||
prompt = {"system": "system", "user_template": "{user_input}"}
|
||||
monkeypatch.setattr(parser, "_call_llm", lambda *_args: None)
|
||||
assert parser.parse("hello", "test", prompt_override=prompt) == {
|
||||
"error": "AI service unavailable",
|
||||
"user_input": "hello",
|
||||
}
|
||||
|
||||
monkeypatch.setitem(parser.AI_CONFIG["validation"], "max_retries", 2)
|
||||
monkeypatch.setattr(parser, "_call_llm", lambda *_args: "not json")
|
||||
exhausted = parser.parse("hello", "test", prompt_override=prompt)
|
||||
assert exhausted["error"] == "Failed to parse after 2 attempts"
|
||||
assert exhausted["validation_errors"] == ["Response was not valid JSON"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_parse_unknown_unavailable_and_exhausted(monkeypatch):
|
||||
assert await parser.parse_async("hello", "missing") == {
|
||||
"error": "Unknown interaction type: missing"
|
||||
}
|
||||
prompt = {"system": "system", "user_template": "{user_input}"}
|
||||
call = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr(parser, "_call_llm_async", call)
|
||||
assert await parser.parse_async("hello", "test", prompt_override=prompt) == {
|
||||
"error": "AI service unavailable",
|
||||
"user_input": "hello",
|
||||
}
|
||||
|
||||
monkeypatch.setitem(parser.AI_CONFIG["validation"], "max_retries", 1)
|
||||
call.return_value = "[]"
|
||||
exhausted = await parser.parse_async(
|
||||
"hello", "test", prompt_override=prompt
|
||||
)
|
||||
assert exhausted["validation_errors"] == [
|
||||
"Response must be a JSON object"
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_command_router_handles_early_and_ambiguous_results(monkeypatch):
|
||||
target = FrameworkRegistry()
|
||||
target.begin_module("one", "modules.one")
|
||||
target.register_command(
|
||||
"one",
|
||||
lambda *_args: None,
|
||||
{"system": "system", "user_template": "{user_input}"},
|
||||
)
|
||||
target.finish_module()
|
||||
|
||||
call = AsyncMock(return_value={"error": "offline"})
|
||||
monkeypatch.setattr(parser, "parse_async", call)
|
||||
assert await parser.parse_command_async("hello", target) == {
|
||||
"error": "offline"
|
||||
}
|
||||
|
||||
call.return_value = {
|
||||
"needs_clarification": "which one?",
|
||||
"confidence": 0.95,
|
||||
}
|
||||
assert await parser.parse_command_async("hello", target) == {
|
||||
"needs_clarification": "which one?",
|
||||
"confidence": 0.95,
|
||||
}
|
||||
|
||||
call.return_value = {"interaction_type": "one", "confidence": 0.2}
|
||||
ambiguous = await parser.parse_command_async("hello", target)
|
||||
assert "needs_clarification" in ambiguous
|
||||
|
||||
call.return_value = {"interaction_type": "missing", "confidence": 1.0}
|
||||
assert await parser.parse_command_async("hello", target) == {
|
||||
"error": "Unknown command type: missing"
|
||||
}
|
||||
|
||||
|
||||
def test_registered_validator_is_used(monkeypatch):
|
||||
prompt = {"system": "system", "user_template": "{user_input}"}
|
||||
parser.register_validator("registered", lambda value: [] if value["ok"] else ["bad"])
|
||||
monkeypatch.setattr(parser, "_call_llm", lambda *_args: '{"ok":true}')
|
||||
assert parser.parse("hello", "registered", prompt_override=prompt) == {"ok": True}
|
||||
Reference in New Issue
Block a user