from __future__ import annotations from collections.abc import Callable from inspect import Parameter, signature from typing import TYPE_CHECKING, Any from core.errors import TaskFailedError from core.models import Scene from runtime.context import TaskContext from runtime.planner import PlannedStep, Planner from runtime.planner_config import PlannerConfig, load_config from runtime.planner_prompts import PLANNER_SYSTEM_PROMPT, planner_user_prompt from runtime.tool_calling_client import ToolCallingClient, build_client from runtime.tool_specs import ALL_TOOL_SPECS if TYPE_CHECKING: from world.models import WorldState FINISH_TASK_TOOL = "finish_task" class AIPlanner(Planner): def __init__( self, *, client: ToolCallingClient | None = None, config: PlannerConfig | None = None, event_logger: Callable[[dict[str, Any]], None] | None = None, ) -> None: self.config = config or load_config() self.client = client or build_client(self.config) self.event_logger = event_logger def plan( self, *, goal: str, scene: Scene, context: TaskContext, world: "WorldState | None" = None, screenshot: bytes | None = None, ) -> list[PlannedStep]: _sync_tool_results(context) scene_json = _without_ocr(scene.to_dict()) if self.config.multimodal else scene.to_dict() user_prompt = planner_user_prompt( goal=goal, scene_json=scene_json, device_platform=context.device_platform, ) if context.step_results and not context.step_results[-1].success: user_prompt += ( "\n\nThe previous tool call failed. Diagnose the failure and choose a " "corrected action or finish the task if it cannot proceed.\n" f"Previous failure: {context.step_results[-1].error or 'unknown error'}" ) self._log({"type": "llm_request", "task_id": context.task_id, "goal": goal, "system_prompt": PLANNER_SYSTEM_PROMPT, "user_prompt": user_prompt, "has_screenshot": screenshot is not None}) try: kwargs = { "system_prompt": PLANNER_SYSTEM_PROMPT, "user_prompt": user_prompt, "screenshot": screenshot, "tools": ALL_TOOL_SPECS, "timeout": self.config.timeout, } if _accepts_history(self.client.decide): kwargs["history"] = context.planner_history[ -self.config.history_max_turns : ] decision = self.client.decide(**kwargs) except Exception as exc: self._log({"type": "agent_error", "task_id": context.task_id, "error": str(exc)}) raise self._log({"type": "llm_response", "task_id": context.task_id, "content": decision.text_output, "thinking": decision.thinking, "tool_name": decision.tool_name, "arguments": decision.arguments, "purpose": decision.purpose, "expected_outcome": decision.expected_outcome}) if decision.tool_name == FINISH_TASK_TOOL: if decision.arguments.get("success"): return [] raise TaskFailedError(decision.arguments.get("reason") or "task failed") context.planner_history.append( { "user_prompt": user_prompt, "tool_name": decision.tool_name, "arguments": _conversation_arguments(decision), "rationale": decision.text_output, "tool_result": None, } ) return [ PlannedStep( action=decision.tool_name, description=( f"AI planner: {decision.purpose}" if decision.purpose else f"AI planner: {decision.tool_name}({decision.arguments})" ), args=dict(decision.arguments), expected_text=decision.expected_outcome, purpose=decision.purpose, expected_outcome=decision.expected_outcome, prompt=decision.user_prompt or user_prompt, rationale=decision.text_output, thinking=decision.thinking, ) ] def goal_reached(self, *, goal: str, scene: Scene, context: TaskContext) -> bool: # Completion is signaled exclusively via the finish_task tool call # (mapped to an empty plan above), never via this hook. return False def _log(self, event: dict[str, Any]) -> None: if self.event_logger is not None: try: self.event_logger(event) except Exception: pass def _without_ocr(scene_json: dict[str, Any]) -> dict[str, Any]: cleaned = dict(scene_json) elements = cleaned.get("elements") if isinstance(elements, list): cleaned["elements"] = [ {key: value for key, value in element.items() if key not in {"source", "confidence", "foreground_color", "background_color"}} for element in elements if isinstance(element, dict) and element.get("source") != "ocr" ] cleaned.pop("ocr_elements", None) return cleaned def _sync_tool_results(context: TaskContext) -> None: for turn, result in zip(context.planner_history, context.step_results, strict=False): if turn.get("tool_result") is None: turn["tool_result"] = result.to_dict() def _conversation_arguments(decision: Any) -> dict[str, Any]: arguments = dict(decision.arguments) if decision.purpose is not None: arguments["purpose"] = decision.purpose if decision.expected_outcome is not None: arguments["expected_outcome"] = decision.expected_outcome return arguments def _accepts_history(method: Any) -> bool: parameters = signature(method).parameters.values() return any( parameter.name == "history" or parameter.kind == Parameter.VAR_KEYWORD for parameter in parameters )