Tests / Test apps.device-host-agent.tests.test_mcp_token.test_load_or_create_concurrent_calls_do_not_corrupt failed
160 lines
6.0 KiB
Python
160 lines
6.0 KiB
Python
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
|
|
)
|