/
/
1"""Tests for the Yandex Music interactive setup flow (run_setup)."""
2
3from __future__ import annotations
4
5import asyncio
6import base64
7import time
8from collections.abc import Awaitable, Callable
9from contextlib import suppress
10from typing import Any
11from unittest import mock
12from urllib.parse import unquote
13
14import pytest
15from music_assistant_models.enums import ConfigEntryType, FlowStepType
16from ya_passport_auth import Credentials, DeviceCodeSession, QrSession, SecretStr
17from ya_passport_auth.exceptions import DeviceCodeTimeoutError, QRTimeoutError
18
19from music_assistant.models.setup_flow import SetupFlowContext, SetupSession, StepExpiredError
20from music_assistant.providers.yandex_music import setup_flow as ym_flow
21from music_assistant.providers.yandex_music.constants import (
22 CONF_REFRESH_TOKEN,
23 CONF_REMEMBER_SESSION,
24 CONF_TOKEN,
25 CONF_X_TOKEN,
26)
27
28
29class _FakeClient:
30 """Canned PassportClient that confirms a QR/device login (optionally after one expiry)."""
31
32 def __init__(
33 self,
34 creds: Credentials,
35 *,
36 qr_fail_first: bool = False,
37 device_fail_first: bool = False,
38 ) -> None:
39 self._creds = creds
40 self.qr_starts = 0
41 self._qr_polls = 0
42 self._qr_fail_first = qr_fail_first
43 self.device_starts = 0
44 self._device_polls = 0
45 self._device_fail_first = device_fail_first
46
47 async def start_qr_login(self) -> QrSession:
48 self.qr_starts += 1
49 return QrSession(track_id="t", csrf_token="c", qr_url="https://passport.yandex.ru/qr/abc")
50
51 async def poll_qr_until_confirmed(self, _qr: QrSession, **_kwargs: Any) -> Credentials:
52 self._qr_polls += 1
53 if self._qr_fail_first and self._qr_polls == 1:
54 raise QRTimeoutError("expired")
55 return self._creds
56
57 async def start_device_login(self, **_kwargs: Any) -> DeviceCodeSession:
58 self.device_starts += 1
59 return DeviceCodeSession(
60 device_code=SecretStr("dc"),
61 user_code="ABCD-1234",
62 verification_url="https://ya.ru/device",
63 expires_in=300,
64 interval=5,
65 )
66
67 async def poll_device_until_confirmed(
68 self, _session: DeviceCodeSession, **_kwargs: Any
69 ) -> Credentials:
70 self._device_polls += 1
71 if self._device_fail_first and self._device_polls == 1:
72 raise DeviceCodeTimeoutError("expired")
73 return self._creds
74
75
76class _HangingClient(_FakeClient):
77 """Passport client whose confirmation polls never finish on their own."""
78
79 async def poll_qr_until_confirmed(self, _qr: QrSession, **_kwargs: Any) -> Credentials:
80 await asyncio.Event().wait()
81 raise AssertionError("unreachable")
82
83 async def poll_device_until_confirmed(
84 self, _session: DeviceCodeSession, **_kwargs: Any
85 ) -> Credentials:
86 await asyncio.Event().wait()
87 raise AssertionError("unreachable")
88
89
90def _async_cm(client: _FakeClient) -> mock.MagicMock:
91 """Wrap a fake client as the async context manager PassportClient.create returns."""
92 ctx = mock.MagicMock()
93 ctx.__aenter__ = mock.AsyncMock(return_value=client)
94 ctx.__aexit__ = mock.AsyncMock(return_value=False)
95 return ctx
96
97
98def _make_session(finish_handler: Any) -> tuple[SetupSession, mock.Mock]:
99 """Build a real SetupSession backed by a Mock mass for driving run_setup directly."""
100 mass = mock.Mock()
101 context = SetupFlowContext(kind="setup", reason="user", domain="yandex_music")
102 return SetupSession(mass, "flow-test", context, finish_handler), mass
103
104
105def _published_steps(mass: mock.Mock) -> list[Any]:
106 """Return the flow steps pushed through mass.signal_event, in order."""
107 return [call.kwargs["data"] for call in mass.signal_event.call_args_list]
108
109
110async def _wait_for(predicate: Any, timeout: float = 5.0) -> Any:
111 """Wait until the predicate returns truthy (or fail the test)."""
112 deadline = time.monotonic() + timeout
113 while time.monotonic() < deadline:
114 if result := predicate():
115 return result
116 await asyncio.sleep(0.01)
117 raise AssertionError("condition not met within timeout")
118
119
120async def _drive(session: SetupSession, submit: dict[str, Any]) -> None:
121 """Wait for the user form, submit the given values, then wait for finish."""
122 await _wait_for(lambda: session.current_step and session.current_step.type == FlowStepType.FORM)
123 session.handle_submit(submit)
124 await _wait_for(lambda: session.finished)
125
126
127async def _assert_login_has_hard_timeout(
128 login: Callable[[SetupSession], Awaitable[Credentials]],
129) -> None:
130 """Assert an abandoned login expires even while Yandex polling remains pending."""
131 creds = Credentials(x_token=SecretStr("XT"), music_token=SecretStr("MT"))
132 client = _HangingClient(creds)
133 session = mock.Mock(spec=SetupSession)
134
135 async def progress_until(awaitable: Awaitable[Credentials], **_kwargs: Any) -> Credentials:
136 return await awaitable
137
138 session.progress_until = mock.AsyncMock(side_effect=progress_until)
139 with (
140 mock.patch.object(ym_flow, "_AUTH_FLOW_TIMEOUT_SECONDS", 0.01),
141 mock.patch.object(ym_flow, "PassportClient") as passport_client,
142 ):
143 passport_client.create.return_value = _async_cm(client)
144 with pytest.raises(StepExpiredError):
145 await asyncio.wait_for(login(session), timeout=0.2)
146
147
148def test_qr_image_has_opaque_white_quiet_zone() -> None:
149 """The QR remains high-contrast against Music Assistant's dark theme."""
150 image = ym_flow._qr_image("https://passport.yandex.ru/qr/test")
151 svg = unquote(image.split(",", 1)[1])
152
153 assert "<path fill='#fff' d='M0 0h37v37h-37z'/>" in svg
154 assert "<path class='qrline' stroke='#000' d='M4 4.5" in svg
155
156
157async def test_qr_login_has_hard_timeout() -> None:
158 """An abandoned QR login cannot refresh codes forever."""
159 await _assert_login_has_hard_timeout(ym_flow._qr_login)
160
161
162async def test_device_login_has_hard_timeout() -> None:
163 """An abandoned Device Code login cannot refresh codes forever."""
164 await _assert_login_has_hard_timeout(ym_flow._device_login)
165
166
167async def test_device_countdown_respects_hard_timeout() -> None:
168 """The displayed Device Code lifetime cannot exceed the whole flow lifetime."""
169 creds = Credentials(x_token=SecretStr("XT"), music_token=SecretStr("MT"))
170 client = _FakeClient(creds)
171 session = mock.Mock(spec=SetupSession)
172 shown_expiry: float | None = None
173
174 async def progress_until(
175 awaitable: Awaitable[Credentials], *, expires_in: float, **_kwargs: Any
176 ) -> Credentials:
177 nonlocal shown_expiry
178 shown_expiry = expires_in
179 return await awaitable
180
181 session.progress_until = mock.AsyncMock(side_effect=progress_until)
182 with (
183 mock.patch.object(ym_flow, "_AUTH_FLOW_TIMEOUT_SECONDS", 10),
184 mock.patch.object(ym_flow, "PassportClient") as passport_client,
185 ):
186 passport_client.create.return_value = _async_cm(client)
187 await ym_flow._device_login(session)
188
189 assert shown_expiry is not None
190 assert 0 < shown_expiry <= 10
191
192
193def test_device_image_makes_verification_address_prominent() -> None:
194 """The non-clickable fallback clearly tells users where to enter the code."""
195 image = ym_flow._device_image("ABCD-1234", "https://ya.ru/device")
196 svg = base64.b64decode(image.split(",", 1)[1]).decode("utf-8")
197
198 assert "Open this address in a browser" in svg
199 assert ">ya.ru/device</text>" in svg
200 assert "https://ya.ru/device" not in svg
201 address = svg.split(">ya.ru/device</text>", 1)[0].rsplit("<text", 1)[1]
202 assert 'font-size="24"' in address
203
204
205async def test_manual_token_is_only_shown_after_selecting_its_method() -> None:
206 """QR and Device Code never render the secure token input on their method form."""
207
208 async def finish(_s: SetupSession, _values: dict[str, Any]) -> dict[str, str]:
209 raise AssertionError("form inspection must not finish the flow")
210
211 session, _mass = _make_session(finish)
212 task = asyncio.create_task(ym_flow.run_setup(session))
213 try:
214 form = await _wait_for(
215 lambda: (
216 session.current_step
217 if session.current_step and session.current_step.type == FlowStepType.FORM
218 else None
219 )
220 )
221 method_entry = next(entry for entry in form.entries if entry.key == ym_flow.CONF_METHOD)
222 assert {option.value for option in method_entry.options} >= {"qr", "device", "token"}
223 assert method_entry.default_value == ym_flow.METHOD_QR
224 assert CONF_TOKEN not in {entry.key for entry in form.entries}
225
226 session.handle_submit({ym_flow.CONF_METHOD: "token", CONF_REMEMBER_SESSION: True})
227 token_form = await _wait_for(
228 lambda: (
229 session.current_step
230 if session.current_step
231 and session.current_step.type == FlowStepType.FORM
232 and session.current_step.step_id == "token_login"
233 else None
234 ),
235 timeout=0.5,
236 )
237 token_entry = next(entry for entry in token_form.entries if entry.key == CONF_TOKEN)
238 assert token_entry.type == ConfigEntryType.SECURE_STRING
239 assert token_entry.required is True
240 assert {entry.key for entry in token_form.entries} == {CONF_TOKEN}
241 finally:
242 task.cancel()
243 with suppress(BaseException):
244 await task
245
246
247async def test_manual_token_login_persists_only_submitted_token() -> None:
248 """Manual login bypasses Passport and clears credentials from a previous session."""
249 creds = Credentials(x_token=SecretStr("unused-XT"), music_token=SecretStr("unused-MT"))
250 client = _FakeClient(creds)
251 collected: dict[str, Any] = {}
252
253 async def finish(_s: SetupSession, values: dict[str, Any]) -> dict[str, str]:
254 collected.update(values)
255 return {"instance_id": "yandex_music--1"}
256
257 session, _mass = _make_session(finish)
258 with mock.patch.object(ym_flow, "PassportClient") as passport_client:
259 passport_client.create.return_value = _async_cm(client)
260 task = asyncio.create_task(ym_flow.run_setup(session))
261 await _wait_for(lambda: session.current_step and session.current_step.step_id == "user")
262 session.handle_submit({ym_flow.CONF_METHOD: "token", CONF_REMEMBER_SESSION: True})
263 await _wait_for(
264 lambda: session.current_step and session.current_step.step_id == "token_login",
265 timeout=0.5,
266 )
267 session.handle_submit({CONF_TOKEN: "manual-token"})
268 await _wait_for(lambda: session.finished)
269 await task
270
271 assert collected == {
272 CONF_TOKEN: "manual-token",
273 CONF_X_TOKEN: None,
274 CONF_REFRESH_TOKEN: None,
275 }
276 passport_client.create.assert_not_called()
277
278
279async def test_manual_token_login_rejects_empty_token() -> None:
280 """Selecting manual login without a token re-renders the form with a field error."""
281
282 async def finish(_s: SetupSession, _values: dict[str, Any]) -> dict[str, str]:
283 raise AssertionError("an empty manual token must not finish the flow")
284
285 session, _mass = _make_session(finish)
286 task = asyncio.create_task(ym_flow.run_setup(session))
287 try:
288 await _wait_for(lambda: session.current_step and session.current_step.step_id == "user")
289 session.handle_submit({ym_flow.CONF_METHOD: "token", CONF_REMEMBER_SESSION: True})
290 await _wait_for(
291 lambda: session.current_step and session.current_step.step_id == "token_login",
292 timeout=0.5,
293 )
294 session.handle_submit({CONF_TOKEN: ""})
295 await _wait_for(
296 lambda: (
297 session.current_step and session.current_step.errors.get(CONF_TOKEN) == "required"
298 ),
299 timeout=0.5,
300 )
301 assert session.current_step is not None
302 assert session.current_step.errors == {CONF_TOKEN: "required"}
303 assert not task.done()
304 finally:
305 task.cancel()
306 with suppress(BaseException):
307 await task
308
309
310async def test_device_login_remember_persists_full_triple() -> None:
311 """Device login with remember on persists music + x + refresh tokens."""
312 creds = Credentials(
313 x_token=SecretStr("XT"),
314 music_token=SecretStr("MT"),
315 refresh_token=SecretStr("RT"),
316 display_login="alice",
317 uid=1,
318 )
319 collected: dict[str, Any] = {}
320
321 async def finish(_s: SetupSession, values: dict[str, Any]) -> dict[str, str]:
322 collected.update(values)
323 return {"instance_id": "yandex_music--1"}
324
325 session, mass = _make_session(finish)
326 client = _FakeClient(creds)
327 with mock.patch.object(ym_flow, "PassportClient") as pc:
328 pc.create.return_value = _async_cm(client)
329 task = asyncio.create_task(ym_flow.run_setup(session))
330 form = await _wait_for(
331 lambda: (
332 session.current_step
333 if session.current_step and session.current_step.type == FlowStepType.FORM
334 else None
335 )
336 )
337 method_entry = next(entry for entry in form.entries if entry.key == ym_flow.CONF_METHOD)
338 assert method_entry.default_value == ym_flow.METHOD_QR
339 await _drive(
340 session, {ym_flow.CONF_METHOD: ym_flow.METHOD_DEVICE, CONF_REMEMBER_SESSION: True}
341 )
342 await task
343
344 assert collected == {CONF_TOKEN: "MT", CONF_X_TOKEN: "XT", CONF_REFRESH_TOKEN: "RT"}
345 progress = [s for s in _published_steps(mass) if s.type == FlowStepType.PROGRESS]
346 assert progress
347 assert progress[0].step_id == "device_login"
348 assert progress[0].image is not None
349 assert progress[0].image.startswith("data:image/svg+xml")
350
351
352async def test_qr_login_without_remember_stores_music_token_only() -> None:
353 """QR login with remember off stores only the music token (x/refresh cleared)."""
354 creds = Credentials(x_token=SecretStr("XT"), music_token=SecretStr("MT"), display_login="bob")
355 collected: dict[str, Any] = {}
356
357 async def finish(_s: SetupSession, values: dict[str, Any]) -> dict[str, str]:
358 collected.update(values)
359 return {"instance_id": "yandex_music--1"}
360
361 session, mass = _make_session(finish)
362 client = _FakeClient(creds)
363 with mock.patch.object(ym_flow, "PassportClient") as pc:
364 pc.create.return_value = _async_cm(client)
365 task = asyncio.create_task(ym_flow.run_setup(session))
366 await _drive(
367 session, {ym_flow.CONF_METHOD: ym_flow.METHOD_QR, CONF_REMEMBER_SESSION: False}
368 )
369 await task
370
371 assert collected == {CONF_TOKEN: "MT", CONF_X_TOKEN: None, CONF_REFRESH_TOKEN: None}
372 scan_steps = [s for s in _published_steps(mass) if s.step_id == "scan_qr"]
373 assert scan_steps
374 assert all(s.image and s.image.startswith("data:image/svg+xml") for s in scan_steps)
375
376
377async def test_qr_login_refreshes_expired_code() -> None:
378 """An expired QR code is minted afresh and the login still completes."""
379 creds = Credentials(x_token=SecretStr("XT"), music_token=SecretStr("MT"))
380 collected: dict[str, Any] = {}
381
382 async def finish(_s: SetupSession, values: dict[str, Any]) -> dict[str, str]:
383 collected.update(values)
384 return {"instance_id": "yandex_music--1"}
385
386 session, _mass = _make_session(finish)
387 client = _FakeClient(creds, qr_fail_first=True)
388 with mock.patch.object(ym_flow, "PassportClient") as pc:
389 pc.create.return_value = _async_cm(client)
390 task = asyncio.create_task(ym_flow.run_setup(session))
391 await _drive(session, {ym_flow.CONF_METHOD: ym_flow.METHOD_QR, CONF_REMEMBER_SESSION: True})
392 await task
393
394 assert collected[CONF_TOKEN] == "MT"
395 # the expired code triggered a second start_qr_login (refresh loop)
396 assert client.qr_starts == 2
397
398
399async def test_device_login_refreshes_expired_code() -> None:
400 """An expired Device Code is minted afresh and the login still completes."""
401 creds = Credentials(x_token=SecretStr("XT"), music_token=SecretStr("MT"))
402 collected: dict[str, Any] = {}
403
404 async def finish(_s: SetupSession, values: dict[str, Any]) -> dict[str, str]:
405 collected.update(values)
406 return {"instance_id": "yandex_music--1"}
407
408 session, _mass = _make_session(finish)
409 client = _FakeClient(creds, device_fail_first=True)
410 with mock.patch.object(ym_flow, "PassportClient") as passport_client:
411 passport_client.create.return_value = _async_cm(client)
412 task = asyncio.create_task(ym_flow.run_setup(session))
413 await _drive(
414 session, {ym_flow.CONF_METHOD: ym_flow.METHOD_DEVICE, CONF_REMEMBER_SESSION: True}
415 )
416 await task
417
418 assert collected[CONF_TOKEN] == "MT"
419 assert client.device_starts == 2
420