diff --git a/backend/src/backend/jupyter.py b/backend/src/backend/jupyter.py index 74d01e1..489bd96 100644 --- a/backend/src/backend/jupyter.py +++ b/backend/src/backend/jupyter.py @@ -8,6 +8,8 @@ """ import re +import threading +import time from common.auth.jwt import JwtError, verify_jwt_token from common.auth.membership import MembershipError, load_active_membership @@ -20,6 +22,26 @@ from sqlalchemy.ext.asyncio import AsyncSession from backend.dependencies import database_session from backend.runtime_client import RuntimeClientError +# --------------------------------------------------------------------------- +# (workspace_id, user_id) -> (expires_at_monotonic, payload) 的 5 秒验证结果缓存。 +# +# payload 至少包含 x-upstream-addr 与 x-jupyter-internal-token,只缓存 +# "membership 校验通过 + 拿到 runtime 信息"后的成功结果;403/500 不写缓存。 +# +# 设计说明: +# * TTL 只有 5 秒,且 Runtime 容器单实例(CLAUDE.md "Service rules": +# "Runtime must stay single-replica while file leases and Jupyter tickets +# use the simplified implementation"),module-level 内存缓存是安全的, +# 无需 Redis 之类的外部存储。 +# * JWT 验签与 lock check 不进缓存:前者是每请求必须的信任边界;后者是 +# per-URI 且 5s 内可能解锁/加锁,跨用户/跨 notebook 不应共享缓存。 +# * 折衷:用户被踢出 workspace / membership 撤销后,最坏 5 秒内本接口仍会 +# 对已缓存的 (workspace_id, user_id) 返回 200,这是可接受的折衷。 +_JUPYTER_AUTH_CACHE: dict[tuple[str, str], tuple[float, dict[str, str]]] = {} +_JUPYTER_AUTH_CACHE_LOCK = threading.Lock() +_JUPYTER_AUTH_CACHE_TTL_SECONDS = 5.0 + + router = APIRouter(tags=["jupyter"]) security = HTTPBearer(auto_error=False) @@ -89,6 +111,24 @@ async def load_active_membership_or_403( ) from exc +def _jupyter_auth_cache_get(workspace_id: str, user_id: str) -> dict[str, str] | None: + with _JUPYTER_AUTH_CACHE_LOCK: + entry = _JUPYTER_AUTH_CACHE.get((workspace_id, user_id)) + if entry is None: + return None + expires_at, payload = entry + if time.monotonic() >= expires_at: + _JUPYTER_AUTH_CACHE.pop((workspace_id, user_id), None) + return None + return payload + + +def _jupyter_auth_cache_put(workspace_id: str, user_id: str, payload: dict[str, str]) -> None: + expires_at = time.monotonic() + _JUPYTER_AUTH_CACHE_TTL_SECONDS + with _JUPYTER_AUTH_CACHE_LOCK: + _JUPYTER_AUTH_CACHE[(workspace_id, user_id)] = (expires_at, payload) + + # 供 Nginx auth_request 调用:验证访问 Jupyter 的身份、成员关系和文件锁, # 再返回应转发到的 Jupyter 地址及内部令牌。 @router.get("/api/v1/auth/jupyter") @@ -132,8 +172,14 @@ async def verify_jupyter_access( detail="Invalid Authentication Token", ) - await load_active_membership_or_403(session, user_id, workspace_id) + # JWT 验签之后、昂贵的 membership/runtime 查找之前先查缓存。命中时跳过 + # membership 与 runtime,但仍要跑下面的 lock check(per-URI,缓存不含它)。 + cached = _jupyter_auth_cache_get(workspace_id, user_id) + if cached is None: + await load_active_membership_or_403(session, user_id, workspace_id) + # lock check 永远执行、不进缓存:同一 (workspace_id, user_id) 的不同 URI + # 状态不同,且 5s 内可能解锁/加锁。 notebook_path = extract_notebook_path(original_uri, workspace_id) if notebook_path and await check_notebook_is_locked( session, @@ -146,6 +192,11 @@ async def verify_jupyter_access( detail=f"Notebook '{notebook_path}' is currently locked", ) + if cached is not None: + response.headers["x-upstream-addr"] = cached["x_upstream_addr"] + response.headers["x-jupyter-internal-token"] = cached["x_jupyter_internal_token"] + return {"status": "ok"} + runtime_client = request.app.state.runtime_client ws_info = await runtime_client.get_workspace(workspace_id) if not ws_info or ws_info.get("status") != "running": @@ -166,6 +217,14 @@ async def verify_jupyter_access( detail="Jupyter instance returned no port", ) - response.headers["x-upstream-addr"] = f"{jupyter_base_url}:{target_port}" - response.headers["x-jupyter-internal-token"] = jupyter_token or "" + headers_payload = { + "x_upstream_addr": f"{jupyter_base_url}:{target_port}", + "x_jupyter_internal_token": jupyter_token or "", + } + # 只缓存成功结果;lock check 失败(403)或 runtime 启动失败(500)在上方 + # 已提前返回,不会走到这里污染缓存。 + _jupyter_auth_cache_put(workspace_id, user_id, headers_payload) + + response.headers["x-upstream-addr"] = headers_payload["x_upstream_addr"] + response.headers["x-jupyter-internal-token"] = headers_payload["x_jupyter_internal_token"] return {"status": "ok"} diff --git a/backend/tests/test_jupyter_auth_cache.py b/backend/tests/test_jupyter_auth_cache.py new file mode 100644 index 0000000..0b33f68 --- /dev/null +++ b/backend/tests/test_jupyter_auth_cache.py @@ -0,0 +1,236 @@ +"""Unit tests for the 5s auth-result cache in ``backend.jupyter``. + +Jupyter 一次会话会触发几十次 Nginx ``auth_request``;本缓存按 +``(workspace_id, user_id)`` 缓存 membership + runtime 的查找结果,避免 +每次都跑 DB JOIN 与跨进程 RPC。JWT 验签与 per-URI 的 lock check **不进 +缓存**,每请求都执行。 + +这些测试直接调用 ``verify_jupyter_access``(不经过 FastAPI TestClient), +用 SimpleNamespace 构造 fake request / response / runtime,并用 +monkeypatch 替换 JWT / membership / lock / runtime 的调用点来统计次数。 +""" + +from __future__ import annotations + +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import HTTPException + +import backend.jupyter as jupyter_module +from backend.jupyter import ( + _JUPYTER_AUTH_CACHE, + verify_jupyter_access, +) +from backend.runtime_client import RuntimeClientError + +WS_ID = "01WS0000000000000000000A" +USER_ID = "01USR0000000000000000000A" +NOTEBOOK_URI = f"/jupyter/{WS_ID}/notebooks/a.ipynb" +NON_NOTEBOOK_URI = f"/jupyter/{WS_ID}/tree" + +_DESCRIPTOR = { + "status": "running", + "workspace_id": WS_ID, + "base_url": "http://runtime", + "port": 34567, + "token": "jupyter-token", +} + + +@pytest.fixture(autouse=True) +def _clear_cache() -> None: + _JUPYTER_AUTH_CACHE.clear() + yield + _JUPYTER_AUTH_CACHE.clear() + + +def _make_context(runtime_client, uri: str = NOTEBOOK_URI) -> tuple[SimpleNamespace, SimpleNamespace]: + request = SimpleNamespace( + headers={ + "X-Original-Workspace-Id": WS_ID, + "X-Original-URI": uri, + }, + cookies={"access_token": "a.b.c"}, + app=SimpleNamespace(state=SimpleNamespace(runtime_client=runtime_client)), + ) + response = SimpleNamespace(headers={}) + return request, response + + +def _make_runtime_client(descriptor=None) -> SimpleNamespace: + client = SimpleNamespace() + client.get_workspace = AsyncMock(return_value=descriptor if descriptor is not None else _DESCRIPTOR) + client.start_workspace = AsyncMock(return_value=_DESCRIPTOR) + return client + + +def _setup_mocks( + monkeypatch: pytest.MonkeyPatch, + *, + user_id: str = USER_ID, + locked: bool = False, + runtime_client: SimpleNamespace | None = None, +) -> tuple[SimpleNamespace, SimpleNamespace, AsyncMock, AsyncMock, AsyncMock, SimpleNamespace]: + """Patch the call points once and return fakes for counting. + + ``verify_jwt_token`` 默认替换为固定 payload;需要计数的测试可在此之后 + 再次 ``monkeypatch.setattr`` 覆盖(后设置者生效)。 + """ + monkeypatch.setattr( + jupyter_module, + "verify_jwt_token", + lambda _token: {"sub": user_id}, + ) + membership = AsyncMock() + monkeypatch.setattr(jupyter_module, "load_active_membership_or_403", membership) + lock_check = AsyncMock(return_value=locked) + monkeypatch.setattr(jupyter_module, "check_notebook_is_locked", lock_check) + runtime = runtime_client if runtime_client is not None else _make_runtime_client() + request, response = _make_context(runtime) + return request, response, membership, lock_check, runtime + + +async def _call_once(request, response) -> None: + await verify_jupyter_access( + request, + response, + auth=None, + session=AsyncMock(), + ) + + +async def test_cache_hit_skips_membership_and_runtime(monkeypatch) -> None: + """第一次跑完整路径,第二次同样 (ws, user) 跳过 membership + runtime。""" + request, response, membership, lock_check, runtime = _setup_mocks(monkeypatch) + + await _call_once(request, response) + assert membership.await_count == 1 + assert runtime.get_workspace.await_count == 1 + assert runtime.start_workspace.await_count == 0 + assert response.headers["x-upstream-addr"] == "http://runtime:34567" + assert response.headers["x-jupyter-internal-token"] == "jupyter-token" + + # 第二次请求:缓存命中,membership / runtime 不再执行。 + response.headers = {} + await _call_once(request, response) + assert membership.await_count == 1 + assert runtime.get_workspace.await_count == 1 + assert runtime.start_workspace.await_count == 0 + # lock check 每请求都跑。 + assert lock_check.await_count == 2 + # 缓存命中也要写 headers。 + assert response.headers["x-upstream-addr"] == "http://runtime:34567" + assert response.headers["x-jupyter-internal-token"] == "jupyter-token" + + +async def test_cache_miss_runs_full_path(monkeypatch) -> None: + """清空缓存后第一次请求必须跑 membership + runtime。""" + _JUPYTER_AUTH_CACHE.clear() + request, response, membership, lock_check, runtime = _setup_mocks(monkeypatch) + + await _call_once(request, response) + + assert membership.await_count == 1 + assert runtime.get_workspace.await_count == 1 + assert lock_check.await_count == 1 + + +async def test_lock_check_runs_every_request_even_on_cache_hit(monkeypatch) -> None: + """缓存命中时仍要执行 lock check(per-URI,不进缓存)。""" + request, response, membership, lock_check, runtime = _setup_mocks(monkeypatch) + + await _call_once(request, response) # 预热缓存 + response.headers = {} + await _call_once(request, response) # 缓存命中 + + assert membership.await_count == 1 + assert lock_check.await_count == 2 + + +async def test_jwt_verify_runs_every_request(monkeypatch) -> None: + """JWT 验签是每请求的安全边界,缓存命中也不能跳过。""" + verify_calls: list[int] = [] + + def _fake_verify(_token): + verify_calls.append(1) + return {"sub": USER_ID} + + request, response, membership, lock_check, runtime = _setup_mocks(monkeypatch) + monkeypatch.setattr(jupyter_module, "verify_jwt_token", _fake_verify) + + await _call_once(request, response) + response.headers = {} + await _call_once(request, response) # 缓存命中 + + assert membership.await_count == 1 + assert len(verify_calls) == 2 + + +async def test_cache_ttl_expires_after_5s(monkeypatch) -> None: + """TTL 用 time.monotonic;5s 后缓存失效,重新走完整路径。""" + now = [100.0] + monkeypatch.setattr("time.monotonic", lambda: now[0]) + request, response, membership, lock_check, runtime = _setup_mocks(monkeypatch) + + await _call_once(request, response) # t=100,写缓存(expires=105) + assert membership.await_count == 1 + + response.headers = {} + await _call_once(request, response) # t=100,缓存命中 + assert membership.await_count == 1 + + now[0] = 105.0 # 恰好到过期时刻 -> 缓存失效 + response.headers = {} + await _call_once(request, response) + assert membership.await_count == 2 + assert runtime.get_workspace.await_count == 2 + + +async def test_lock_check_failure_does_not_populate_cache(monkeypatch) -> None: + """lock 403 不进缓存。""" + request, response, membership, lock_check, runtime = _setup_mocks(monkeypatch, locked=True) + + with pytest.raises(HTTPException) as excinfo: + await _call_once(request, response) + + assert excinfo.value.status_code == 403 + assert _JUPYTER_AUTH_CACHE == {} + assert jupyter_module._jupyter_auth_cache_get(WS_ID, USER_ID) is None + + +async def test_runtime_start_failure_does_not_populate_cache(monkeypatch) -> None: + """runtime 启动失败(500)不进缓存。""" + runtime = _make_runtime_client() + runtime.get_workspace = AsyncMock(return_value=None) # 未运行 -> 走 start + runtime.start_workspace = AsyncMock( + side_effect=RuntimeClientError(500, {"code": "JUPYTER_START_FAILED"}) + ) + request, response, membership, lock_check, _rt = _setup_mocks( + monkeypatch, runtime_client=runtime + ) + + with pytest.raises(HTTPException) as excinfo: + await _call_once(request, response) + + assert excinfo.value.status_code == 500 + assert _JUPYTER_AUTH_CACHE == {} + + +async def test_cache_keyed_per_user_and_workspace(monkeypatch) -> None: + """不同 user 共享 workspace 时不串缓存。""" + request, response, membership, lock_check, runtime = _setup_mocks(monkeypatch) + await _call_once(request, response) # USER_A 预热缓存 + assert membership.await_count == 1 + + # 换一个 user_id,同一个 workspace -> 缓存 key 不同,必须重新跑完整路径。 + request2, response2 = _make_context(runtime) + monkeypatch.setattr( + jupyter_module, + "verify_jwt_token", + lambda _token: {"sub": "01USR0000000000000000000B"}, + ) + await _call_once(request2, response2) + assert membership.await_count == 2 + assert runtime.get_workspace.await_count == 2