@@ -15,8 +15,8 @@ use std::time::Duration;
1515use tokio_util:: sync:: CancellationToken ;
1616
1717const CLIENT_ID : & str = "b1a00492-073a-47ea-816f-4c329264a828" ;
18- const DEVICE_AUTHORIZATION_URL : & str = "https://auth.x.ai/oauth2/device/code " ;
19- const TOKEN_URL : & str = "https://auth.x.ai/oauth2/token " ;
18+ const ISSUER : & str = "https://auth.x.ai" ;
19+ const DISCOVERY_URL : & str = "https://auth.x.ai/.well-known/openid-configuration " ;
2020const DEVICE_CODE_GRANT_TYPE : & str = "urn:ietf:params:oauth:grant-type:device_code" ;
2121const SCOPE : & str = "openid profile email offline_access grok-cli:access api:access" ;
2222const XAI_BASE_URL : & str = "https://api.x.ai/v1" ;
@@ -30,6 +30,63 @@ const SHORT_TOKEN_REFRESH_LEEWAY_MS: i64 = 2 * 60 * 1000;
3030const LONG_TOKEN_REFRESH_LEEWAY_MS : i64 = 60 * 60 * 1000 ;
3131const SHORT_TOKEN_THRESHOLD_MS : i64 = 45 * 60 * 1000 ;
3232
33+ #[ derive( Debug , Deserialize ) ]
34+ struct DiscoveryDocument {
35+ issuer : String ,
36+ device_authorization_endpoint : String ,
37+ token_endpoint : String ,
38+ }
39+
40+ struct OAuthEndpoints {
41+ device_authorization : reqwest:: Url ,
42+ token : reqwest:: Url ,
43+ }
44+
45+ fn trusted_auth_endpoint ( value : & str ) -> Result < reqwest:: Url > {
46+ let url = reqwest:: Url :: parse ( value) . context ( "parse xAI authentication endpoint" ) ?;
47+ if value. chars ( ) . any ( |character| character. is_ascii_control ( ) )
48+ || url. scheme ( ) != "https"
49+ || url. host_str ( ) != Some ( "auth.x.ai" )
50+ || url. port_or_known_default ( ) != Some ( 443 )
51+ || !url. username ( ) . is_empty ( )
52+ || url. password ( ) . is_some ( )
53+ || url. fragment ( ) . is_some ( )
54+ {
55+ return Err ( anyhow ! (
56+ "xAI discovery returned an untrusted authentication endpoint"
57+ ) ) ;
58+ }
59+ Ok ( url)
60+ }
61+
62+ fn discovery_endpoints ( document : DiscoveryDocument ) -> Result < OAuthEndpoints > {
63+ if document. issuer . trim_end_matches ( '/' ) != ISSUER {
64+ return Err ( anyhow ! (
65+ "xAI discovery issuer does not match the expected issuer"
66+ ) ) ;
67+ }
68+ Ok ( OAuthEndpoints {
69+ device_authorization : trusted_auth_endpoint ( & document. device_authorization_endpoint ) ?,
70+ token : trusted_auth_endpoint ( & document. token_endpoint ) ?,
71+ } )
72+ }
73+
74+ async fn discover_endpoints ( options : & SubscriptionHttpOptions ) -> Result < OAuthEndpoints > {
75+ let response = oauth_request ( http_client ( options) ?. get ( DISCOVERY_URL ) )
76+ . send ( )
77+ . await
78+ . context ( "fetch xAI OpenID discovery document" ) ?;
79+ if !response. status ( ) . is_success ( ) {
80+ return Err ( anyhow ! ( "xAI discovery failed: HTTP {}" , response. status( ) ) ) ;
81+ }
82+ discovery_endpoints (
83+ response
84+ . json ( )
85+ . await
86+ . context ( "parse xAI discovery document" ) ?,
87+ )
88+ }
89+
3390#[ derive( Debug , Deserialize ) ]
3491struct DeviceCodeResponse {
3592 device_code : String ,
@@ -141,9 +198,12 @@ fn validate_verification_url(url: &str) -> Result<()> {
141198 Ok ( ( ) )
142199}
143200
144- async fn request_device_code ( options : & SubscriptionHttpOptions ) -> Result < DeviceCodeResponse > {
201+ async fn request_device_code (
202+ options : & SubscriptionHttpOptions ,
203+ endpoints : & OAuthEndpoints ,
204+ ) -> Result < DeviceCodeResponse > {
145205 let client = http_client ( options) ?;
146- let response = oauth_request ( client. post ( DEVICE_AUTHORIZATION_URL ) )
206+ let response = oauth_request ( client. post ( endpoints . device_authorization . clone ( ) ) )
147207 . form ( & [
148208 ( "client_id" , CLIENT_ID ) ,
149209 ( "scope" , SCOPE ) ,
@@ -210,9 +270,10 @@ fn classify_device_poll_error(
210270async fn poll_once (
211271 device_code : & str ,
212272 options : & SubscriptionHttpOptions ,
273+ endpoints : & OAuthEndpoints ,
213274) -> Result < DevicePoll < TokenResponse > > {
214275 let client = http_client ( options) ?;
215- let response = oauth_request ( client. post ( TOKEN_URL ) )
276+ let response = oauth_request ( client. post ( endpoints . token . clone ( ) ) )
216277 . form ( & [
217278 ( "grant_type" , DEVICE_CODE_GRANT_TYPE ) ,
218279 ( "client_id" , CLIENT_ID ) ,
@@ -295,8 +356,9 @@ async fn persist_tokens(tokens: TokenResponse, expected_revision: u64) -> Result
295356}
296357
297358async fn refresh ( refresh_token : & str , options : & SubscriptionHttpOptions ) -> Result < TokenResponse > {
359+ let endpoints = discover_endpoints ( options) . await ?;
298360 let client = http_client ( options) ?;
299- let response = oauth_request ( client. post ( TOKEN_URL ) )
361+ let response = oauth_request ( client. post ( endpoints . token ) )
300362 . form ( & [
301363 ( "grant_type" , "refresh_token" ) ,
302364 ( "refresh_token" , refresh_token) ,
@@ -323,7 +385,8 @@ pub(crate) async fn begin_login(
323385 expected_revision : u64 ,
324386 options : SubscriptionHttpOptions ,
325387) -> Result < StartedLogin > {
326- let device = request_device_code ( & options) . await ?;
388+ let endpoints = discover_endpoints ( & options) . await ?;
389+ let device = request_device_code ( & options, & endpoints) . await ?;
327390 let interval = positive_seconds ( device. interval , DEFAULT_POLL_INTERVAL_SECS ) ;
328391 let expires_in = positive_seconds ( device. expires_in , DEFAULT_DEVICE_LIFETIME_SECS )
329392 . min ( super :: LOGIN_TIMEOUT . as_secs ( ) as i64 ) ;
@@ -344,7 +407,7 @@ pub(crate) async fn begin_login(
344407 Duration :: from_secs ( expires_in as u64 ) ,
345408 Duration :: from_secs ( 3 ) ,
346409 true ,
347- || poll_once ( & device_code, & options) ,
410+ || poll_once ( & device_code, & options, & endpoints ) ,
348411 )
349412 . await
350413 . context ( "complete xAI device authorization" )
@@ -485,6 +548,61 @@ mod tests {
485548 XAI_BASE_URL , XAI_REQUEST_URL ,
486549 } ;
487550
551+ #[ test]
552+ fn discovery_accepts_changed_paths_and_ignores_unrelated_metadata ( ) {
553+ let document = serde_json:: from_value ( serde_json:: json!( {
554+ "issuer" : "https://auth.x.ai/" ,
555+ "device_authorization_endpoint" : "https://auth.x.ai/new/device" ,
556+ "token_endpoint" : "https://auth.x.ai/new/token" ,
557+ "authorization_endpoint" : "https://auth.x.ai/authorize"
558+ } ) )
559+ . unwrap ( ) ;
560+ let endpoints = super :: discovery_endpoints ( document) . unwrap ( ) ;
561+ assert_eq ! (
562+ endpoints. device_authorization. as_str( ) ,
563+ "https://auth.x.ai/new/device"
564+ ) ;
565+ assert_eq ! ( endpoints. token. as_str( ) , "https://auth.x.ai/new/token" ) ;
566+ }
567+
568+ #[ test]
569+ fn discovery_rejects_wrong_issuer_missing_fields_and_untrusted_endpoints ( ) {
570+ for invalid in [
571+ "http://auth.x.ai/token" ,
572+ "https://auth.x.ai:8443/token" ,
573+ "https://user@auth.x.ai/token" ,
574+ "https://auth.x.ai.evil.test/token" ,
575+ "https://attacker.example/token" ,
576+ "https://auth.x.ai/token#fragment" ,
577+ "https://auth.x.ai/\n token" ,
578+ ] {
579+ for field in [ "token_endpoint" , "device_authorization_endpoint" ] {
580+ let mut value = serde_json:: json!( {
581+ "issuer" : "https://auth.x.ai" ,
582+ "device_authorization_endpoint" : "https://auth.x.ai/device" ,
583+ "token_endpoint" : "https://auth.x.ai/token"
584+ } ) ;
585+ value[ field] = serde_json:: json!( invalid) ;
586+ assert ! (
587+ super :: discovery_endpoints( serde_json:: from_value( value) . unwrap( ) ) . is_err( )
588+ ) ;
589+ }
590+ }
591+ let document = serde_json:: from_value ( serde_json:: json!( {
592+ "issuer" : "https://attacker.example" ,
593+ "device_authorization_endpoint" : "https://auth.x.ai/device" ,
594+ "token_endpoint" : "https://auth.x.ai/token"
595+ } ) )
596+ . unwrap ( ) ;
597+ assert ! ( super :: discovery_endpoints( document) . is_err( ) ) ;
598+ assert ! (
599+ serde_json:: from_value:: <super :: DiscoveryDocument >( serde_json:: json!( {
600+ "issuer" : "https://auth.x.ai" , "token_endpoint" : "https://auth.x.ai/token"
601+ } ) )
602+ . is_err( )
603+ ) ;
604+ }
605+
488606 #[ test]
489607 fn accepts_https_verification_urls_and_rejects_unsafe_urls ( ) {
490608 validate_verification_url ( "https://accounts.x.ai/oauth2/device?user_code=ABCD-EFGH" )
0 commit comments