from __future__ import annotations from collections.abc import Callable from dataclasses import fields from typing import Any from device.manager import DriverFactory 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) SUPPORTED_DRIVER_TYPES: dict[str, DriverFactoryBuilder] = { "wda": build_wda_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)