Files
teng-lin--notebooklm-py/tests/unit/test_artifacts_polling_retries.py
wehub-resource-sync 09e9f3545f
Test / Code Quality (push) Has been cancelled
Test / Test (macos-latest, Python 3.10) (push) Has been cancelled
Test / Test (macos-latest, Python 3.11) (push) Has been cancelled
Test / Test (macos-latest, Python 3.12) (push) Has been cancelled
Test / Test (macos-latest, Python 3.13) (push) Has been cancelled
Test / Test (macos-latest, Python 3.14) (push) Has been cancelled
Test / Test (ubuntu-latest, Python 3.10) (push) Has been cancelled
Test / Test (ubuntu-latest, Python 3.11) (push) Has been cancelled
Test / Test (ubuntu-latest, Python 3.12) (push) Has been cancelled
Test / Test (ubuntu-latest, Python 3.13) (push) Has been cancelled
Test / Test (ubuntu-latest, Python 3.14) (push) Has been cancelled
Test / Test (windows-latest, Python 3.10) (push) Has been cancelled
Test / Test (windows-latest, Python 3.11) (push) Has been cancelled
Test / Test (windows-latest, Python 3.12) (push) Has been cancelled
Test / Test (windows-latest, Python 3.13) (push) Has been cancelled
Test / Test (windows-latest, Python 3.14) (push) Has been cancelled
CodeQL / Analyze (push) Has been cancelled
dependency-audit / pip-audit (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 13:30:13 +08:00

553 lines
18 KiB
Python

import asyncio
from contextlib import asynccontextmanager
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from notebooklm._artifact.polling import ArtifactPollingService
from notebooklm._artifacts import ArtifactsAPI, GenerationStatus
from notebooklm._polling_registry import PollRegistry
from notebooklm.exceptions import ArtifactPendingTimeoutError
from notebooklm.rpc import AuthError, NetworkError, RPCTimeoutError
class _FakeTransportProvider:
# ``ArtifactPollingService.wait_for_completion`` calls
# the injected loop guard. ``None`` is the documented
# silent-no-op value for the affinity helper, so this stub stays correct
# without binding to a real loop.
bound_loop = None
def assert_bound_loop(self) -> None:
return None
def __init__(
self,
*,
token: object | None = None,
begin_error: BaseException | None = None,
yield_before_begin_error: bool = False,
begin_release: asyncio.Event | None = None,
finish_release: asyncio.Event | None = None,
) -> None:
self.poll_registry = PollRegistry()
self.token = object() if token is None else token
self.begin_error = begin_error
self.yield_before_begin_error = yield_before_begin_error
self.begin_release = begin_release
self.finish_release = finish_release
self.begin_tasks: list[asyncio.Task[object]] = []
self.begin_labels: list[str] = []
self.begin_task_done_states: list[bool] = []
self.begin_started = asyncio.Event()
self.finish_tokens: list[object] = []
self.finish_started = asyncio.Event()
self.finish_finished = asyncio.Event()
async def rpc_call(self, *args, **kwargs):
raise AssertionError("unexpected rpc_call")
async def transport_post(self, *args, **kwargs):
raise AssertionError("unexpected transport_post")
async def next_reqid(self, step: int = 100000) -> int:
return step
def operation_scope(self, log_label: str):
provider = self
class _Scope:
async def __aenter__(self) -> None:
await provider._enter_scope(log_label)
return None
async def __aexit__(self, exc_type, exc, tb) -> None:
await provider._exit_scope()
return None
return _Scope()
async def _enter_scope(self, log_label: str) -> None:
task = asyncio.current_task()
assert task is not None
self.begin_tasks.append(task)
self.begin_labels.append(log_label)
self.begin_task_done_states.append(task.done())
self.begin_started.set()
if self.begin_release is not None:
await self.begin_release.wait()
if self.yield_before_begin_error:
await asyncio.sleep(0)
if self.begin_error is not None:
raise self.begin_error
async def _exit_scope(self) -> None:
self.finish_tokens.append(self.token)
self.finish_started.set()
if self.finish_release is not None:
await self.finish_release.wait()
self.finish_finished.set()
@pytest.fixture
def api():
from notebooklm._mind_map import NoteBackedMindMapService
from notebooklm._note_service import NoteService
core = _make_session_core()
mock_notebooks = MagicMock()
mock_notebooks.get_source_ids = AsyncMock(return_value=[])
return ArtifactsAPI(
rpc=core,
drain=core,
lifecycle=core,
notebooks=mock_notebooks,
mind_maps=MagicMock(spec=NoteBackedMindMapService),
note_service=MagicMock(spec=NoteService),
)
def _make_session_core() -> MagicMock:
core = MagicMock()
# Real registry backing so wait_for_completion can ``dict.get(key)``.
core.assert_bound_loop = MagicMock(return_value=None)
core.operation_scope = MagicMock(side_effect=lambda _label: _noop_operation_scope())
return core
@asynccontextmanager
async def _noop_operation_scope():
yield None
@pytest.mark.asyncio
async def test_wait_for_completion_retry_success(api):
# Mock poll_status to fail twice then succeed
status_ready = GenerationStatus(task_id="task1", status="completed")
api.poll_status = AsyncMock()
api.poll_status.side_effect = [
NetworkError("transient net"),
RPCTimeoutError("transient timeout"),
status_ready,
]
with patch("asyncio.sleep", AsyncMock()) as mock_sleep:
# Also need to patch asyncio.get_running_loop().time() to avoid timeout
# but here we just test the retry logic
result = await api.wait_for_completion("nb1", "task1", timeout=60.0)
assert result == status_ready
assert api.poll_status.call_count == 3
assert mock_sleep.call_count == 2
# Backoff: 2^1=2, 2^2=4
mock_sleep.assert_any_call(2.0)
mock_sleep.assert_any_call(4.0)
@pytest.mark.asyncio
async def test_polling_service_clamps_transient_retry_sleep_to_remaining_timeout() -> None:
provider = _FakeTransportProvider()
clock = 0.0
sleeps: list[float] = []
def monotonic() -> float:
return clock
async def sleep(seconds: float) -> None:
nonlocal clock
sleeps.append(seconds)
clock += seconds
service = ArtifactPollingService(
loop_guard=provider,
op_scope=provider,
poll_registry=provider.poll_registry,
sleep=sleep,
monotonic=monotonic,
)
poll_status = AsyncMock(side_effect=NetworkError("transient net"))
with pytest.raises(ArtifactPendingTimeoutError) as exc_info:
await service.wait_for_completion("nb1", "task1", timeout=1.0, poll_status=poll_status)
assert isinstance(exc_info.value.__cause__, NetworkError)
assert "transient net" in str(exc_info.value.__cause__)
assert poll_status.await_count == 1
assert sleeps == [1.0]
assert clock == 1.0
@pytest.mark.asyncio
async def test_polling_service_clamps_poll_interval_to_remaining_timeout() -> None:
provider = _FakeTransportProvider()
clock = 0.0
sleeps: list[float] = []
def monotonic() -> float:
return clock
async def sleep(seconds: float) -> None:
nonlocal clock
sleeps.append(seconds)
clock += seconds
service = ArtifactPollingService(
loop_guard=provider,
op_scope=provider,
poll_registry=provider.poll_registry,
sleep=sleep,
monotonic=monotonic,
)
poll_status = AsyncMock(return_value=GenerationStatus(task_id="task1", status="pending"))
with pytest.raises(ArtifactPendingTimeoutError):
await service.wait_for_completion(
"nb1",
"task1",
initial_interval=10.0,
timeout=1.0,
poll_status=poll_status,
)
assert poll_status.await_count == 2
assert sleeps == [1.0]
assert clock == 1.0
@pytest.mark.asyncio
async def test_wait_for_completion_retry_exhausted(api):
api.poll_status = AsyncMock()
api.poll_status.side_effect = NetworkError("persistent fail")
with patch("asyncio.sleep", AsyncMock()):
with pytest.raises(NetworkError, match="persistent fail"):
await api.wait_for_completion("nb1", "task1", timeout=60.0)
# Initial call + 3 retries = 4 total calls
assert api.poll_status.call_count == 4
@pytest.mark.asyncio
async def test_wait_for_completion_no_retry_on_auth_error(api):
api.poll_status = AsyncMock()
api.poll_status.side_effect = AuthError("auth fail")
with patch("asyncio.sleep", AsyncMock()) as mock_sleep:
with pytest.raises(AuthError, match="auth fail"):
await api.wait_for_completion("nb1", "task1", timeout=60.0)
assert api.poll_status.call_count == 1
assert mock_sleep.call_count == 0
@pytest.mark.asyncio
async def test_polling_service_operation_scope_wraps_spawned_poll_task() -> None:
token = object()
provider = _FakeTransportProvider(token=token)
service = ArtifactPollingService(
loop_guard=provider, op_scope=provider, poll_registry=provider.poll_registry
)
async def poll_status(notebook_id: str, task_id: str) -> GenerationStatus:
assert (notebook_id, task_id) == ("nb1", "task1")
return GenerationStatus(task_id=task_id, status="completed")
result = await service.wait_for_completion(
"nb1",
"task1",
initial_interval=0.0,
max_interval=0.0,
timeout=1.0,
poll_status=poll_status,
)
assert result.status == "completed"
assert len(provider.begin_tasks) == 1
poll_task = provider.begin_tasks[0]
assert isinstance(poll_task, asyncio.Task)
assert poll_task.done()
assert poll_task.get_name() == "artifact-poll-nb1-task1"
assert provider.begin_task_done_states == [False]
assert provider.begin_labels == ["artifact wait task1"]
await asyncio.wait_for(provider.finish_finished.wait(), timeout=1.0)
assert provider.finish_tokens == [token]
@pytest.mark.asyncio
async def test_polling_service_registers_pending_before_transport_begin_completes() -> None:
begin_release = asyncio.Event()
provider = _FakeTransportProvider(begin_release=begin_release)
service = ArtifactPollingService(
loop_guard=provider, op_scope=provider, poll_registry=provider.poll_registry
)
poll_call_count = 0
async def poll_status(notebook_id: str, task_id: str) -> GenerationStatus:
nonlocal poll_call_count
poll_call_count += 1
return GenerationStatus(task_id=task_id, status="completed")
leader = asyncio.create_task(
service.wait_for_completion(
"nb1",
"task1",
initial_interval=0.0,
max_interval=0.0,
timeout=1.0,
poll_status=poll_status,
)
)
follower: asyncio.Task[GenerationStatus] | None = None
key = ("nb1", "task1")
try:
await asyncio.wait_for(provider.begin_started.wait(), timeout=1.0)
assert provider.poll_registry.get(key) is not None
follower = asyncio.create_task(
service.wait_for_completion(
"nb1",
"task1",
initial_interval=0.0,
max_interval=0.0,
timeout=1.0,
poll_status=poll_status,
)
)
await asyncio.sleep(0)
begin_release.set()
leader_result = await asyncio.wait_for(leader, timeout=1.0)
follower_result = await asyncio.wait_for(follower, timeout=1.0)
assert leader_result.status == "completed"
assert follower_result.status == "completed"
assert poll_call_count == 1
await asyncio.wait_for(provider.finish_finished.wait(), timeout=1.0)
assert provider.poll_registry.get(key) is None
finally:
begin_release.set()
cleanup_tasks = [
task for task in (leader, follower) if task is not None and not task.done()
]
for task in cleanup_tasks:
task.cancel()
if cleanup_tasks:
await asyncio.gather(*cleanup_tasks, return_exceptions=True)
@pytest.mark.asyncio
async def test_polling_service_resolves_wait_before_slow_transport_finish() -> None:
token = object()
finish_release = asyncio.Event()
provider = _FakeTransportProvider(token=token, finish_release=finish_release)
service = ArtifactPollingService(
loop_guard=provider, op_scope=provider, poll_registry=provider.poll_registry
)
async def poll_status(notebook_id: str, task_id: str) -> GenerationStatus:
return GenerationStatus(task_id=task_id, status="completed")
waiter = asyncio.create_task(
service.wait_for_completion(
"nb1",
"task1",
initial_interval=0.0,
max_interval=0.0,
timeout=1.0,
poll_status=poll_status,
)
)
try:
await asyncio.wait_for(provider.finish_started.wait(), timeout=1.0)
assert not waiter.done()
assert provider.finish_tokens == [token]
finish_release.set()
result = await asyncio.wait_for(waiter, timeout=1.0)
finally:
finish_release.set()
if not waiter.done():
waiter.cancel()
await asyncio.gather(waiter, return_exceptions=True)
assert result.status == "completed"
assert provider.finish_finished.is_set()
@pytest.mark.asyncio
async def test_polling_service_drain_waits_for_bookkeeping_without_active_polls() -> None:
token = object()
finish_release = asyncio.Event()
provider = _FakeTransportProvider(token=token, finish_release=finish_release)
service = ArtifactPollingService(
loop_guard=provider, op_scope=provider, poll_registry=provider.poll_registry
)
async def poll_status(notebook_id: str, task_id: str) -> GenerationStatus:
return GenerationStatus(task_id=task_id, status="completed")
waiter = asyncio.create_task(
service.wait_for_completion(
"nb1",
"task1",
initial_interval=0.0,
max_interval=0.0,
timeout=1.0,
poll_status=poll_status,
)
)
await asyncio.wait_for(provider.finish_started.wait(), timeout=1.0)
assert not waiter.done()
drain_task = asyncio.create_task(service.drain())
try:
await asyncio.sleep(0)
assert not drain_task.done()
finish_release.set()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(waiter, timeout=1.0)
await asyncio.wait_for(drain_task, timeout=1.0)
finally:
finish_release.set()
if not waiter.done():
waiter.cancel()
await asyncio.gather(waiter, return_exceptions=True)
if not drain_task.done():
drain_task.cancel()
await asyncio.gather(drain_task, return_exceptions=True)
assert provider.finish_tokens == [token]
@pytest.mark.asyncio
async def test_polling_service_finishes_transport_token_once_after_poll_failure() -> None:
token = object()
provider = _FakeTransportProvider(token=token)
service = ArtifactPollingService(
loop_guard=provider, op_scope=provider, poll_registry=provider.poll_registry
)
async def poll_status(notebook_id: str, task_id: str) -> GenerationStatus:
raise ValueError(f"poll failed: {notebook_id}/{task_id}")
with pytest.raises(ValueError, match="poll failed: nb1/task1"):
await service.wait_for_completion(
"nb1",
"task1",
initial_interval=0.0,
max_interval=0.0,
timeout=1.0,
poll_status=poll_status,
)
assert len(provider.begin_tasks) == 1
assert provider.begin_tasks[0].done()
await asyncio.wait_for(provider.finish_finished.wait(), timeout=1.0)
assert provider.finish_tokens == [token]
@pytest.mark.asyncio
async def test_polling_service_cancels_and_drains_spawned_poll_task_if_begin_fails() -> None:
begin_error = RuntimeError("draining")
provider = _FakeTransportProvider(
begin_error=begin_error,
yield_before_begin_error=True,
)
service = ArtifactPollingService(
loop_guard=provider, op_scope=provider, poll_registry=provider.poll_registry
)
async def poll_status(notebook_id: str, task_id: str) -> GenerationStatus:
raise AssertionError("poll should not start when operation admission fails")
with pytest.raises(RuntimeError, match="draining"):
await service.wait_for_completion(
"nb1",
"task1",
initial_interval=0.0,
max_interval=0.0,
timeout=1.0,
poll_status=poll_status,
)
assert len(provider.begin_tasks) == 1
assert provider.begin_tasks[0].done()
assert provider.poll_registry.get(("nb1", "task1")) is None
assert not provider.finish_started.is_set()
assert provider.finish_tokens == []
@pytest.mark.asyncio
async def test_wait_for_completion_follower_cancellation_does_not_cancel_leader_or_later_waiter():
from notebooklm._mind_map import NoteBackedMindMapService
from notebooklm._note_service import NoteService
core = _make_session_core()
api = ArtifactsAPI(
rpc=core,
drain=core,
lifecycle=core,
notebooks=MagicMock(),
mind_maps=MagicMock(spec=NoteBackedMindMapService),
note_service=MagicMock(spec=NoteService),
)
poll_started = asyncio.Event()
release_poll = asyncio.Event()
status_ready = GenerationStatus(task_id="task1", status="completed")
poll_call_count = 0
test_timeout = 1.0
async def poll_status(notebook_id: str, task_id: str) -> GenerationStatus:
nonlocal poll_call_count
assert (notebook_id, task_id) == ("nb1", "task1")
poll_call_count += 1
poll_started.set()
await release_poll.wait()
return status_ready
api.poll_status = AsyncMock(side_effect=poll_status)
leader = asyncio.create_task(api.wait_for_completion("nb1", "task1", timeout=60.0))
key = ("nb1", "task1")
later_waiter: asyncio.Task[GenerationStatus] | None = None
try:
await asyncio.wait_for(poll_started.wait(), timeout=test_timeout)
for _ in range(10):
if api._poll_registry.get(key) is not None:
break
await asyncio.sleep(0)
assert api._poll_registry.get(key) is not None
follower = asyncio.create_task(api.wait_for_completion("nb1", "task1", timeout=60.0))
await asyncio.sleep(0)
follower.cancel()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(follower, timeout=test_timeout)
assert not leader.done()
assert api._poll_registry.get(key) is not None
assert poll_call_count == 1
later_waiter = asyncio.create_task(api.wait_for_completion("nb1", "task1", timeout=60.0))
await asyncio.sleep(0)
release_poll.set()
assert await asyncio.wait_for(leader, timeout=test_timeout) == status_ready
assert await asyncio.wait_for(later_waiter, timeout=test_timeout) == status_ready
assert poll_call_count == 1
assert api._poll_registry.get(key) is None
finally:
release_poll.set()
cleanup_tasks = []
for task in (leader, later_waiter):
if task is not None and not task.done():
task.cancel()
cleanup_tasks.append(task)
if cleanup_tasks:
await asyncio.gather(*cleanup_tasks, return_exceptions=True)