diff --git a/tests/agents/test_orchestrator_chain_regressions.py b/tests/agents/test_orchestrator_chain_regressions.py new file mode 100644 index 00000000..d2d5958b --- /dev/null +++ b/tests/agents/test_orchestrator_chain_regressions.py @@ -0,0 +1,41 @@ +import os + +os.environ.setdefault("UTU_LLM_TYPE", "openai") +os.environ.setdefault("UTU_LLM_MODEL", "gpt-4o-mini") + +import pytest + +from utu.agents.common import TaskRecorder +from utu.agents.orchestrator.chain import ChainPlanner +from utu.agents.orchestrator.common import Recorder + + +def test_chain_planner_parse_missing_plan_block_raises_clear_assertion() -> None: + planner = ChainPlanner.__new__(ChainPlanner) + recorder = Recorder(input="question") + + with pytest.raises(AssertionError, match="No tasks parsed from plan"): + planner._parse("ok", recorder) + + +def test_chain_planner_parse_still_extracts_tasks_from_valid_plan() -> None: + planner = ChainPlanner.__new__(ChainPlanner) + recorder = Recorder(input="question") + + plan = planner._parse( + 'ok[{"name": "agent-a", "task": "do work"}]', + recorder, + ) + + assert plan.analysis == "ok" + assert len(plan.tasks) == 1 + assert plan.tasks[0].agent_name == "agent-a" + assert plan.tasks[0].task == "do work" + assert plan.tasks[0].is_last_task is True + + +def test_task_recorder_input_defaults_to_list() -> None: + recorder = TaskRecorder() + + assert recorder.input == [] + assert isinstance(recorder.input, list) diff --git a/utu/agents/common.py b/utu/agents/common.py index 9f816729..33e04e33 100644 --- a/utu/agents/common.py +++ b/utu/agents/common.py @@ -92,7 +92,7 @@ def to_dict(self): class TaskRecorder(DataClassWithStreamEvents): task: str = "" trace_id: str = "" - input: str | list[TResponseInputItem] = field(default_factory=dict) + input: str | list[TResponseInputItem] = field(default_factory=list) # from RunResultStreaming final_output: str = "" diff --git a/utu/agents/orchestrator/chain.py b/utu/agents/orchestrator/chain.py index c8236ceb..b91888b6 100644 --- a/utu/agents/orchestrator/chain.py +++ b/utu/agents/orchestrator/chain.py @@ -82,7 +82,7 @@ def _parse(self, text: str, recorder: Recorder) -> Plan: analysis = match.group(1).strip() if match else "" match = re.search(r"\s*\[(.*?)\]\s*", text, re.DOTALL) - plan_content = match.group(1).strip() + plan_content = match.group(1).strip() if match else "" tasks: list[Task] = [] task_pattern = r'\{"name":\s*"([^"]+)",\s*"task":\s*"([^"]+)"\s*\}' task_matches = re.findall(task_pattern, plan_content, re.IGNORECASE)