from __future__ import annotations from collections.abc import Callable from dataclasses import fields from typing import Any from device.manager import DriverFactory from driver.android_driver import AndroidDriver, AndroidDriverConfig from driver.wda_driver import WDADriver, WDADriverConfig DriverFactoryBuilder = Callable[[dict[str, Any]], DriverFactory] def build_wda_driver_factory(connection_info: dict[str, Any]) -> DriverFactory: config_fields = {field.name for field in fields(WDADriverConfig)} data = dict(connection_info) raw_extra_capabilities = data.pop("extra_capabilities", {}) if not isinstance(raw_extra_capabilities, dict): raise ValueError("extra_capabilities must be an object") config_values: dict[str, Any] = {} for key in list(data): if key in config_fields and key != "extra_capabilities": config_values[key] = data.pop(key) config = WDADriverConfig( **config_values, extra_capabilities={**raw_extra_capabilities, **data}, ) return lambda: WDADriver(config) def build_android_driver_factory(connection_info: dict[str, Any]) -> DriverFactory: config_fields = {field.name for field in fields(AndroidDriverConfig)} data = dict(connection_info) raw_extra_capabilities = data.pop("extra_capabilities", {}) if not isinstance(raw_extra_capabilities, dict): raise ValueError("extra_capabilities must be an object") config_values: dict[str, Any] = {} for key in list(data): if key in config_fields and key != "extra_capabilities": config_values[key] = data.pop(key) config = AndroidDriverConfig( **config_values, extra_capabilities={**raw_extra_capabilities, **data}, ) return lambda: AndroidDriver(config) SUPPORTED_DRIVER_TYPES: dict[str, DriverFactoryBuilder] = { "wda": build_wda_driver_factory, "uiautomator2": build_android_driver_factory, } def build_driver_factory( driver_type: str, connection_info: dict[str, Any], ) -> DriverFactory: builder = SUPPORTED_DRIVER_TYPES.get(driver_type) if not builder: raise ValueError(f"unsupported driver_type: {driver_type}") return builder(connection_info) def register_driver_type(driver_type: str, factory: DriverFactoryBuilder) -> None: """Register a new ``driver_type`` -> factory-builder mapping. This is the extension point external code (e.g. ``cloud.plugins``' driver-kind plugin wiring) uses to add a new driver type without editing this module. ``factory`` must be a callable accepting a ``connection_info`` dict and returning a ``DriverFactory`` (the same shape as :func:`build_wda_driver_factory`); once registered, ``build_driver_factory(driver_type, ...)`` can construct drivers of the new type. """ if not driver_type: raise ValueError("driver_type must be a non-empty string") if not callable(factory): raise ValueError(f"factory for driver_type {driver_type!r} must be callable") if driver_type in SUPPORTED_DRIVER_TYPES: raise ValueError(f"driver_type {driver_type!r} is already registered") SUPPORTED_DRIVER_TYPES[driver_type] = factory