@@ -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
10221031def 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"
0 commit comments