314 lines
11 KiB
Python
314 lines
11 KiB
Python
"""K1.3 端点测试:Kimi Code OAuth start/disconnect/status(内存替身,无 DB/无网络)。
|
||
|
||
覆盖:
|
||
- POST .../start → 202 + user_code/verification_uri + 创建 job;后台轮询 poll_token 成功
|
||
→ 存加密 oauth_enc + job done;**token 不进 job 结果/响应**;
|
||
- 后台轮询 authorization_pending → 继续轮询直到成功;
|
||
- POST .../disconnect → 清除凭据行;
|
||
- GET .../status → connected 真值(无 token 本体)。
|
||
|
||
后台 work 用 `run_job` 自建独立 session(请求 session 已关闭)——测试经依赖 override
|
||
注 fake http + fake store(共享实例,断后台落库)+ FakeSessionFactory;monkeypatch
|
||
`kimi_oauth.asyncio.sleep` 为 no-op(避免真等 interval 秒)。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
from collections.abc import AsyncIterator
|
||
from contextlib import asynccontextmanager
|
||
from typing import Any
|
||
|
||
import pytest
|
||
from cryptography.fernet import Fernet
|
||
from fakes_providers import FakeCredentialStore
|
||
from fastapi.testclient import TestClient
|
||
from ww_api.services.credentials import STUB_OWNER_ID
|
||
from ww_api.services.kimi_oauth import decrypt_oauth_bundle
|
||
from ww_db import get_session
|
||
|
||
|
||
class _FakeResponse:
|
||
def __init__(self, status_code: int, body: dict[str, Any]) -> None:
|
||
self.status_code = status_code
|
||
self._body = body
|
||
|
||
def json(self) -> Any:
|
||
return self._body
|
||
|
||
|
||
class _ScriptedHttp:
|
||
"""按调用序号返回预设响应(device auth → poll attempts)。aclose 无操作。"""
|
||
|
||
def __init__(self, responses: list[_FakeResponse]) -> None:
|
||
self._responses = responses
|
||
self.calls: list[tuple[str, dict[str, str]]] = []
|
||
|
||
async def post(self, url: str, *, data: dict[str, str]) -> _FakeResponse:
|
||
idx = len(self.calls)
|
||
self.calls.append((url, data))
|
||
return self._responses[idx]
|
||
|
||
async def aclose(self) -> None:
|
||
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: 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
|
||
|
||
get_settings.cache_clear()
|
||
|
||
from fakes_projects import FakeJobRepo, FakeSession, FakeSessionFactory
|
||
from ww_api import routers
|
||
from ww_api.main import create_app
|
||
from ww_api.services import job_runner
|
||
from ww_api.services.project_deps import (
|
||
get_job_repo,
|
||
get_session_factory,
|
||
)
|
||
from ww_api.services.provider_deps import get_credential_store
|
||
|
||
# 后台轮询不真等 interval 秒。
|
||
async def _no_sleep(_seconds: float) -> None:
|
||
return None
|
||
|
||
monkeypatch.setattr("ww_api.routers.kimi_oauth.asyncio.sleep", _no_sleep)
|
||
# 后台 work 自建 SqlCredentialStore(session)——指到共享 fake store(断后台落库)。
|
||
monkeypatch.setattr(routers.kimi_oauth, "SqlCredentialStore", lambda _session: store)
|
||
# 后台 work 内调模块级 `_default_http_client()` 自建 http(非 dep)——monkeypatch 它指到
|
||
# scripted fake,否则后台轮询会真联网(见 device-flow 注入纪律)。
|
||
monkeypatch.setattr(routers.kimi_oauth, "_default_http_client", lambda: http)
|
||
# `run_job` 默认 repo_factory 构 `SqlJobRepo(session)`——FakeSession 无 .execute,
|
||
# 故 monkeypatch 模块级 `job_runner.SqlJobRepo` 指到共享 fake(见 memory/gotchas)。
|
||
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: sf
|
||
app.dependency_overrides[routers.kimi_oauth._default_http_client] = lambda: http
|
||
return TestClient(app)
|
||
|
||
|
||
def _device_auth_response() -> _FakeResponse:
|
||
return _FakeResponse(
|
||
200,
|
||
{
|
||
"device_code": "dev-1",
|
||
"user_code": "WXYZ-9999",
|
||
"verification_uri": "https://kimi.com/device",
|
||
"verification_uri_complete": "https://kimi.com/device?code=WXYZ-9999",
|
||
"expires_in": 600,
|
||
"interval": 5,
|
||
},
|
||
)
|
||
|
||
|
||
def _token_response() -> _FakeResponse:
|
||
return _FakeResponse(
|
||
200, {"access_token": "secret-acc", "refresh_token": "secret-ref", "expires_in": 900}
|
||
)
|
||
|
||
|
||
def test_start_returns_202_and_polls_to_success(monkeypatch: pytest.MonkeyPatch) -> None:
|
||
enc_key = Fernet.generate_key().decode()
|
||
store = FakeCredentialStore()
|
||
# device auth → poll #1 success(TestClient 同步跑完 background task)。
|
||
http = _ScriptedHttp([_device_auth_response(), _token_response()])
|
||
client = _make_client(store=store, http=http, enc_key=enc_key, monkeypatch=monkeypatch)
|
||
|
||
resp = client.post("/settings/providers/kimi-code/oauth/start")
|
||
|
||
assert resp.status_code == 202
|
||
body = resp.json()
|
||
assert body["user_code"] == "WXYZ-9999"
|
||
assert body["verification_uri"] == "https://kimi.com/device"
|
||
assert body["interval"] == 5
|
||
assert "job_id" in body
|
||
# 响应绝不含 token。
|
||
assert "secret-acc" not in resp.text
|
||
assert "secret-ref" not in resp.text
|
||
|
||
# 后台轮询成功 → 加密 oauth 凭据落库(共享 store)。
|
||
cred = store.oauth.get((STUB_OWNER_ID, "kimi-code"))
|
||
assert cred is not None
|
||
token = decrypt_oauth_bundle(cred, key=enc_key)
|
||
assert token.access_token == "secret-acc"
|
||
assert token.refresh_token == "secret-ref"
|
||
|
||
|
||
def test_start_polls_through_authorization_pending(monkeypatch: pytest.MonkeyPatch) -> None:
|
||
enc_key = Fernet.generate_key().decode()
|
||
store = FakeCredentialStore()
|
||
http = _ScriptedHttp(
|
||
[
|
||
_device_auth_response(),
|
||
_FakeResponse(400, {"error": "authorization_pending"}),
|
||
_token_response(),
|
||
]
|
||
)
|
||
client = _make_client(store=store, http=http, enc_key=enc_key, monkeypatch=monkeypatch)
|
||
|
||
resp = client.post("/settings/providers/kimi-code/oauth/start")
|
||
|
||
assert resp.status_code == 202
|
||
# pending 后继续轮询 → 第二次成功落库。
|
||
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()
|
||
store.oauth[(STUB_OWNER_ID, "kimi-code")] = b"blob"
|
||
http = _ScriptedHttp([])
|
||
client = _make_client(store=store, http=http, enc_key=enc_key, monkeypatch=monkeypatch)
|
||
|
||
resp = client.post("/settings/providers/kimi-code/oauth/disconnect")
|
||
|
||
assert resp.status_code == 200
|
||
assert resp.json()["disconnected"] is True
|
||
assert (STUB_OWNER_ID, "kimi-code") not in store.oauth
|
||
|
||
|
||
def test_status_reflects_connection(monkeypatch: pytest.MonkeyPatch) -> None:
|
||
from datetime import UTC, datetime, timedelta
|
||
|
||
from ww_api.services.kimi_oauth import TokenSet, encrypt_oauth_bundle
|
||
|
||
enc_key = Fernet.generate_key().decode()
|
||
store = FakeCredentialStore()
|
||
expires = datetime.now(UTC) + timedelta(seconds=900)
|
||
token = TokenSet(access_token="a", refresh_token="r", expires_at=expires)
|
||
store.oauth[(STUB_OWNER_ID, "kimi-code")] = encrypt_oauth_bundle(token, key=enc_key)
|
||
http = _ScriptedHttp([])
|
||
client = _make_client(store=store, http=http, enc_key=enc_key, monkeypatch=monkeypatch)
|
||
|
||
resp = client.get("/settings/providers/kimi-code/oauth/status")
|
||
|
||
assert resp.status_code == 200
|
||
body = resp.json()
|
||
assert body["connected"] is True
|
||
assert body["expires_at"] is not None
|
||
# 状态绝不回 token 本体。
|
||
assert "access_token" not in body
|
||
assert "refresh_token" not in body
|
||
|
||
|
||
def test_status_not_connected_when_no_credential(monkeypatch: pytest.MonkeyPatch) -> None:
|
||
enc_key = Fernet.generate_key().decode()
|
||
store = FakeCredentialStore()
|
||
http = _ScriptedHttp([])
|
||
client = _make_client(store=store, http=http, enc_key=enc_key, monkeypatch=monkeypatch)
|
||
|
||
resp = client.get("/settings/providers/kimi-code/oauth/status")
|
||
|
||
assert resp.status_code == 200
|
||
assert resp.json()["connected"] is False
|