diff --git a/python/tests/test_client_ws.py b/python/tests/test_client_ws.py index b6c2ad16..828cc830 100644 --- a/python/tests/test_client_ws.py +++ b/python/tests/test_client_ws.py @@ -65,7 +65,29 @@ _FIXTURE_MESSAGES = [ ] +#: Upgrade-request headers captured by the in-process server, so tests +#: can assert what the client actually put on the handshake. +_CAPTURED_UPGRADE_HEADERS: dict[str, str] = {} + + +def _upgrade_headers(websocket: Any) -> Any: + """Read the handshake request headers across websockets versions + (`websocket.request.headers` >= 14, `websocket.request_headers` <= 13).""" + req = getattr(websocket, "request", None) + if req is not None and getattr(req, "headers", None) is not None: + return req.headers + return getattr(websocket, "request_headers", {}) + + async def _handler(websocket: Any) -> None: + _CAPTURED_UPGRADE_HEADERS.clear() + try: + headers = _upgrade_headers(websocket) + auth = headers.get("Authorization") if hasattr(headers, "get") else None + if auth is not None: + _CAPTURED_UPGRADE_HEADERS["Authorization"] = auth + except Exception: + pass for msg in _FIXTURE_MESSAGES: await websocket.send(json.dumps(msg)) # Send one malformed frame to assert the client logs+drops it @@ -174,6 +196,160 @@ def test_sensing_client_decoder_directly() -> None: assert msg.rssi is None +# ─── Auth: bearer token on the WS upgrade (issue #1395) ────────────── + + +def _auth_header_from_kwargs(kwargs: dict) -> Any: + """Pull the Authorization value out of whichever header kwarg the + installed `websockets` uses (`additional_headers` >= 14, + `extra_headers` <= 13). Returns None if no header kwarg was passed.""" + for key in ("additional_headers", "extra_headers"): + if key in kwargs: + return dict(kwargs[key]).get("Authorization") + return None + + +class _DummyWS: + async def close(self) -> None: + pass + + +class _CapturingConnect: + """Stand-in for `websockets.connect` that records the kwargs it was + called with and returns an awaitable yielding a dummy connection.""" + + def __init__(self) -> None: + self.calls: list[tuple[str, dict]] = [] + + def __call__(self, url: str, **kwargs: Any) -> Any: + self.calls.append((url, kwargs)) + + async def _coro() -> _DummyWS: + return _DummyWS() + + return _coro() + + @property + def last_kwargs(self) -> dict: + return self.calls[-1][1] + + +async def test_token_from_constructor_sets_auth_header(monkeypatch: Any) -> None: + from wifi_densepose.client import ws as ws_mod + + fake = _CapturingConnect() + monkeypatch.setattr(ws_mod.websockets, "connect", fake) + + async with SensingClient("ws://x/ws/sensing", token="tok-abc"): + pass + + assert _auth_header_from_kwargs(fake.last_kwargs) == "Bearer tok-abc" + + +async def test_token_from_env_sets_auth_header(monkeypatch: Any) -> None: + from wifi_densepose.client import ws as ws_mod + + fake = _CapturingConnect() + monkeypatch.setattr(ws_mod.websockets, "connect", fake) + monkeypatch.setenv("RUVIEW_API_TOKEN", "env-tok-123") + + async with SensingClient("ws://x/ws/sensing"): + pass + + assert _auth_header_from_kwargs(fake.last_kwargs) == "Bearer env-tok-123" + + +async def test_constructor_token_overrides_env(monkeypatch: Any) -> None: + from wifi_densepose.client import ws as ws_mod + + fake = _CapturingConnect() + monkeypatch.setattr(ws_mod.websockets, "connect", fake) + monkeypatch.setenv("RUVIEW_API_TOKEN", "env-tok") + + async with SensingClient("ws://x/ws/sensing", token="ctor-tok"): + pass + + assert _auth_header_from_kwargs(fake.last_kwargs) == "Bearer ctor-tok" + + +async def test_no_token_sends_no_auth_header(monkeypatch: Any) -> None: + from wifi_densepose.client import ws as ws_mod + + fake = _CapturingConnect() + monkeypatch.setattr(ws_mod.websockets, "connect", fake) + monkeypatch.delenv("RUVIEW_API_TOKEN", raising=False) + + async with SensingClient("ws://x/ws/sensing"): + pass + + assert _auth_header_from_kwargs(fake.last_kwargs) is None + # Auth-disabled path must not smuggle either header kwarg in. + assert "additional_headers" not in fake.last_kwargs + assert "extra_headers" not in fake.last_kwargs + + +async def test_empty_token_sends_no_auth_header(monkeypatch: Any) -> None: + """An explicitly empty token (or empty env var) means 'no auth'.""" + from wifi_densepose.client import ws as ws_mod + + fake = _CapturingConnect() + monkeypatch.setattr(ws_mod.websockets, "connect", fake) + monkeypatch.setenv("RUVIEW_API_TOKEN", "") + + async with SensingClient("ws://x/ws/sensing"): + pass + + assert _auth_header_from_kwargs(fake.last_kwargs) is None + + +@pytest.mark.parametrize( + "header_param,expected", + [ + ("additional_headers", "additional_headers"), # websockets >= 14 + ("extra_headers", "extra_headers"), # websockets <= 13 + ], +) +def test_select_header_kwarg_across_websockets_versions( + header_param: str, expected: str +) -> None: + """Version-compat: the kwarg is chosen by inspecting the installed + `connect` signature, so both the pre-14 (`extra_headers`) and + post-14 (`additional_headers`) conventions resolve correctly without + two websockets installs.""" + from wifi_densepose.client.ws import _select_header_kwarg + + # Build a fake `connect` whose signature carries only the one kwarg + # the emulated websockets version would expose. + ns: dict = {} + exec( + f"def fake_connect(uri, *, {header_param}=None, ping_interval=None): ...", + ns, + ) + assert _select_header_kwarg(ns["fake_connect"]) == expected + + +def test_select_header_kwarg_matches_installed_websockets() -> None: + """On whatever `websockets` is actually installed, the chosen kwarg + must be a real parameter of `websockets.connect`.""" + import inspect + + import websockets + + from wifi_densepose.client.ws import _select_header_kwarg + + chosen = _select_header_kwarg(websockets.connect) + assert chosen in inspect.signature(websockets.connect).parameters + + +async def test_auth_header_reaches_server_end_to_end(ws_server: str) -> None: + """Real in-process server: assert the bearer actually arrives on the + upgrade request (proves the header is wired to the live handshake, + not just the connect kwargs).""" + async with SensingClient(ws_server, token="e2e-token") as client: + await client.recv_one(timeout=2.0) + assert _CAPTURED_UPGRADE_HEADERS.get("Authorization") == "Bearer e2e-token" + + def test_sensing_client_decoder_handles_None_subfields() -> None: """When the sensing-server explicitly emits null for HR/BR (no measurement yet), the client should propagate None, not crash.""" diff --git a/python/wifi_densepose/client/ws.py b/python/wifi_densepose/client/ws.py index 385ff7f0..d9a712cd 100644 --- a/python/wifi_densepose/client/ws.py +++ b/python/wifi_densepose/client/ws.py @@ -31,8 +31,10 @@ asyncio.run(main()) from __future__ import annotations import asyncio +import inspect import json import logging +import os from dataclasses import dataclass, field from typing import Any, AsyncIterator, Optional @@ -48,6 +50,33 @@ except ImportError: # pragma: no cover log = logging.getLogger(__name__) +#: Environment variable the sensing-server bearer token is read from by +#: default. Mirrors the TypeScript MCP client (tools/ruview-mcp). +TOKEN_ENV_VAR = "RUVIEW_API_TOKEN" + + +def _select_header_kwarg(connect_fn: Any) -> str: + """Return the ``websockets.connect`` keyword for extra request headers. + + The keyword was renamed inside the ``websockets>=12`` range this + package supports: ``<= 13`` accepts ``extra_headers``, ``>= 14`` + accepts ``additional_headers``. We inspect the actual signature of + the installed ``connect`` rather than guessing from ``__version__``, + so a version bump that renames the kwarg again is handled by + detection instead of raising ``TypeError`` at connect time. + """ + try: + params = inspect.signature(connect_fn).parameters + except (TypeError, ValueError): # pragma: no cover — no introspectable sig + return "additional_headers" + if "additional_headers" in params: + return "additional_headers" + if "extra_headers" in params: + return "extra_headers" + # Neither present (unexpected) — prefer the newer convention. + return "additional_headers" + + # ─── Typed messages ────────────────────────────────────────────────── @@ -172,12 +201,18 @@ class SensingClient: the ``async with`` in your own retry loop. Auto-reconnect logic is application-specific (e.g., "retry forever" for a long-running automation vs "fail fast" for a CLI tool that should exit). + + Auth: pass ``token=`` to send ``Authorization: Bearer `` on + the WS upgrade, for sensing-servers started with ``RUVIEW_API_TOKEN`` + set. If ``token`` is omitted it defaults to the ``RUVIEW_API_TOKEN`` + environment variable; when neither is set, no header is sent. """ def __init__( self, url: str, *, + token: Optional[str] = None, ping_interval: float = 20.0, ping_timeout: float = 20.0, max_size: int = 16 * 1024 * 1024, @@ -188,18 +223,28 @@ class SensingClient: "`pip install \"wifi-densepose[client]\"` to enable the client extras." ) self.url = url + # Bearer token for auth-enabled sensing-servers. Explicit + # constructor argument wins; otherwise fall back to the + # RUVIEW_API_TOKEN environment variable. An empty value (unset + # env, or "") means "no auth" — no Authorization header is sent. + self._token = token if token is not None else os.environ.get(TOKEN_ENV_VAR) self._ping_interval = ping_interval self._ping_timeout = ping_timeout self._max_size = max_size self._ws: Any = None # websockets.WebSocketClientProtocol — typed Any to avoid import cost async def __aenter__(self) -> "SensingClient": - self._ws = await websockets.connect( - self.url, + connect_kwargs: dict[str, Any] = dict( ping_interval=self._ping_interval, ping_timeout=self._ping_timeout, max_size=self._max_size, ) + if self._token: + # Python (unlike the browser UI) can set Authorization + # directly on the WS upgrade — no ticket workaround needed. + header_kwarg = _select_header_kwarg(websockets.connect) + connect_kwargs[header_kwarg] = {"Authorization": f"Bearer {self._token}"} + self._ws = await websockets.connect(self.url, **connect_kwargs) return self async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: