Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 55 additions & 0 deletions backend/node_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,59 @@ def _normalise_numeric(value: Optional[str], default: int) -> int:
raise NodeRegistrationError(f"Invalid numeric value '{value}'") from exc


def _session_value_as_str(session_data: Dict, *keys: str) -> Optional[str]:
for key in keys:
if key not in session_data:
continue
value = session_data.get(key)
if isinstance(value, str):
value = value.strip()
if value:
return value
elif isinstance(value, (int, float)):
return str(value)
return None


def _populate_metadata_from_session(node: Dict) -> None:
resources = node.get("resources")
if not isinstance(resources, dict):
return

session_data = resources.get("session_data")
if not isinstance(session_data, dict):
return

if not _strip_or_none(node.get("platform")):
platform = _session_value_as_str(session_data, "platformName", "appium:platformName")
if platform:
node["platform"] = platform

if not _strip_or_none(node.get("platform_version")):
version = _session_value_as_str(
session_data, "platformVersion", "appium:platformVersion"
)
if version:
node["platform_version"] = version

device_name = node.get("device_name") or node.get("deviceName")
if not _strip_or_none(device_name):
name = _session_value_as_str(
session_data,
"device_name",
"deviceName",
"appium:deviceName",
)
if name:
node.setdefault("deviceName", name)
node.setdefault("device_name", name)

if not _strip_or_none(node.get("udid")):
udid = _session_value_as_str(session_data, "udid", "appium:udid")
if udid:
node["udid"] = udid


def _normalise_node_payload(raw_node: Dict) -> Dict:
node = dict(raw_node)
node_id = node.get("id") or str(uuid.uuid4())
Expand All @@ -101,6 +154,8 @@ def _normalise_node_payload(raw_node: Dict) -> Dict:
logger.warning("Failed to decode resources JSON for node %s", node_id)
node.pop("resources", None)

_populate_metadata_from_session(node)

return node


Expand Down
75 changes: 71 additions & 4 deletions backend/proxy_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,21 @@ def _extract_requested_platform(payload: Dict) -> Optional[str]:
return _extract_capability_value(payload, ("platformName", "appium:platformName"))


def _extract_requested_platform_version(payload: Dict) -> Optional[str]:
"""Return the requested platform version, normalised for comparison."""

value = _extract_capability_value(payload, ("platformVersion", "appium:platformVersion"))
return _normalise_version(value)


def _extract_requested_avd(payload: Dict) -> Optional[str]:
"""Return the requested Android Virtual Device name if present."""

value = _extract_capability_value(payload, ("appium:avd", "avd"))
value = _normalise_str(value)
return value.lower() if value else None


def _normalise_str(value: Optional[str]) -> Optional[str]:
if isinstance(value, str):
value = value.strip()
Expand All @@ -138,6 +153,19 @@ def _normalise_str(value: Optional[str]) -> Optional[str]:
return None


def _normalise_version(value: Optional[str]) -> Optional[str]:
if value is None:
return None
if not isinstance(value, str):
value = str(value)
value = value.strip()
if not value:
return None
while value.endswith(".0"):
value = value[:-2]
return value.lower()


def _merge_session_capabilities(
body: bytes, headers: Dict[str, str], session_data: Dict, *, payload: Optional[Dict] = None
) -> Tuple[bytes, Dict[str, str], Optional[Dict]]:
Expand Down Expand Up @@ -231,6 +259,8 @@ async def forward_request(request: Request, path: str):
requested_platform: Optional[str] = None
requested_device_name: Optional[str] = None
requested_udid: Optional[str] = None
requested_avd: Optional[str] = None
requested_platform_version: Optional[str] = None

if not session_id and request.method == "POST":
payload = _parse_session_payload(body)
Expand All @@ -244,6 +274,12 @@ async def forward_request(request: Request, path: str):
udid = _extract_capability_value(payload, ("appium:udid", "udid"))
if udid:
requested_udid = udid.lower()
avd = _extract_requested_avd(payload)
if avd:
requested_avd = avd
platform_version = _extract_requested_platform_version(payload)
if platform_version:
requested_platform_version = platform_version

if session_id:
target_node_id = await redis_client.hget(SESSION_MAP_KEY, session_id)
Expand Down Expand Up @@ -282,10 +318,13 @@ async def forward_request(request: Request, path: str):
session_data = None

if requested_platform:
node_platform = (node.get("platform") or "").strip()
if not node_platform:
continue
if node_platform.lower() != requested_platform:
node_platform = _normalise_str(node.get("platform"))
if not node_platform and session_data:
node_platform = _normalise_str(
session_data.get("platformName")
or session_data.get("appium:platformName")
)
if not node_platform or node_platform.lower() != requested_platform:
continue

if requested_udid:
Expand All @@ -310,6 +349,34 @@ async def forward_request(request: Request, path: str):
if not node_device_name or node_device_name.lower() != requested_device_name:
continue

if requested_avd:
node_avd = _normalise_str(node.get("avd"))
if node_avd:
node_avd = node_avd.lower()
if not node_avd and session_data:
node_avd = _normalise_str(
session_data.get("appium:avd") or session_data.get("avd")
)
if node_avd:
node_avd = node_avd.lower()
if not node_avd or node_avd != requested_avd:
continue

if requested_platform_version:
node_platform_version = _normalise_version(
node.get("platform_version") or node.get("platformVersion")
)
if not node_platform_version and session_data:
node_platform_version = _normalise_version(
session_data.get("appium:platformVersion")
or session_data.get("platformVersion")
)
if (
not node_platform_version
or node_platform_version != requested_platform_version
):
continue

if request.method == "DELETE" or (status == "online" and active_sessions < max_sessions):
target_node = node
target_node_id = node_id
Expand Down