#!/usr/bin/env python3
"""Shared WebKit WebDriver helpers for browser smoke tests."""
from __future__ import annotations

import json
import os
import socket
import subprocess
import time
from typing import Any
from urllib.request import Request, urlopen


class WebDriverError(RuntimeError):
    pass


def browser_display_available() -> bool:
    return bool(os.environ.get("DISPLAY") or os.environ.get("WAYLAND_DISPLAY"))


def mini_browser_session_requested(path: str, payload: dict[str, Any] | None) -> bool:
    if path != "/session" or not isinstance(payload, dict):
        return False
    capabilities = payload.get("capabilities") or {}
    always_match = capabilities.get("alwaysMatch") or {}
    return always_match.get("browserName") == "MiniBrowser"


def free_port() -> int:
    with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
        sock.bind(("127.0.0.1", 0))
        return int(sock.getsockname()[1])


def request_json(
    port: int,
    method: str,
    path: str,
    payload: dict[str, Any] | None = None,
    timeout: float = 10,
) -> dict[str, Any] | None:
    if mini_browser_session_requested(path, payload) and not browser_display_available():
        raise WebDriverError(
            "MiniBrowser requires DISPLAY/WAYLAND_DISPLAY. "
            "Run browser smoke tests with `xvfb-run -a ...` in this headless VM; "
            "otherwise WebKitWebDriver can hang until the /session timeout."
        )
    data = None if payload is None else json.dumps(payload).encode("utf-8")
    request = Request(
        f"http://127.0.0.1:{port}{path}",
        data=data,
        method=method,
        headers={"Content-Type": "application/json"},
    )
    try:
        with urlopen(request, timeout=timeout) as response:
            body = response.read().decode("utf-8")
    except TimeoutError as exc:
        raise WebDriverError(f"WebDriver request timed out after {timeout}s: {method} {path}") from exc
    return json.loads(body) if body else None


def wait_for_driver(port: int, proc: subprocess.Popen[bytes], timeout: float = 8) -> None:
    deadline = time.monotonic() + timeout
    while time.monotonic() < deadline:
        if proc.poll() is not None:
            raise WebDriverError(f"WebKitWebDriver exited early with status {proc.returncode}")
        try:
            status = request_json(port, "GET", "/status", timeout=0.5)
        except Exception:
            time.sleep(0.1)
            continue
        if status and status.get("value", {}).get("ready") is True:
            return
        time.sleep(0.1)
    raise WebDriverError("WebKitWebDriver did not become ready")


def execute_js(port: int, session_id: str, script: str, timeout: float = 8) -> Any:
    response = request_json(
        port,
        "POST",
        f"/session/{session_id}/execute/sync",
        {"script": script, "args": []},
        timeout=timeout,
    )
    if not response or "value" not in response:
        raise WebDriverError(f"bad execute/sync response: {response!r}")
    return response["value"]


def execute_async_js(port: int, session_id: str, script: str, timeout: float = 12) -> Any:
    response = request_json(
        port,
        "POST",
        f"/session/{session_id}/execute/async",
        {"script": script, "args": []},
        timeout=timeout,
    )
    if not response or "value" not in response:
        raise WebDriverError(f"bad execute/async response: {response!r}")
    return response["value"]


def find_element_id(port: int, session_id: str, selector: str, timeout: float = 8) -> str:
    response = request_json(
        port,
        "POST",
        f"/session/{session_id}/element",
        {"using": "css selector", "value": selector},
        timeout=timeout,
    )
    value = (response or {}).get("value") or {}
    element_id = value.get("element-6066-11e4-a52e-4f735466cecf") or value.get("ELEMENT")
    if not element_id:
        raise WebDriverError(f"could not find element {selector!r}: {response!r}")
    return str(element_id)


def click_element(port: int, session_id: str, element_id: str, timeout: float = 8) -> None:
    request_json(port, "POST", f"/session/{session_id}/element/{element_id}/click", {}, timeout=timeout)


def click_selector(port: int, session_id: str, selector: str, timeout: float = 8) -> bool:
    click_element(port, session_id, find_element_id(port, session_id, selector, timeout=timeout), timeout=timeout)
    return True
