Files
writer-work-flow/apps/api/tests/test_kimi_oauth_endpoints.py

356 lines
13 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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共享实例断后台落库+ FakeSessionFactorymonkeypatch
`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
# 请求级 http 现为生成器依赖 `_http_client`CR-H2——override 它注 fake否则真生成器联网
app.dependency_overrides[routers.kimi_oauth._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 successTestClient 同步跑完 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
async def test_http_client_dependency_is_generator_and_closes_client(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""CR-H2请求级 `_http_client` 是**生成器依赖**,请求结束时 aclose 客户端(不泄漏连接)。
驱动生成器一轮enter→yield→exit断言是 async 生成器yield 时未关闭;耗尽时恰好
关闭一次。RED修复前 `_http_client` 不存在AttributeError
"""
import inspect
from ww_api.routers import kimi_oauth
closes = {"count": 0}
class _SpyClient:
def __init__(self, *args: Any, **kwargs: Any) -> None:
pass
async def __aenter__(self) -> _SpyClient:
return self
async def __aexit__(self, *exc: Any) -> None:
closes["count"] += 1
async def aclose(self) -> None:
closes["count"] += 1
monkeypatch.setattr("ww_api.routers.kimi_oauth.httpx.AsyncClient", _SpyClient)
# 是 async 生成器依赖(有 teardown非普通函数普通函数依赖无 teardown → 泄漏)。
assert inspect.isasyncgenfunction(kimi_oauth._http_client)
gen = kimi_oauth._http_client()
client = await gen.__anext__()
assert client is not None
assert closes["count"] == 0 # yield 时尚未关闭
with pytest.raises(StopAsyncIteration):
await gen.__anext__()
assert closes["count"] == 1 # 请求结束恰好关闭一次
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