perf(jupyter): 5s (workspace_id, user_id) validation cache
Jupyter 一次会话会拉几十次 auth_request(HTML shell / static / WebSocket / api/contents / autosave / kernels),每次都跑 JWT verify + WorkspaceMembers JOIN + Scripts.is_locked 查 + runtime RPC,重复开销大。 新增 module-level (workspace_id, user_id) -> payload 缓存: * 只缓存 membership 校验通过 + 拿到 runtime 信息的成功结果 (x-upstream-addr、x-jupyter-internal-token) * JWT 验签、lock check 仍每请求执行(前者是信任边界,后者 per-URI 状态易变) * TTL 5s,time.monotonic(),threading.Lock 保护 * 失败结果(lock 403 / runtime 500)不写缓存 折衷:被踢出 workspace 后最坏 5s 仍返 200;Runtime 单实例下 无需 Redis。新增 8 个 case 覆盖 hit/miss/TTL/lock-every-request/ jwt-every-request/failure-does-not-populate。
This commit is contained in:
@@ -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"}
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user