fix(backend): Kimi OAuth 轮询移出 DB 会话——只在最终持久化开短会话(CR-H1)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user