Files
agentic-mobile-control/tests/test_cloud_planner_decision_endpoint.py
T
2026-07-15 18:14:28 +08:00

385 lines
13 KiB
Python

from __future__ import annotations
from datetime import UTC, datetime, timedelta
from fastapi import FastAPI
from fastapi.testclient import TestClient
from sqlalchemy.orm import Session
from cloud.auth import BearerCredential, ConfiguredBearerAuthProvider
from cloud.config import CloudConfig
from cloud.db_models import TaskAttemptRow
from cloud.internal_api.api import create_internal_router
from cloud.pool import DevicePool
from cloud.store import CloudStore
from runtime.tool_calling_client import ToolCallDecision, ToolCallUnavailable
from runtime.tool_specs import ToolSpec
class _FakeToolCallingClient:
"""Stand-in ``ToolCallingClient`` for endpoint tests -- never calls a real
LLM provider."""
def __init__(
self, decision: ToolCallDecision | None = None, error: str | None = None
):
self.decision = decision
self.error = error
self.calls: list[dict[str, object]] = []
def decide(
self,
*,
system_prompt: str,
user_prompt: str,
screenshot: bytes | None,
tools: list[ToolSpec],
timeout: float,
) -> ToolCallDecision:
self.calls.append(
{
"system_prompt": system_prompt,
"user_prompt": user_prompt,
"screenshot": screenshot,
"tools": tools,
"timeout": timeout,
}
)
if self.error is not None:
raise ToolCallUnavailable(self.error)
assert self.decision is not None
return self.decision
def _build_client(
tmp_path, *, fake_client: _FakeToolCallingClient
) -> tuple[TestClient, _FakeToolCallingClient, DevicePool]:
pool = DevicePool(
CloudStore(tmp_path / "internal.sqlite3"),
CloudConfig(stale_after_seconds=60),
)
auth_provider = ConfiguredBearerAuthProvider(
[
BearerCredential(principal_id="agent-a", token="token-a", host_id="host-a"),
BearerCredential(principal_id="agent-b", token="token-b", host_id="host-b"),
]
)
app = FastAPI()
app.include_router(
create_internal_router(
pool=pool,
auth_provider=auth_provider,
planner_client_factory=lambda: fake_client,
)
)
return TestClient(app), fake_client, pool
def _decision_payload(**overrides: object) -> dict[str, object]:
payload = {
"host_id": "host-a",
"system_prompt": "you are a planner",
"user_prompt": "tap the login button",
"screenshot_base64": None,
"tools": [
{
"name": "tap",
"description": "tap an element",
"parameters": {"type": "object", "properties": {}},
}
],
"timeout_seconds": 30.0,
}
payload.update(overrides)
return payload
def test_authenticated_host_resolves_planner_decision(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(
tool_name="tap",
arguments={"x": 1, "y": 2},
text_output="The login button is visible. Opening it.",
thinking="The next screen should be the account page.",
purpose="Open the login screen.",
expected_outcome="The login form is visible.",
)
)
client, fake_client, _pool = _build_client(tmp_path, fake_client=fake_client)
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(),
)
assert response.status_code == 200
assert response.json() == {
"tool_name": "tap",
"arguments": {"x": 1, "y": 2},
"rationale": "The login button is visible. Opening it.",
"thinking": "The next screen should be the account page.",
"purpose": "Open the login screen.",
"expected_outcome": "The login form is visible.",
}
assert len(fake_client.calls) == 1
assert fake_client.calls[0]["system_prompt"] == "you are a planner"
assert fake_client.calls[0]["timeout"] == 30.0
def test_planner_decision_decodes_screenshot_base64(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(tool_name="tap", arguments={})
)
client, fake_client, _pool = _build_client(tmp_path, fake_client=fake_client)
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(screenshot_base64="aGVsbG8="),
)
assert response.status_code == 200
assert fake_client.calls[0]["screenshot"] == b"hello"
def test_invalid_screenshot_base64_is_rejected(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(tool_name="tap", arguments={})
)
client, fake_client, _pool = _build_client(tmp_path, fake_client=fake_client)
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(screenshot_base64="not-valid-base64!!"),
)
assert response.status_code == 422
assert fake_client.calls == []
def test_unauthenticated_request_is_rejected(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(tool_name="tap", arguments={})
)
client, fake_client, _pool = _build_client(tmp_path, fake_client=fake_client)
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
json=_decision_payload(),
)
assert response.status_code == 401
assert fake_client.calls == []
def test_foreign_host_token_cannot_request_another_hosts_decision(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(tool_name="tap", arguments={})
)
client, fake_client, _pool = _build_client(tmp_path, fake_client=fake_client)
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-b"},
json=_decision_payload(),
)
assert response.status_code == 403
assert fake_client.calls == []
def test_path_and_payload_host_id_mismatch_is_rejected(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(tool_name="tap", arguments={})
)
client, fake_client, _pool = _build_client(tmp_path, fake_client=fake_client)
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(host_id="host-b"),
)
assert response.status_code == 422
assert fake_client.calls == []
def test_provider_failure_returns_structured_error_without_crashing(tmp_path) -> None:
fake_client = _FakeToolCallingClient(error="anthropic timed out")
client, fake_client, _pool = _build_client(tmp_path, fake_client=fake_client)
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(),
)
assert response.status_code == 502
assert response.json() == {
"code": "planner_unavailable",
"detail": "anthropic timed out",
}
def _seed_attempt(pool: DevicePool, task_id: str, host_id: str = "host-a") -> None:
"""Insert a task_attempts row so _validate_planner_context passes."""
now = datetime.now(tz=UTC)
with Session(pool.store.engine) as session, session.begin():
session.add(
TaskAttemptRow(
task_id=task_id,
attempt=1,
lease_id="lease-x",
host_id=host_id,
device_id="device-x",
status="dispatched",
lease_expires_at=(now + timedelta(minutes=5)).isoformat(),
created_at=now.isoformat(),
completed_at=None,
failure_reason=None,
result_json=None,
)
)
def test_successful_decision_is_persisted_with_correct_fields(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(
tool_name="tap",
arguments={"x": 10, "y": 20},
text_output="The settings tab is visible. Opening it.",
thinking="A tap should navigate to settings.",
purpose="Open settings.",
expected_outcome="The settings page is visible.",
)
)
client, fake_client, pool = _build_client(tmp_path, fake_client=fake_client)
_seed_attempt(pool, "task-log-1")
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(
task_id="task-log-1",
attempt=1,
lease_id="lease-x",
system_prompt="you are a test planner",
user_prompt="tap the button now",
),
)
assert response.status_code == 200
decisions = pool.store.list_planner_decisions(task_id="task-log-1", attempt=1)
assert len(decisions) == 1
row = decisions[0]
assert row.step_index == 1
assert row.system_prompt == "you are a test planner"
assert row.user_prompt == "tap the button now"
assert row.tool_name == "tap"
assert '"x": 10' in row.arguments_json
assert '"y": 20' in row.arguments_json
assert row.rationale == "The settings tab is visible. Opening it."
assert row.thinking == "A tap should navigate to settings."
assert row.purpose == "Open settings."
assert row.expected_outcome == "The settings page is visible."
def test_failed_decision_persists_nothing(tmp_path) -> None:
fake_client = _FakeToolCallingClient(error="model is down")
client, fake_client, pool = _build_client(tmp_path, fake_client=fake_client)
_seed_attempt(pool, "task-log-fail")
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(
task_id="task-log-fail",
attempt=1,
lease_id="lease-x",
),
)
assert response.status_code == 502
assert pool.store.list_planner_decisions(task_id="task-log-fail", attempt=1) == []
def test_screenshot_bytes_are_never_persisted_to_decision_log(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(tool_name="tap", arguments={})
)
client, fake_client, pool = _build_client(tmp_path, fake_client=fake_client)
_seed_attempt(pool, "task-log-screenshot")
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(
task_id="task-log-screenshot",
attempt=1,
lease_id="lease-x",
screenshot_base64="aGVsbG8gd29ybGQ=",
),
)
assert response.status_code == 200
decisions = pool.store.list_planner_decisions(
task_id="task-log-screenshot", attempt=1
)
assert len(decisions) == 1
# The row has no screenshot column at all — verify the raw row text
# does not contain the screenshot bytes.
row_repr = repr(decisions[0])
assert "hello world" not in row_repr
assert "aGVsbG8" not in row_repr
def test_decision_without_task_context_is_not_persisted(tmp_path) -> None:
"""When task_id/attempt are absent, the log insert is skipped."""
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(tool_name="tap", arguments={})
)
client, fake_client, pool = _build_client(tmp_path, fake_client=fake_client)
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(), # no task_id / attempt / lease_id
)
assert response.status_code == 200
# No task_id to query by — verify no rows exist at all via direct SQL.
from sqlalchemy import func, select
from cloud.db_models import PlannerDecisionLogRow
with Session(pool.store.engine) as session:
count = session.scalar(select(func.count()).select_from(PlannerDecisionLogRow))
assert count == 0
def test_multiple_decisions_get_incrementing_step_index(tmp_path) -> None:
fake_client = _FakeToolCallingClient(
decision=ToolCallDecision(tool_name="tap", arguments={})
)
client, fake_client, pool = _build_client(tmp_path, fake_client=fake_client)
_seed_attempt(pool, "task-log-multi")
for i in range(3):
response = client.post(
"/internal/v1/hosts/host-a/planner/decide",
headers={"Authorization": "Bearer token-a"},
json=_decision_payload(
task_id="task-log-multi",
attempt=1,
lease_id="lease-x",
user_prompt=f"step {i}",
),
)
assert response.status_code == 200
decisions = pool.store.list_planner_decisions(task_id="task-log-multi", attempt=1)
assert [d.step_index for d in decisions] == [1, 2, 3]
assert [d.user_prompt for d in decisions] == ["step 0", "step 1", "step 2"]