fix(backend): Kimi OAuth 轮询移出 DB 会话——只在最终持久化开短会话(CR-H1)

This commit is contained in:
Yaojia Wang
2026-07-08 10:37:42 +02:00
parent ed0a679322
commit ca8fbc89a1
2 changed files with 176 additions and 37 deletions

View File

@@ -15,6 +15,8 @@
from __future__ import annotations
import os
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import Any
import pytest
@@ -51,12 +53,65 @@ class _ScriptedHttp:
return None
class _TrackingSessionFactory:
"""记录同一时刻在持会话数open_count+ 累计开启次数total_opened的 session 工厂替身。
CR-H1 断言用:轮询期间 open_count 应恒 0无会话被长时间占用持久化时 total_opened≥1。
"""
def __init__(self) -> None:
self.open_count = 0
self.total_opened = 0
def __call__(self) -> Any:
from fakes_projects import FakeSession
@asynccontextmanager
async def _cm() -> AsyncIterator[Any]:
self.open_count += 1
self.total_opened += 1
try:
yield FakeSession()
finally:
self.open_count -= 1
return _cm()
class _ObservingHttp:
"""观测型 scripted http每次 POST 到 TOKEN_URL轮询时记录当下在持会话数。aclose 无操作。"""
def __init__(
self,
responses: list[_FakeResponse],
*,
factory: _TrackingSessionFactory,
token_url: str,
) -> None:
self._responses = responses
self._factory = factory
self._token_url = token_url
self.calls: list[tuple[str, dict[str, str]]] = []
self.open_during_poll: list[int] = []
async def post(self, url: str, *, data: dict[str, str]) -> _FakeResponse:
idx = len(self.calls)
self.calls.append((url, data))
if url == self._token_url: # 轮询 POST区别于 device_authorization POST
self.open_during_poll.append(self._factory.open_count)
return self._responses[idx]
async def aclose(self) -> None:
return None
def _make_client(
*,
store: FakeCredentialStore,
http: _ScriptedHttp,
http: Any,
enc_key: str,
monkeypatch: pytest.MonkeyPatch,
session_factory: Any = None,
) -> TestClient:
os.environ["CREDENTIAL_ENC_KEY"] = enc_key
from ww_config import get_settings
@@ -88,11 +143,13 @@ def _make_client(
job_repo = FakeJobRepo()
monkeypatch.setattr(job_runner, "SqlJobRepo", lambda _session: job_repo)
sf = session_factory if session_factory is not None else FakeSessionFactory()
app = create_app()
app.dependency_overrides[get_credential_store] = lambda: store
app.dependency_overrides[get_job_repo] = lambda: job_repo
app.dependency_overrides[get_session] = lambda: FakeSession()
app.dependency_overrides[get_session_factory] = lambda: FakeSessionFactory()
app.dependency_overrides[get_session_factory] = lambda: sf
app.dependency_overrides[routers.kimi_oauth._default_http_client] = lambda: http
return TestClient(app)
@@ -163,6 +220,49 @@ def test_start_polls_through_authorization_pending(monkeypatch: pytest.MonkeyPat
assert (STUB_OWNER_ID, "kimi-code") in store.oauth
def test_poll_loop_does_not_hold_db_session(monkeypatch: pytest.MonkeyPatch) -> None:
"""CR-H1设备授权轮询期间**不得**占用 DB 会话(否则连接池被长时间空占→饥饿)。
轮询循环只做 HTTP无 DB仅最终持久化才开一个**短**会话。断言:轮询 POST 期间
session_factory 未开会话open_count==0而持久化确开过短会话凭据已落库、响应不泄 token。
"""
from ww_api.services.kimi_oauth import TOKEN_URL
enc_key = Fernet.generate_key().decode()
store = FakeCredentialStore()
factory = _TrackingSessionFactory()
# device auth → poll #1 pending → poll #2 success≥2 次轮询 POST
http = _ObservingHttp(
[
_device_auth_response(),
_FakeResponse(400, {"error": "authorization_pending"}),
_token_response(),
],
factory=factory,
token_url=TOKEN_URL,
)
client = _make_client(
store=store,
http=http,
enc_key=enc_key,
monkeypatch=monkeypatch,
session_factory=factory,
)
resp = client.post("/settings/providers/kimi-code/oauth/start")
assert resp.status_code == 202
# 核心断言:轮询 POST 期间无 session_factory 会话在持。
assert http.open_during_poll # 确有轮询 POST 发生
assert all(c == 0 for c in http.open_during_poll)
# 持久化时确开过一个短会话。
assert factory.total_opened >= 1
# 凭据已落库;响应绝不含 token。
assert (STUB_OWNER_ID, "kimi-code") in store.oauth
assert "secret-acc" not in resp.text
assert "secret-ref" not in resp.text
def test_disconnect_clears_credential(monkeypatch: pytest.MonkeyPatch) -> None:
enc_key = Fernet.generate_key().decode()
store = FakeCredentialStore()