项目文件夹

文件
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

570 行
25 KiB
Python

"""Snapshot-invariant for the shared POST helper.
httpx merges the cookie jar into the outgoing ``httpx.Request`` synchronously
when ``Kernel.post`` opens the stream. ``RuntimeTransport.perform_authed_post``
materializes a ``RpcRequest`` from an auth snapshot before the middleware chain,
and the ``Kernel.post`` terminal refreshes that envelope if auth changed while
the request waited in the chain. Therefore, within a single terminal attempt
the entire ``(csrf, session_id, cookies)`` snapshot is atomic from a concurrent
coroutine standpoint: no other task can mutate state between the terminal
freshness check and the wire.
The POST lives in the shared authed transport path used by
``RuntimeTransport.perform_authed_post`` so chat can share the same transport pipeline.
The AST guard below follows the POST; the invariant still belongs at the
shared site.
The auth-snapshot lock hardened the invariant by:
- making ``AuthRefreshCoordinator.snapshot()`` ``async def`` and acquiring a dedicated
``_auth_snapshot_lock`` for the read, so the four scalar fields
(``csrf_token``/``session_id``/``authuser``/``account_email``) are
observed atomically with respect to ``refresh_auth``'s
write-block; and
- refactoring ``RpcExecutor.build_url()`` to consume the resulting
``AuthSnapshot`` rather than reading ``self.auth`` LIVE — that
prior live-read was the actual torn-read hazard, since it let a
refresh's write to ``self.auth.session_id`` slip into the URL between
snapshot capture and request build.
This file *locks* the invariant in four ways:
1. ``test_kernel_post_terminal_has_no_await_before_post_per_attempt`` —
static AST guard against an ``await`` inside the terminal's ``try`` body
before ``Kernel.post``. The freshness check lives before the try block
and is guarded separately so the lock acquisition itself is not a
regression.
2. ``test_build_url_does_not_read_self_auth`` — static AST guard
against any ``self.auth.<field>`` attribute access in
``RpcExecutor.build_url``. The method MUST consume only its
``snapshot: AuthSnapshot`` parameter; reverting to ``self.auth``
would silently un-do the atomicity fix.
3. ``test_concurrent_refresh_does_not_corrupt_inflight_rpc_request`` —
runtime self-consistency. Drives concurrent ``refresh_auth`` against
an in-flight ``rpc_call`` (both orderings) and asserts the captured
``httpx.Request`` is never observed with mixed-generation (csrf,
session_id, cookies) state.
4. ``test_auth_refresh_rebuild_has_no_await_after_snapshot_capture`` —
static guard on the auth-refresh retry rebuild: once the post-refresh
snapshot has been captured, pairing ``context["auth_snapshot"]`` with
the rebuilt envelope must remain synchronous. The terminal still owns
the final before-wire freshness check because inner middlewares may
await after this rebuild.
"""
from __future__ import annotations
import ast
import asyncio
import contextlib
import inspect
import json
import textwrap
import httpx
import pytest
from notebooklm._middleware.auth_refresh import AuthRefreshMiddleware
from notebooklm._rpc_executor import RpcExecutor
from notebooklm._runtime.auth import AuthRefreshCoordinator
from notebooklm._runtime.transport import RuntimeTransport
from notebooklm.auth import AuthTokens
from notebooklm.rpc import RPCMethod
from tests._fixtures.kernel_test_helpers import install_http_client_for_test
from tests._helpers.client_factory import build_client_shell_for_tests
# Test-side deadline for any single asyncio.Event in the race scaffolding.
# Generous enough not to flake on slow CI, tight enough that a regression
# (e.g., POST never reached the transport) fails fast instead of hanging.
EVENT_TIMEOUT_S = 5.0
def test_kernel_post_terminal_has_no_await_before_post_per_attempt():
"""No ``await`` may sit inside the terminal ``try`` before ``Kernel.post``.
The terminal first refreshes the envelope if auth changed while the
request waited behind middlewares, then enters a small ``try`` whose first
await is the actual ``Kernel.post`` send. An await in that try prologue
would let a concurrent refresh change the cookie jar between request
rebuild and wire send.
The terminal body lives on :meth:`RuntimeTransport.terminal`; the
chain leaf (:meth:`MiddlewareChainHost._authed_post_chain_terminal`)
forwards to it. The AST guard inspects the collaborator method
that carries the actual body.
"""
src = textwrap.dedent(inspect.getsource(RuntimeTransport.terminal))
tree = ast.parse(src)
func = next(n for n in ast.walk(tree) if isinstance(n, ast.AsyncFunctionDef))
# Locate the ``try`` block guarding the POST. Post-PR-12.9 the leaf
# has no ``while`` retry loop and no ``async with`` semaphore wrap;
# the try sits at the top of the function body (the semaphore is
# held by ``SemaphoreMiddleware`` higher up the chain).
def _find_first_try(parent: ast.AST) -> ast.Try | None:
for child in ast.iter_child_nodes(parent):
if isinstance(child, ast.Try):
return child
if isinstance(child, ast.AsyncWith | ast.With):
found = _find_first_try(child)
if found is not None:
return found
if isinstance(child, ast.While):
# Tolerate a re-introduced while (e.g. if a future PR
# adds a retry loop back to the leaf) — walk into it.
found = _find_first_try(child)
if found is not None:
return found
return None
found_try = _find_first_try(func)
assert found_try is not None, (
"Could not locate the ``try:`` block guarding the POST in "
"RuntimeTransport.terminal. Update this guard to match."
)
def is_post_await(node):
"""Match the single per-iteration POST await.
Accepts either historical shape:
- ``await client.post(...)`` (pre-streaming),
- ``await stream_post_with_size_cap(...)`` (the helper performs the
streaming POST internally, so it's the same conceptual POST site for
the purposes of this concurrency invariant),
- ``await self._kernel.post(...)`` (current terminal shape).
"""
if not isinstance(node, ast.Await):
return False
call = node.value
if not isinstance(call, ast.Call):
return False
func = call.func
if isinstance(func, ast.Attribute) and func.attr == "post":
return True
return isinstance(func, ast.Name) and func.id == "stream_post_with_size_cap"
def _walk_outer(parent):
"""Yield nodes lexically inside ``parent`` itself (skip nested defs).
``ast.walk`` descends into nested ``FunctionDef`` / ``AsyncFunctionDef``
/ ``Lambda`` bodies — that would let a future helper coroutine
smuggle the matching ``await ...post(...)`` past this guard. We only
want statements at this lexical level. We DO descend into the loop's
own ``try`` / ``except`` / ``if`` blocks so awaits in retry-branch
bookkeeping (post-error handlers) remain visible to the guard.
"""
boundaries = (ast.FunctionDef, ast.AsyncFunctionDef, ast.Lambda, ast.ClassDef)
for child in ast.iter_child_nodes(parent):
if isinstance(child, boundaries):
continue
yield child
yield from _walk_outer(child)
# Walk only the ``try`` body — that's the critical prologue → POST
# window. Awaits in ``except`` handlers are by definition AFTER the
# POST and don't violate the invariant. ``self._snapshot()`` and
# ``build_request(snapshot)`` are synchronous assignments inside the
# ``async with semaphore:`` body before the try, so they're
# irrelevant to this guard.
try_node = found_try
# We only walk the try body, NOT its handlers.
try_body_nodes: list[ast.AST] = []
for stmt in try_node.body:
try_body_nodes.append(stmt)
try_body_nodes.extend(_walk_outer(stmt))
post_await_positions = [(n.lineno, n.col_offset) for n in try_body_nodes if is_post_await(n)]
post_await_position = min(post_await_positions, default=None)
assert post_await_position is not None, (
"Could not locate `await ...post(...)` in the try body of "
"RuntimeTransport.terminal. If the call site was refactored (e.g. to "
"``client.request(...)``), update this guard to match — the "
"invariant is 'no await between snapshot read and the POST per "
"iteration', not specifically the `.post` attribute."
)
earlier_awaits = [
n
for n in try_body_nodes
if isinstance(n, ast.Await) and (n.lineno, n.col_offset) < post_await_position
]
assert not earlier_awaits, (
f"RuntimeTransport.terminal gained an await before the per-attempt POST "
f"at {post_await_position}: "
f"{[(n.lineno, ast.dump(n)) for n in earlier_awaits]}. "
"This breaks the snapshot-invariant — auth state could be mutated "
"between the snapshot read and the actual send."
)
def test_terminal_freshness_check_has_no_await_after_materialization():
"""Freshness rebuild must not yield after materializing a new envelope.
The freshness-rebuild body lives on
:meth:`RuntimeTransport.refresh_request_for_current_auth`, so the AST guard
inspects the collaborator method that carries the actual body.
"""
src = textwrap.dedent(inspect.getsource(RuntimeTransport.refresh_request_for_current_auth))
tree = ast.parse(src)
func = next(n for n in ast.walk(tree) if isinstance(n, ast.AsyncFunctionDef))
materialize_positions = [
(n.lineno, n.col_offset)
for n in ast.walk(func)
if isinstance(n, ast.Call)
and isinstance(n.func, ast.Name)
and n.func.id == "materialize_rpc_request"
]
assert materialize_positions, "Could not locate materialize_rpc_request in freshness check"
materialize_position = min(materialize_positions)
later_awaits = [
n
for n in ast.walk(func)
if isinstance(n, ast.Await) and (n.lineno, n.col_offset) > materialize_position
]
assert later_awaits == [], (
"RuntimeTransport.refresh_request_for_current_auth must not await after "
"materialize_rpc_request; that would let auth/cookies move between "
"request rebuild and Kernel.post."
)
def test_auth_refresh_rebuild_has_no_await_after_snapshot_capture():
"""Auth-refresh retry rebuild pairs fresh snapshot and envelope atomically."""
src = textwrap.dedent(inspect.getsource(AuthRefreshMiddleware._rebuild_request_after_refresh))
tree = ast.parse(src)
func = next(n for n in ast.walk(tree) if isinstance(n, ast.AsyncFunctionDef))
snapshot_awaits = [
(n.lineno, n.col_offset)
for n in ast.walk(func)
if isinstance(n, ast.Await)
and isinstance(n.value, ast.Call)
and isinstance(n.value.func, ast.Attribute)
and n.value.func.attr == "_snapshot_provider"
]
assert snapshot_awaits, "Could not locate await self._snapshot_provider()"
snapshot_position = min(snapshot_awaits)
materialize_positions = [
(n.lineno, n.col_offset)
for n in ast.walk(func)
if isinstance(n, ast.Call)
and isinstance(n.func, ast.Name)
and n.func.id == "materialize_rpc_request"
]
assert materialize_positions, "Could not locate materialize_rpc_request in refresh rebuild"
assert snapshot_position < min(materialize_positions), (
"AuthRefreshMiddleware._rebuild_request_after_refresh must capture "
"the fresh snapshot before rebuilding the retry envelope."
)
later_awaits = [
n
for n in ast.walk(func)
if isinstance(n, ast.Await) and (n.lineno, n.col_offset) > snapshot_position
]
assert later_awaits == [], (
"AuthRefreshMiddleware._rebuild_request_after_refresh must not await "
"after capturing the fresh snapshot; context['auth_snapshot'] and the "
"rebuilt RpcRequest must stay paired until the terminal freshness check."
)
def test_build_url_does_not_read_self_auth():
"""``RpcExecutor.build_url`` must consume only its ``snapshot`` parameter.
pre-fix, ``RpcExecutor.build_url`` reached into ``self.auth``
on every call to read ``session_id``, ``authuser``, and
``account_email``. With ``AuthRefreshCoordinator.snapshot()`` and
``RpcExecutor.build_url()`` running on separate Python statements, a concurrent ``refresh_auth`` could
flip ``self.auth.session_id`` between snapshot capture and URL build
— producing a request whose URL was stamped with the *new*
generation while the body still carried the *old* CSRF.
The fix made ``_build_url`` accept ``snapshot: AuthSnapshot`` and
read every auth scalar off the snapshot. This guard asserts that
contract statically so a future "convenience" refactor (e.g.
"let's just read ``self.auth`` again, it's right there") can't
silently re-introduce the torn read.
Allowed reads inside ``_build_url``: ``snapshot.session_id``,
``snapshot.authuser``, ``snapshot.account_email``, anything not
rooted at ``self.auth``. Forbidden: any ``self.auth.<field>``
attribute access, regardless of which field.
"""
src = textwrap.dedent(inspect.getsource(RpcExecutor.build_url))
tree = ast.parse(src)
# ``_build_url`` is a sync method, not async.
func = next(n for n in ast.walk(tree) if isinstance(n, ast.FunctionDef))
forbidden: list[tuple[int, str]] = []
for node in ast.walk(func):
if not isinstance(node, ast.Attribute):
continue
# Looking for ``self.auth`` (Attribute whose .value is Name "self"
# and .attr is "auth"). That's the immediate parent of any
# ``self.auth.<field>`` read.
if isinstance(node.value, ast.Name) and node.value.id == "self" and node.attr == "auth":
forbidden.append((node.lineno, ast.dump(node)))
assert not forbidden, (
f"RpcExecutor.build_url reads self.auth — torn-read regression. "
f"Read every auth scalar off the ``snapshot`` parameter instead. "
f"Occurrences: {forbidden}"
)
def test_snapshot_acquires_auth_snapshot_lock():
"""``AuthRefreshCoordinator.snapshot`` must acquire ``_auth_snapshot_lock``.
The lock-acquisition contract belongs to the coordinator, so this guard
inspects the coordinator method directly.
The lock is the only thing that serializes the four-scalar snapshot
read with the matching two-scalar write in
``NotebookLMClient.refresh_auth``. Removing the ``async with`` block
here would re-open the torn-read window between
``host.auth.csrf_token`` and ``host.auth.session_id`` reads, even
though those two attribute reads are individually atomic at the
Python bytecode level.
This guard asserts that ``snapshot``'s body contains an
``async with`` whose context expression resolves to
``self.get_auth_snapshot_lock()`` (or, defensively, anything
referencing ``_auth_snapshot_lock`` so a maintainer who inlines the
lazy accessor doesn't trip the guard).
"""
src = textwrap.dedent(inspect.getsource(AuthRefreshCoordinator.snapshot))
tree = ast.parse(src)
func = next(n for n in ast.walk(tree) if isinstance(n, ast.AsyncFunctionDef))
has_lock_acquisition = False
for node in ast.walk(func):
if not isinstance(node, ast.AsyncWith):
continue
# Each ``async with X`` may chain multiple items; check each.
for item in node.items:
ctx = item.context_expr
# Match both call form ``self.get_auth_snapshot_lock()`` and
# direct attribute ``self._auth_snapshot_lock``.
if isinstance(ctx, ast.Call):
ctx = ctx.func
if isinstance(ctx, ast.Attribute) and "auth_snapshot_lock" in ctx.attr:
has_lock_acquisition = True
break
assert has_lock_acquisition, (
"AuthRefreshCoordinator.snapshot no longer acquires "
"_auth_snapshot_lock. Atomicity contract broken — the four-scalar "
"snapshot read is no longer atomic with the refresh-side write block "
"in NotebookLMClient.refresh_auth, exposing torn (csrf, sid) reads."
)
def test_update_auth_tokens_has_no_await_inside_mutation_block():
"""``AuthRefreshCoordinator.update_auth_tokens`` must not await mid-mutation.
The no-await invariant inside the csrf/session_id mutation block belongs to
the coordinator, so this guard inspects the coordinator method directly.
Lock acquisition may await; the mutation block itself may
not, because any yield inside it would let a snapshot observe a
torn (csrf, session_id) pair.
"""
src = textwrap.dedent(inspect.getsource(AuthRefreshCoordinator.update_auth_tokens))
tree = ast.parse(src)
func = next(n for n in ast.walk(tree) if isinstance(n, ast.AsyncFunctionDef))
mutation_try = next(
(
node
for node in ast.walk(func)
if isinstance(node, ast.Try)
and any(
isinstance(stmt, ast.Assign)
and any(
isinstance(target, ast.Attribute)
and target.attr in {"csrf_token", "session_id"}
for target in stmt.targets
)
for stmt in node.body
)
),
None,
)
assert mutation_try is not None, (
"Could not locate the guarded csrf/session_id mutation block in "
"AuthRefreshCoordinator.update_auth_tokens."
)
awaits = [node for node in ast.walk(mutation_try) if isinstance(node, ast.Await)]
assert awaits == [], (
"AuthRefreshCoordinator.update_auth_tokens must not await inside "
"the critical mutation block; doing so would let snapshots observe "
"torn auth tokens."
)
def _synthetic_rpc_response_text(rpc_id: str) -> str:
"""Build a minimal valid batchexecute response that decodes to []."""
inner = json.dumps([])
chunk = json.dumps([["wrb.fr", rpc_id, inner, None, None]])
return f")]}}'\n{len(chunk)}\n{chunk}\n"
@pytest.mark.parametrize("rpc_first", [True, False], ids=["rpc-first", "refresh-first"])
async def test_concurrent_refresh_does_not_corrupt_inflight_rpc_request(rpc_first):
"""Every outgoing RPC must carry a coherent (csrf, session_id, cookies) tuple.
On current code both parameterizations observe OLD/OLD/OLD: the RPC's
request is fully built (synchronously) while refresh is still suspended
in its GET, so all three values are captured from the pre-rotation state.
The assertion below catches the broken case where a future refactor
introduces a yield point between auth read and ``post()`` — letting
refresh complete in between would produce mixed generations.
"""
captured_post: list[dict] = []
rpc_send_entered = asyncio.Event()
let_rpc_send_complete = asyncio.Event()
get_entered = asyncio.Event()
let_get_complete = asyncio.Event()
rpc_method_id = RPCMethod.LIST_NOTEBOOKS.value
async def handler(request: httpx.Request) -> httpx.Response:
if request.method == "POST":
captured_post.append(
{
"url": str(request.url),
"cookie": request.headers.get("cookie", ""),
"body": bytes(request.content),
}
)
rpc_send_entered.set()
await let_rpc_send_complete.wait()
return httpx.Response(200, text=_synthetic_rpc_response_text(rpc_method_id))
get_entered.set()
await let_get_complete.wait()
body = '<script>"SNlM0e":"CSRF_NEW","FdrFJe":"SID_NEW"</script>'
return httpx.Response(
200,
text=body,
headers={"set-cookie": "SID=new_sid_cookie; Path=/; Domain=.google.com"},
)
transport = httpx.MockTransport(handler)
# Build the same auth scaffold the unit conftest's ``make_core`` produces
# (CSRF_OLD / SID_OLD / old_sid_cookie) so the OLD/NEW marker assertions
# below stay valid. Routed through :func:`build_client_shell_for_tests` so the
# refresh client shell and composed runtime share one auth snapshot provider.
auth = AuthTokens(
csrf_token="CSRF_OLD",
session_id="SID_OLD",
cookies={"SID": "old_sid_cookie"},
)
client = build_client_shell_for_tests(auth=auth, refresh_retry_delay=0.0)
core = client
await core.__aenter__()
try:
# Swap the auto-built http client for one that uses the test
# transport so we can observe the real ``httpx.Request`` (cookie
# merge, headers, body, URL). Mirrors the ``make_core`` post-open
# transport-install dance verbatim.
prior_cookies = core._collaborators.kernel.get_http_client().cookies
await core._collaborators.kernel.get_http_client().aclose()
install_http_client_for_test(
core._collaborators.kernel,
httpx.AsyncClient(
cookies=prior_cookies,
transport=transport,
timeout=httpx.Timeout(connect=1.0, read=5.0, write=5.0, pool=1.0),
),
)
# try/finally ensures the mock-transport handlers are unblocked even
# if a wait_for times out — otherwise pending tasks dangle in the
# event loop and the test hangs until pytest's own timeout fires.
rpc_task: asyncio.Task | None = None
refresh_task: asyncio.Task | None = None
try:
if rpc_first:
rpc_task = asyncio.create_task(
core._rpc_executor.rpc_call(RPCMethod.LIST_NOTEBOOKS, [])
)
await asyncio.wait_for(rpc_send_entered.wait(), EVENT_TIMEOUT_S)
refresh_task = asyncio.create_task(client.refresh_auth())
await asyncio.wait_for(get_entered.wait(), EVENT_TIMEOUT_S)
let_get_complete.set()
await asyncio.wait_for(refresh_task, EVENT_TIMEOUT_S)
let_rpc_send_complete.set()
await asyncio.wait_for(rpc_task, EVENT_TIMEOUT_S)
else:
refresh_task = asyncio.create_task(client.refresh_auth())
await asyncio.wait_for(get_entered.wait(), EVENT_TIMEOUT_S)
rpc_task = asyncio.create_task(
core._rpc_executor.rpc_call(RPCMethod.LIST_NOTEBOOKS, [])
)
await asyncio.wait_for(rpc_send_entered.wait(), EVENT_TIMEOUT_S)
let_get_complete.set()
await asyncio.wait_for(refresh_task, EVENT_TIMEOUT_S)
let_rpc_send_complete.set()
await asyncio.wait_for(rpc_task, EVENT_TIMEOUT_S)
finally:
# Always release the mock-transport gates so any in-flight handlers
# can return — even if the test errored above.
let_get_complete.set()
let_rpc_send_complete.set()
pending = [t for t in (rpc_task, refresh_task) if t is not None and not t.done()]
for t in pending:
t.cancel()
# Bounded join so cancelled tasks actually settle before the
# outer ``core.close()`` runs. Narrow to ``(CancelledError,
# Exception)`` so KeyboardInterrupt / SystemExit during the
# test still propagate.
if pending:
with contextlib.suppress(asyncio.CancelledError, Exception):
await asyncio.wait_for(
asyncio.gather(*pending, return_exceptions=True),
EVENT_TIMEOUT_S,
)
finally:
await core.close()
assert len(captured_post) == 1, (
f"Expected exactly one POST on the wire, got {len(captured_post)}: {captured_post!r}"
)
seen = captured_post[0]
cookie_is_new = "new_sid_cookie" in seen["cookie"]
cookie_is_old = "old_sid_cookie" in seen["cookie"]
csrf_is_new = b"CSRF_NEW" in seen["body"]
csrf_is_old = b"CSRF_OLD" in seen["body"]
sid_is_new = "SID_NEW" in seen["url"]
sid_is_old = "SID_OLD" in seen["url"]
# Sanity: each indicator is unambiguous (exactly one of old/new per axis).
# Without this, the coherence check below could false-pass when both
# "is_new" indicators are False simply because the markers weren't injected.
assert cookie_is_old ^ cookie_is_new, (
f"Cookie axis ambiguous (old={cookie_is_old}, new={cookie_is_new}): {seen['cookie']!r}"
)
assert csrf_is_old ^ csrf_is_new, (
f"CSRF axis ambiguous (old={csrf_is_old}, new={csrf_is_new}): body did not contain "
f"a recognizable CSRF marker"
)
assert sid_is_old ^ sid_is_new, (
f"Session-ID axis ambiguous (old={sid_is_old}, new={sid_is_new}): {seen['url']!r}"
)
# The invariant: all three axes must agree (all-OLD or all-NEW). Any mix
# indicates an unexpected yield in the prologue.
assert cookie_is_new == csrf_is_new == sid_is_new, (
f"Mixed-generation request observed (cookie_new={cookie_is_new}, "
f"csrf_new={csrf_is_new}, sid_new={sid_is_new}). A yield point was "
f"introduced between auth read and post() in the RPC transport path — re-run "
f"the AST guard above to find the offending await."
)