"""基线取证测试:mock ADB/uiautomator2,不连接手机。""" from __future__ import annotations from pathlib import Path import base64 from io import BytesIO import sys import tempfile import unittest from PIL import Image from uiautomator2.exceptions import HTTPTimeoutError CLIENT_ROOT = Path(__file__).resolve().parents[2] sys.path.insert(0, str(CLIENT_ROOT / "src")) sys.path.insert(0, str(CLIENT_ROOT / "scripts")) from cmbuyer_client.device.adb import AdbDevice, DeviceInspection from cmbuyer_client.device.baseline import ( BaselineCaptureError, BaselineCaptureTimeoutError, DeviceBaselineCapturer, NoReconnectUiautomatorConnector, PDD_PACKAGE, ) from capture_device_baseline import parse_arguments, validate_arguments SERIAL = "USB-serial-for-test" class StaticAdbClient: def __init__(self) -> None: self.serials: list[str] = [] def inspect(self, serial: str) -> DeviceInspection: self.serials.append(serial) return DeviceInspection( device=AdbDevice(serial=serial, state="device", model="Test Model"), model="Test Model", android_version="16", ) class FakeUiDevice: def __init__(self, fail_dump: bool = False) -> None: self.fail_dump = fail_dump self.rpc_calls: list[tuple[str, object, float]] = [] self.app_info_calls: list[str] = [] def app_info(self, package_name: str) -> dict[str, str]: self.app_info_calls.append(package_name) return {"versionName": "8.17.0"} def jsonrpc_call(self, method: str, params: object = None, timeout: float = 10) -> str: self.rpc_calls.append((method, params, timeout)) if method == "takeScreenshot": image_data = BytesIO() Image.new("RGB", (1, 1), color="white").save(image_data, format="PNG") return base64.b64encode(image_data.getvalue()).decode("ascii") if method != "dumpWindowHierarchy": raise AssertionError(f"unexpected method: {method}") if self.fail_dump: raise RuntimeError("mock dump failed") return "" class BaselineCaptureTests(unittest.TestCase): def test_capture_writes_hashes_without_xml_or_raw_serial_in_manifest(self) -> None: adb = StaticAdbClient() device = FakeUiDevice() capturer = DeviceBaselineCapturer(adb, lambda serial: device, timeout_seconds=7.5) with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "baseline" result = capturer.capture(SERIAL, output) manifest = result.manifest_path.read_text(encoding="utf-8") self.assertEqual(adb.serials, [SERIAL]) self.assertEqual(device.app_info_calls, [PDD_PACKAGE]) self.assertEqual( device.rpc_calls, [ ("takeScreenshot", [1, 80], 7.5), ("dumpWindowHierarchy", [False, 50], 7.5), ], ) self.assertTrue(result.screenshot_path.is_file()) self.assertTrue(result.hierarchy_path.is_file()) self.assertIn('"sha256"', manifest) self.assertNotIn("page body must stay out of manifest", manifest) self.assertNotIn(SERIAL, manifest) self.assertIn('"channel": "usb"', manifest) def test_capture_failure_cleans_staging_and_does_not_publish_partial_output(self) -> None: device = FakeUiDevice(fail_dump=True) capturer = DeviceBaselineCapturer(StaticAdbClient(), lambda serial: device, timeout_seconds=5) with tempfile.TemporaryDirectory() as directory: parent = Path(directory) output = parent / "baseline" with self.assertRaises(BaselineCaptureError) as raised: capturer.capture(SERIAL, output) self.assertFalse(output.exists()) self.assertEqual(list(parent.iterdir()), []) self.assertNotIn("mock dump failed", str(raised.exception)) def test_existing_output_is_never_overwritten(self) -> None: with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "baseline" output.mkdir() sentinel = output / "keep.txt" sentinel.write_text("preserve", encoding="utf-8") capturer = DeviceBaselineCapturer(StaticAdbClient(), lambda serial: FakeUiDevice(), timeout_seconds=5) with self.assertRaises(BaselineCaptureError): capturer.capture(SERIAL, output) self.assertEqual(sentinel.read_text(encoding="utf-8"), "preserve") def test_no_reconnect_connector_passes_only_current_adb_device_object(self) -> None: class ListedDevice: serial = SERIAL listed = ListedDevice() connected: list[object] = [] connector = NoReconnectUiautomatorConnector(lambda: [listed], lambda device: connected.append(device) or FakeUiDevice()) connector(SERIAL) self.assertEqual(connected, [listed]) def test_no_reconnect_connector_refuses_disappeared_serial(self) -> None: connector = NoReconnectUiautomatorConnector(lambda: [], lambda device: FakeUiDevice()) with self.assertRaises(BaselineCaptureError) as raised: connector(SERIAL) self.assertIn("拒绝自动重连", str(raised.exception)) def test_connector_exception_is_redacted_and_publishes_no_partial_output(self) -> None: def failing_connector(serial: str) -> FakeUiDevice: raise RuntimeError(f"third party leaked {serial}") capturer = DeviceBaselineCapturer(StaticAdbClient(), failing_connector, timeout_seconds=5) with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "baseline" with self.assertRaises(BaselineCaptureError) as raised: capturer.capture(SERIAL, output) self.assertNotIn(SERIAL, str(raised.exception)) self.assertFalse(output.exists()) self.assertEqual(list(Path(directory).iterdir()), []) def test_invalid_screenshot_base64_syntax_fails_closed_without_partial_output(self) -> None: class InvalidScreenshotDevice(FakeUiDevice): def jsonrpc_call(self, method: str, params: object = None, timeout: float = 10) -> str: if method == "takeScreenshot": valid = super().jsonrpc_call(method, params, timeout) return valid[:12] + "!" + valid[12:] return super().jsonrpc_call(method, params, timeout) capturer = DeviceBaselineCapturer(StaticAdbClient(), lambda serial: InvalidScreenshotDevice(), timeout_seconds=5) with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "baseline" with self.assertRaises(BaselineCaptureError) as raised: capturer.capture(SERIAL, output) self.assertIn("Base64 语法无效", str(raised.exception)) self.assertNotIn("!", str(raised.exception)) self.assertFalse(output.exists()) self.assertEqual(list(Path(directory).iterdir()), []) def test_invalid_padding_and_unapproved_ascii_whitespace_fail_closed(self) -> None: invalid_insertions = { "padding": lambda value: value[:-1], "vertical-tab": lambda value: value[:12] + "\v" + value[12:], "form-feed": lambda value: value[:12] + "\f" + value[12:], } for name, make_invalid in invalid_insertions.items(): with self.subTest(name=name), tempfile.TemporaryDirectory() as directory: class InvalidScreenshotDevice(FakeUiDevice): def jsonrpc_call(self, method: str, params: object = None, timeout: float = 10) -> str: value = super().jsonrpc_call(method, params, timeout) if method == "takeScreenshot": return make_invalid(value) return value output = Path(directory) / "baseline" capturer = DeviceBaselineCapturer( StaticAdbClient(), lambda serial: InvalidScreenshotDevice(), timeout_seconds=5, ) with self.assertRaises(BaselineCaptureError) as raised: capturer.capture(SERIAL, output) self.assertIn("Base64 语法无效", str(raised.exception)) self.assertNotIn(SERIAL, str(raised.exception)) self.assertFalse(output.exists()) self.assertEqual(list(Path(directory).iterdir()), []) def test_base64_decoded_nonimage_fails_closed_without_partial_output(self) -> None: class NonImageScreenshotDevice(FakeUiDevice): def jsonrpc_call(self, method: str, params: object = None, timeout: float = 10) -> str: if method == "takeScreenshot": return base64.b64encode(b"not an image").decode("ascii") return super().jsonrpc_call(method, params, timeout) capturer = DeviceBaselineCapturer(StaticAdbClient(), lambda serial: NonImageScreenshotDevice(), timeout_seconds=5) with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "baseline" with self.assertRaises(BaselineCaptureError) as raised: capturer.capture(SERIAL, output) self.assertIn("图像数据无效", str(raised.exception)) self.assertNotIn("not an image", str(raised.exception)) self.assertFalse(output.exists()) self.assertEqual(list(Path(directory).iterdir()), []) def test_ascii_base64_whitespace_is_normalized_before_strict_decode(self) -> None: class WhitespaceScreenshotDevice(FakeUiDevice): def jsonrpc_call(self, method: str, params: object = None, timeout: float = 10) -> str: value = super().jsonrpc_call(method, params, timeout) if method == "takeScreenshot": return value[:10] + " \t\r\n" + value[10:30] + "\n" + value[30:] return value capturer = DeviceBaselineCapturer(StaticAdbClient(), lambda serial: WhitespaceScreenshotDevice(), timeout_seconds=5) with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "baseline" result = capturer.capture(SERIAL, output) self.assertTrue(result.screenshot_path.is_file()) with Image.open(result.screenshot_path) as image: self.assertEqual(image.size, (1, 1)) def test_invalid_or_non_hierarchy_xml_fails_closed_without_partial_output(self) -> None: class InvalidHierarchyDevice(FakeUiDevice): def jsonrpc_call(self, method: str, params: object = None, timeout: float = 10) -> str: if method == "dumpWindowHierarchy": return "" return super().jsonrpc_call(method, params, timeout) capturer = DeviceBaselineCapturer(StaticAdbClient(), lambda serial: InvalidHierarchyDevice(), timeout_seconds=5) with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "baseline" with self.assertRaises(BaselineCaptureError) as raised: capturer.capture(SERIAL, output) self.assertNotIn("not-hierarchy", str(raised.exception)) self.assertFalse(output.exists()) self.assertEqual(list(Path(directory).iterdir()), []) def test_rpc_timeout_is_distinct_redacted_and_does_not_publish_partial_output(self) -> None: class TimeoutRpcDevice(FakeUiDevice): def jsonrpc_call(self, method: str, params: object = None, timeout: float = 10) -> str: raise HTTPTimeoutError(f"raw serial={SERIAL} xml=") capturer = DeviceBaselineCapturer(StaticAdbClient(), lambda serial: TimeoutRpcDevice(), timeout_seconds=5) with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "baseline" with self.assertRaises(BaselineCaptureTimeoutError) as raised: capturer.capture(SERIAL, output) self.assertIn("超时", str(raised.exception)) self.assertNotIn(SERIAL, str(raised.exception)) self.assertNotIn("hierarchy", str(raised.exception)) self.assertFalse(output.exists()) self.assertEqual(list(Path(directory).iterdir()), []) def test_cli_validation_rejects_empty_serial_and_nonpositive_timeout(self) -> None: empty_serial = parse_arguments(["--serial", "", "--output-dir", "baseline"]) with self.assertRaisesRegex(ValueError, "非空 --serial"): validate_arguments(empty_serial) nonpositive_timeout = parse_arguments(["--serial", SERIAL, "--output-dir", "baseline", "--timeout", "0"]) with self.assertRaisesRegex(ValueError, "必须大于 0"): validate_arguments(nonpositive_timeout)