Skip to content

Commit 2d37ead

Browse files
committed
feat: 接线 TRAE callback 登录闭环
复核「写了但没接线」的模块时发现:TRAE 的登录函数全都只有定义没有调用, README 却写了「登录 TRAE 可用」——用户点进去会发现无路可走。 - TraeProvider.start_auth:生成登录 URL(machine/device id 随 state 下发, 保证落盘凭证与登录时用的设备标识一致) - TraeProvider.complete_callback:回调链接 → ExchangeToken → 归一化凭证; 回调只给 refreshToken,accessToken 必须换出来,绝不拿 refreshToken 充数 - POST /api/auth/upstream/start 支持 callback 轨道(CodeBuddy 仍走 poll) - GET /authorize 真正完成兑换并落库(原来只存 URL 不落盘),并新增: · state 必须与 start 发放的一致,且消费后即失效 —— 不允许回退 URL 里的 state,否则旧链接可以被重复兑换 · 响应与凭证列表都不回传 token · 兑换失败归一为 invalid_credential,不漏 HTTP 细节到 API 层 · 入库后立即触发额度探测 后端 541 测试 / 覆盖 100%。真实进程验证:flow=callback、state 校验、 假 token 兑换失败不入库。
1 parent 9690d88 commit 2d37ead

4 files changed

Lines changed: 307 additions & 17 deletions

File tree

‎src/main.py‎

Lines changed: 42 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -83,6 +83,8 @@ async def lifespan(app: FastAPI):
8383
app.state.stats_query = StatsQuery(db)
8484
app.state.upstream_auth = _upstream_auth(registry, config)
8585
app.state.pending_probes = []
86+
app.state.pending_callback_state = None
87+
app.state.pending_callback_user = None
8688

8789
# ------------------------------------------------------------ 鉴权依赖
8890

@@ -270,12 +272,19 @@ async def upstream_auth_start(payload: dict,
270272
principal: Principal = Depends(principal_from_request)):
271273
require_admin(principal)
272274
provider_id = str(payload.get("provider") or "")
273-
oauth = app.state.upstream_auth.get(provider_id)
274-
if oauth is None:
275-
raise InvalidRequest(f"provider {provider_id!r} does not support polling login")
276-
session = await oauth.start(principal.username)
275+
if provider_id in app.state.upstream_auth:
276+
session = await app.state.upstream_auth[provider_id].start(principal.username)
277+
else:
278+
provider = registry.get(provider_id)
279+
builder = getattr(provider, "start_auth", None)
280+
if not callable(builder):
281+
raise InvalidRequest(f"provider {provider_id!r} does not support login")
282+
session = builder(resolve_public_callback_url(config))
283+
app.state.pending_callback_state = session.state
284+
app.state.pending_callback_user = principal.username
285+
# 回调轨道没有本地轮询:登录结果由 /authorize 落库后由前端查凭证列表
277286
return {"flow": session.flow, "state": session.state, "auth_url": session.auth_url,
278-
"interval": session.interval}
287+
"interval": session.interval, "callback_url": session.callback_url}
279288

