198 lines
7.0 KiB
Python
198 lines
7.0 KiB
Python
import json
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from PIL import Image
|
|
|
|
from cockpit_agent.common.enums import (
|
|
ControlType,
|
|
Domain,
|
|
Operation,
|
|
Unit,
|
|
Zone,
|
|
)
|
|
from cockpit_agent.execution.mock_executor import MockExecutor
|
|
from cockpit_agent.grounding.adapter import GroundingResult
|
|
from cockpit_agent.grounding.semantic_target import SemanticTargetBuilder
|
|
from cockpit_agent.intent.schemas import SemanticAction
|
|
from cockpit_agent.perception.schemas import UIState
|
|
from cockpit_agent.pipeline import TaskPipeline
|
|
from cockpit_agent.planning.planner import Planner
|
|
|
|
|
|
class FakeIntentParser:
|
|
last_raw_output = '{"intent": "raw"}'
|
|
|
|
def parse(self, instruction: str) -> SemanticAction:
|
|
del instruction
|
|
return SemanticAction(
|
|
domain=Domain.CLIMATE,
|
|
function="temperature",
|
|
zone=Zone.DRIVER,
|
|
operation=Operation.SET,
|
|
state=None,
|
|
value=23,
|
|
unit=Unit.CELSIUS,
|
|
)
|
|
|
|
|
|
class FakeUIUnderstanding:
|
|
last_raw_output = '{"ui": "raw"}'
|
|
|
|
def __init__(self) -> None:
|
|
self.image_path: str | None = None
|
|
self.function_query: str | None = None
|
|
|
|
def understand(
|
|
self,
|
|
roi_image_path: str,
|
|
intent: SemanticAction,
|
|
function_grounding: GroundingResult,
|
|
) -> UIState:
|
|
self.image_path = roi_image_path
|
|
self.function_query = function_grounding.semantic_query
|
|
return UIState(
|
|
function=intent.function,
|
|
zone=intent.zone,
|
|
control_type=ControlType.STEPPER,
|
|
current_value=26,
|
|
step_value=1,
|
|
current_state=None,
|
|
min_value=None,
|
|
max_value=None,
|
|
options=None,
|
|
)
|
|
|
|
|
|
class FakeGrounding:
|
|
last_raw_output = '{"bbox_2d": [100, 200, 300, 400]}'
|
|
|
|
def __init__(self) -> None:
|
|
self.calls: list[tuple[str, str]] = []
|
|
|
|
def ground(self, image_path: str, semantic_query: str) -> GroundingResult:
|
|
self.calls.append((image_path, semantic_query))
|
|
with Image.open(image_path) as image:
|
|
width, height = image.size
|
|
if len(self.calls) == 1:
|
|
bbox = (20, 20, 80, 60)
|
|
else:
|
|
bbox = (6, 8, 18, 24)
|
|
x1, y1, x2, y2 = bbox
|
|
return GroundingResult(
|
|
semantic_query=semantic_query,
|
|
bbox_relative=(
|
|
x1 / width * 1000,
|
|
y1 / height * 1000,
|
|
x2 / width * 1000,
|
|
y2 / height * 1000,
|
|
),
|
|
bbox_pixel=bbox,
|
|
center_pixel=(round((x1 + x2) / 2), round((y1 + y2) / 2)),
|
|
image_width=width,
|
|
image_height=height,
|
|
raw_output=self.last_raw_output,
|
|
)
|
|
|
|
|
|
class FailingGrounding:
|
|
last_raw_output = "model returned no bbox"
|
|
|
|
def ground(self, image_path: str, semantic_query: str) -> GroundingResult:
|
|
del image_path, semantic_query
|
|
raise ValueError("cannot parse bbox")
|
|
|
|
|
|
class PipelineTest(unittest.TestCase):
|
|
TEMPERATURE_FUNCTION_TARGET = (
|
|
"主驾驶温度调节控件整体,"
|
|
"包含当前温度显示以及用于升高和降低温度的控制元素"
|
|
)
|
|
|
|
def _pipeline(self, grounding, ui=None) -> TaskPipeline:
|
|
return TaskPipeline(
|
|
intent_parser=FakeIntentParser(),
|
|
ui_understanding=ui or FakeUIUnderstanding(),
|
|
planner=Planner(),
|
|
target_builder=SemanticTargetBuilder(),
|
|
grounding=grounding,
|
|
executor=MockExecutor(),
|
|
roi_padding_ratio=0.10,
|
|
)
|
|
|
|
def test_function_grounding_precedes_roi_ui_and_action_grounding(self) -> None:
|
|
grounding = FakeGrounding()
|
|
ui = FakeUIUnderstanding()
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
root = Path(directory)
|
|
image = root / "input.png"
|
|
output = root / "output"
|
|
Image.new("RGB", (200, 100), color="white").save(image)
|
|
run = self._pipeline(grounding, ui).run(
|
|
image_path=str(image),
|
|
instruction="set temperature",
|
|
output_dir=str(output),
|
|
)
|
|
roi_exists = (output / "roi.jpg").is_file()
|
|
function_file_exists = (output / "function_grounding.json").is_file()
|
|
action_file_exists = (output / "action_grounding.json").is_file()
|
|
|
|
self.assertTrue(run["success"])
|
|
self.assertEqual(run["function_target"], self.TEMPERATURE_FUNCTION_TARGET)
|
|
self.assertEqual(grounding.calls[0][1], self.TEMPERATURE_FUNCTION_TARGET)
|
|
self.assertEqual(grounding.calls[1][1], "主驾驶温度降低控制")
|
|
self.assertTrue(grounding.calls[1][0].endswith("roi.jpg"))
|
|
self.assertTrue((ui.image_path or "").endswith("roi.jpg"))
|
|
self.assertEqual(ui.function_query, self.TEMPERATURE_FUNCTION_TARGET)
|
|
self.assertEqual(run["action_plan"]["function"], "temperature_decrease")
|
|
self.assertEqual(run["action_target"], "主驾驶温度降低控制")
|
|
self.assertEqual(
|
|
run["action_grounding"]["source"],
|
|
"roi_action_grounding",
|
|
)
|
|
self.assertEqual(run["mock_action"]["status"], "PROPOSED")
|
|
self.assertTrue(roi_exists)
|
|
self.assertTrue(function_file_exists)
|
|
self.assertTrue(action_file_exists)
|
|
self.assertEqual(run["raw_model_outputs"]["intent"], '{"intent": "raw"}')
|
|
self.assertEqual(run["raw_model_outputs"]["ui_understanding"], '{"ui": "raw"}')
|
|
self.assertIn("function_grounding", run["raw_model_outputs"])
|
|
self.assertIn("action_grounding", run["raw_model_outputs"])
|
|
|
|
def test_function_grounding_failure_is_recorded_before_ui(self) -> None:
|
|
ui = FakeUIUnderstanding()
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
root = Path(directory)
|
|
image = root / "input.png"
|
|
output = root / "output"
|
|
Image.new("RGB", (200, 100), color="white").save(image)
|
|
with self.assertRaisesRegex(ValueError, "cannot parse bbox"):
|
|
self._pipeline(FailingGrounding(), ui).run(
|
|
image_path=str(image),
|
|
instruction="set temperature",
|
|
output_dir=str(output),
|
|
)
|
|
failure = json.loads((output / "failure.json").read_text())
|
|
run = json.loads((output / "run.json").read_text())
|
|
|
|
self.assertEqual(failure["failed_stage"], "function_grounding")
|
|
self.assertEqual(failure["error_type"], "ValueError")
|
|
self.assertEqual(
|
|
failure["raw_model_output"],
|
|
"model returned no bbox",
|
|
)
|
|
self.assertEqual(run["function_target"], self.TEMPERATURE_FUNCTION_TARGET)
|
|
self.assertEqual(
|
|
run["raw_model_outputs"]["function_grounding"],
|
|
"model returned no bbox",
|
|
)
|
|
self.assertIsNone(run["function_grounding"])
|
|
self.assertIsNone(run["roi"])
|
|
self.assertIsNone(run["ui_state"])
|
|
self.assertIsNone(ui.image_path)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|