280289
@app.post("/api/auth/upstream/poll")
281290
async def upstream_auth_poll(payload: dict,
@@ -412,9 +421,36 @@ async def stats_by_provider(principal: Principal = Depends(principal_from_reques
412421

413422
@app.get("/authorize")
414423
async def authorize(request: Request):
415-
"""TRAE 浏览器 302 落点:只捕获 query,不落盘。"""
424+
"""TRAE 浏览器 302 落点。
425+
426+
回调不需要 API Key(浏览器不会带),因此这里不做鉴权,但:
427+
- 只接受带 refreshToken 的链接,其他一律拒绝
428+
- 换到的凭证直接落库,响应里绝不回传 token
429+
- state 必须与 start_auth 发放的一致,防止任意回调被塞进池子
430+
"""
416431
raw = str(request.url)
417432
app.state.last_callback_url = raw
433+
state = request.query_params.get("state")
434+
provider = registry.get("trae")
435+
pending = app.state.pending_callback_state
436+
if provider is None or pending is None or state != pending:
437+
# 没有进行中的登录,或 state 不匹配(含已消费后的重放)→ 拒绝。
438+
# 不允许回退到 URL 里的 state,否则旧链接可以被重复兑换。
439+
return JSONResponse(status_code=400, content=error_payload(
440+
"no pending TRAE login in progress", "invalid_request", 400))
441+
if not request.query_params.get("refreshToken") and not request.query_params.get("userJwt"):
442+
return JSONResponse(status_code=400, content=error_payload(
443+
"callback missing refreshToken", "invalid_request", 400))
444+
try:
445+
credential_data = await provider.complete_callback(raw, state)
446+
except UpstreamProtocolViolation as error:
447+
return JSONResponse(status_code=400, content=error_payload(
448+
str(error), "invalid_credential", 400))
449+
app.state.pending_callback_state = None
450+
credential_id = credentials.add(provider="trae", credential_data=credential_data,
451+
nickname=str(credential_data.get("nickname") or ""),
452+
added_by=app.state.pending_callback_user or "")
453+
schedule_probe(credential_id)
418454
return {"ok": True, "captured": True, "at": int(time.time())}
419455

420456
# 静态资源必须最后注册:catch-all 会匹配所有未命中的路径

‎src/provider/trae/client.py‎

Lines changed: 55 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -13,9 +13,15 @@
1313
import httpx
1414

1515
from ...engine.sse import iter_frames
16-
from ...provider.base import ErrKind, Event, Model, Quota
16+
from ...provider.base import AuthSession, ErrKind, Event, Model, Quota
1717
from . import events as trae_events
18-
from .callback import build_login_url, parse_callback_url
18+
from .callback import (
19+
CallbackInfo,
20+
build_login_url,
21+
credential_from_callback,
22+
new_machine_identity,
23+
parse_callback_url,
24+
)
1925
from .credential import TraeCredential, parse_credential
2026
from .events import (
2127
AGENT_HOST,
@@ -323,10 +329,55 @@ async def aclose(self) -> None:
323329
"""释放内部 HTTP 连接池。"""
324330
await self.client.aclose()
325331

326-
def parse_login_url(self, raw_url: str) -> dict:
327-
"""解析 TRAE 登录回调链接(Q17=C callback 轨道)。"""
332+
# ------------------------------------------------- callback 轨道(Q17=C)
333+
334+
def parse_login_url(self, raw_url: str) -> CallbackInfo:
335+
"""解析 TRAE 登录回调链接。"""
328336
return parse_callback_url(raw_url)
329337

330338
def build_login_url(self, callback_url: str, *, machine_id: str,
331339
device_id: str) -> str:
332340
return build_login_url(callback_url, machine_id=machine_id, device_id=device_id)
341+
342+
def start_auth(self, callback_url: str) -> AuthSession:
343+
"""生成登录 URL。machine/device id 由调用方保管,落盘凭证必须复用同一对。"""
344+
machine_id, device_id = new_machine_identity()
345+
return AuthSession(
346+
flow="callback", state=f"{machine_id}:{device_id}",
347+
callback_url=callback_url,
348+
auth_url=build_login_url(callback_url, machine_id=machine_id,
349+
device_id=device_id),
350+
)
351+
352+
async def complete_callback(self, raw_url: str, state: str) -> dict:
353+
"""回调链接 → ExchangeToken → 归一化凭证。
354+
355+
state 形如 ``machine_id:device_id``,用于保证落盘凭证与登录时用的
356+
设备标识一致(原实现里这两者不一致会导致登录态与凭证不匹配)。
357+
"""
358+
info = parse_callback_url(raw_url)
359+
machine_id, _, device_id = state.partition(":")
360+
if not machine_id or not device_id:
361+
raise UpstreamProtocolViolation("auth state missing machine/device id")
362+
363+
# 回调只给 refreshToken,accessToken 必须由 ExchangeToken 换出来。
364+
# 换失败就没有可用凭证,绝不能用 refreshToken 充当 accessToken 混进池子。
365+
credential = credential_from_callback(
366+
info, "", machine_id=machine_id, device_id=device_id)
367+
try:
368+
credential = await self.client.refresh_token(credential)
369+
except UpstreamHTTPError as error:
370+
# 上游拒绝兑换(refreshToken 过期/失效)→ 归一为协议违规,
371+
# 让调用方按「凭证无效」处理,而不是把 HTTP 细节漏到 API 层
372+
raise UpstreamProtocolViolation(
373+
f"token exchange rejected: {error.status}") from error
374+
if not credential.uid:
375+
uid, nickname = await self.client.get_user_info(credential)
376+
credential = TraeCredential(
377+
uid=uid or credential.uid, access_token=credential.access_token,
378+
refresh_token=credential.refresh_token, expires_at=credential.expires_at,
379+
domain=credential.domain, api_host=credential.api_host,
380+
machine_id=credential.machine_id, device_id=credential.device_id,
381+
enterprise_id=credential.enterprise_id,
382+
nickname=nickname or credential.nickname)
383+
return credential.to_dict()

‎tests/test_m15_operations.py‎

Lines changed: 206 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1012,11 +1012,20 @@ def admin_client(tmp_path):
10121012
yield app, client
10131013

10141014

1015-
def test_upstream_auth_start_requires_admin_and_known_provider(admin_client):
1016-
_app, client = admin_client
1017-
started = client.post("/api/auth/upstream/start", json={"provider": "codebuddy"})
1018-
assert started.status_code in (200, 400, 502)
1019-
assert client.post("/api/auth/upstream/start", json={"provider": "trae"}).status_code == 400
1015+
def test_upstream_auth_start_supports_both_providers(admin_client):
1016+
"""CodeBuddy 走 poll 轨道,TRAE 走 callback 轨道,unknown provider 报 400。"""
1017+
app, client = admin_client
1018+
# TRAE 的 callback 轨道不需要出网,必定成功并带入参回调地址
1019+
trae = client.post("/api/auth/upstream/start", json={"provider": "trae"})
1020+
assert trae.status_code == 200
1021+
body = trae.json()
1022+
assert body["flow"] == "callback"
1023+
assert body["callback_url"].endswith("/authorize")
1024+
assert "auth_callback_url" in body["auth_url"]
1025+
assert app.state.pending_callback_state == body["state"]
1026+
1027+
assert client.post("/api/auth/upstream/start",
1028+
json={"provider": "unknown"}).status_code == 400
10201029

10211030

10221031
def test_upstream_auth_poll_unknown_state(admin_client):
@@ -2052,3 +2061,195 @@ def test_schedule_probe_returns_early_when_credential_unreadable(admin_client):
20522061
assert created.status_code == 200
20532062
time.sleep(0.05)
20542063
assert len(app.state.pending_probes) == before
2064+
2065+
2066+
# ------------------------------------------- TRAE callback 登录闭环
2067+
2068+
def test_trae_start_auth_builds_login_url_with_public_callback():
2069+
from src.main import resolve_public_callback_url
2070+
from src.provider.trae.client import TraeProvider
2071+
2072+
settings = Settings(_env_file=None, APP_SECRET="s", PUBLIC_BASE_URL="https://gw.example")
2073+
provider = TraeProvider()
2074+
session = provider.start_auth(resolve_public_callback_url(settings))
2075+
2076+
assert session.flow == "callback"
2077+
assert session.callback_url == "https://gw.example/authorize"
2078+
assert "auth_callback_url=https%3A%2F%2Fgw.example%2Fauthorize" in session.auth_url
2079+
machine_id, _, device_id = session.state.partition(":")
2080+
assert len(machine_id) == 32 and len(device_id) == 32
2081+
2082+
2083+
async def test_trae_complete_callback_exchanges_token():
2084+
"""回调链接必须真的换 token,而不是只存 refreshToken。"""
2085+
import httpx as _httpx
2086+
2087+
from src.provider.trae.client import TraeClient, TraeProvider
2088+
2089+
def handler(request: _httpx.Request) -> _httpx.Response:
2090+
if request.url.path.endswith("ExchangeToken"):
2091+
return _httpx.Response(200, json={"Result": {
2092+
"Token": "ACCESS", "RefreshToken": "RT2", "TokenExpireAt": 1_800_000_000_000}})
2093+
return _httpx.Response(200, json={"Result": {"UserID": "uid-9", "ScreenName": "昵称"}})
2094+
2095+
transport = _httpx.MockTransport(handler)
2096+
provider = TraeProvider(client=TraeClient(
2097+
stream_client=_httpx.AsyncClient(transport=transport, timeout=None),
2098+
short_client=_httpx.AsyncClient(transport=transport, timeout=None)))
2099+
2100+
session = provider.start_auth("https://gw.example/authorize")
2101+
url = ("https://gw.example/authorize?refreshToken=RT1&userInfo="
2102+
"%7B%22uid%22%3A%22%22%7D")
2103+
data = await provider.complete_callback(url, session.state)
2104+
2105+
assert data["accessToken"] == "ACCESS"
2106+
assert data["refreshToken"] == "RT2"
2107+
assert data["expiresAt"] == 1_800_000_000
2108+
assert data["uid"] == "uid-9" # 回调没给 uid 时回退 GetUserInfo
2109+
assert data["nickname"] == "昵称"
2110+
assert data["machineId"] == session.state.partition(":")[0]
2111+
2112+
2113+
async def test_trae_complete_callback_rejects_bad_state():
2114+
from src.provider.trae.client import TraeProvider
2115+
from src.provider.trae.events import UpstreamProtocolViolation
2116+
2117+
provider = TraeProvider()
2118+
with pytest.raises(UpstreamProtocolViolation):
2119+
await provider.complete_callback("https://x/authorize?refreshToken=RT", "no-colon")
2120+
2121+
2122+
def test_authorize_completes_trae_login_end_to_end(tmp_path):
2123+
"""完整闭环:start → 浏览器回调 → 凭证入库 → 立即探测。"""
2124+
import httpx as _httpx
2125+
2126+
from src.provider.trae.client import TraeClient, TraeProvider
2127+
2128+
def handler(request: _httpx.Request) -> _httpx.Response:
2129+
if request.url.path.endswith("ExchangeToken"):
2130+
return _httpx.Response(200, json={"Result": {"Token": "ACCESS",
2131+
"RefreshToken": "RT2"}})
2132+
if request.url.endswith("ide_user_ent_usage"):
2133+
return _httpx.Response(200, json={"user_entitlement_pack_list": [
2134+
{"entitlement_base_info": {"quota": {"credits_limit": 100}},
2135+
"usage": {"credits_amount": 25}}]})
2136+
return _httpx.Response(200, json={"Result": {"UserID": "u1", "ScreenName": "n"}})
2137+
2138+
transport = _httpx.MockTransport(handler)
2139+
trae = TraeProvider(client=TraeClient(
2140+
stream_client=_httpx.AsyncClient(transport=transport, timeout=None),
2141+
short_client=_httpx.AsyncClient(transport=transport, timeout=None)))
2142+
2143+
settings = Settings(_env_file=None, APP_SECRET="s", DATA_DIR=str(tmp_path),
2144+
ADMIN_USERNAMES="root")
2145+
app = build_app(settings, providers={"trae": trae})
2146+
2147+
with TestClient(app) as client:
2148+
client.post("/api/auth/login", json={"username": "root", "password": "rootpw"})
2149+
started = client.post("/api/auth/upstream/start", json={"provider": "trae"}).json()
2150+
2151+
callback = ("/authorize?refreshToken=RT1&userInfo=%7B%22uid%22%3A%22u1%22%7D"
2152+
f"&state={started['state']}")
2153+
response = client.get(callback)
2154+
assert response.status_code == 200 and response.json()["captured"] is True
2155+
2156+
listed = client.get("/api/credentials").json()["credentials"]
2157+
assert len(listed) == 1 and listed[0]["provider"] == "trae"
2158+
# 响应与列表都不得泄漏 token
2159+
assert "ACCESS" not in response.text
2160+
assert "data_enc" not in listed[0]
2161+
2162+
# state 已消费,同一回调不可重放
2163+
with TestClient(app) as replay:
2164+
assert replay.get(callback).status_code == 400
2165+
2166+
2167+
def test_authorize_rejects_callback_without_token(tmp_path):
2168+
"""带 state 但没有 refreshToken/userJwt 的回调必须拒绝。"""
2169+
settings = Settings(_env_file=None, APP_SECRET="s", DATA_DIR=str(tmp_path),
2170+
ADMIN_USERNAMES="root")
2171+
app = build_app(settings)
2172+
with TestClient(app) as client:
2173+
client.post("/api/auth/login", json={"username": "root", "password": "rootpw"})
2174+
started = client.post("/api/auth/upstream/start", json={"provider": "trae"}).json()
2175+
response = client.get(f"/authorize?state={started['state']}")
2176+
assert response.status_code == 400
2177+
assert response.json()["error"]["code"] == "invalid_request"
2178+
2179+
2180+
def test_upstream_start_uses_poll_track_for_codebuddy(admin_client):
2181+
"""provider 在 upstream_auth 里时走 poll 轨道(main 275-276)。"""
2182+
app, client = admin_client
2183+
2184+
class FakeOAuth:
2185+
def __init__(self) -> None:
2186+
self.store = type("S", (), {"cancel": staticmethod(lambda *_: True)})()
2187+
2188+
async def start(self, username):
2189+
from src.provider.base import AuthSession
2190+
2191+
return AuthSession(flow="poll", state="local-reservation",
2192+
auth_url="https://auth.example/x", interval=5)
2193+
2194+
app.state.upstream_auth["codebuddy"] = FakeOAuth()
2195+
body = client.post("/api/auth/upstream/start", json={"provider": "codebuddy"}).json()
2196+
assert body["flow"] == "poll"
2197+
assert body["state"] == "local-reservation"
2198+
assert body["callback_url"] is None
2199+
2200+
2201+
async def test_complete_callback_does_not_use_refresh_token_as_access_token():
2202+
"""ExchangeToken 失败时不得把 refreshToken 当 accessToken 塞进池子。"""
2203+
import httpx as _httpx
2204+
2205+
from src.provider.trae.client import TraeClient, TraeProvider
2206+
from src.provider.trae.events import UpstreamProtocolViolation
2207+
2208+
def handler(request: _httpx.Request) -> _httpx.Response:
2209+
if request.url.path.endswith("ExchangeToken"):
2210+
return _httpx.Response(400, content=b"bad refresh token")
2211+
return _httpx.Response(200, json={"Result": {"UserID": "u", "ScreenName": "n"}})
2212+
2213+
transport = _httpx.MockTransport(handler)
2214+
provider = TraeProvider(client=TraeClient(
2215+
stream_client=_httpx.AsyncClient(transport=transport, timeout=None),
2216+
short_client=_httpx.AsyncClient(transport=transport, timeout=None)))
2217+
2218+
session = provider.start_auth("https://gw.example/authorize")
2219+
with pytest.raises(UpstreamProtocolViolation):
2220+
await provider.complete_callback(
2221+
"https://gw.example/authorize?refreshToken=RT", session.state)
2222+
2223+
2224+
async def test_complete_callback_raises_when_no_token_available():
2225+
import httpx as _httpx
2226+
2227+
from src.provider.trae.client import TraeClient, TraeProvider
2228+
from src.provider.trae.events import UpstreamProtocolViolation
2229+
2230+
def handler(_request: _httpx.Request) -> _httpx.Response:
2231+
return _httpx.Response(400, content=b"rejected")
2232+
2233+
transport = _httpx.MockTransport(handler)
2234+
provider = TraeProvider(client=TraeClient(
2235+
stream_client=_httpx.AsyncClient(transport=transport, timeout=None),
2236+
short_client=_httpx.AsyncClient(transport=transport, timeout=None)))
2237+
session = provider.start_auth("https://gw.example/authorize")
2238+
with pytest.raises(UpstreamProtocolViolation):
2239+
await provider.complete_callback(
2240+
"https://gw.example/authorize?refreshToken=RT", session.state)
2241+
2242+
2243+
def test_authorize_reports_invalid_credential(tmp_path):
2244+
"""回调结构损坏 → 400 invalid_credential(main 446-448)。"""
2245+
settings = Settings(_env_file=None, APP_SECRET="s", DATA_DIR=str(tmp_path),
2246+
ADMIN_USERNAMES="root")
2247+
app = build_app(settings)
2248+
with TestClient(app) as client:
2249+
client.post("/api/auth/login", json={"username": "root", "password": "rootpw"})
2250+
client.post("/api/auth/upstream/start", json={"provider": "trae"})
2251+
# 有 refreshToken 但状态串损坏 → complete_callback 抛协议违规
2252+
app.state.pending_callback_state = "broken"
2253+
response = client.get("/authorize?refreshToken=RT&state=broken")
2254+
assert response.status_code == 400
2255+
assert response.json()["error"]["code"] == "invalid_credential"

‎tests/test_m1a_trae.py‎

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -876,7 +876,9 @@ def test_non_admin_cannot_write(client):
876876
"credential": {}}).status_code == 403
877877

878878

879-
def test_authorize_callback_capture(client):
879+
def test_authorize_rejects_callback_without_pending_login(client):
880+
"""没有进行中的登录时,任意回调不得被塞进凭证池。"""
880881
response = client.get("/authorize?refreshToken=RT")
881-
assert response.status_code == 200 and response.json()["captured"] is True
882+
assert response.status_code == 400
883+
assert response.json()["error"]["code"] == "invalid_request"
882884
assert "refreshToken=RT" in client.app.state.last_callback_url

0 commit comments

Comments
 (0)