diff --git a/lib/crates/fabro-openai-oauth/src/lib.rs b/lib/crates/fabro-openai-oauth/src/lib.rs index 542c718e9..0ed1eccad 100644 --- a/lib/crates/fabro-openai-oauth/src/lib.rs +++ b/lib/crates/fabro-openai-oauth/src/lib.rs @@ -360,14 +360,16 @@ pub async fn poll_device_flow( #[derive(Deserialize)] struct CallbackParams { - code: String, + code: Option, state: String, + error: Option, + error_description: Option, } pub async fn start_callback_server( port: u16, expected_state: String, -) -> Result<(u16, tokio::sync::oneshot::Receiver), String> { +) -> Result<(u16, tokio::sync::oneshot::Receiver>), String> { let listener = tokio::net::TcpListener::bind(format!("localhost:{port}")) .await .map_err(|e| format!("Failed to bind callback server: {e}"))?; @@ -376,7 +378,7 @@ pub async fn start_callback_server( .map_err(|e| format!("Failed to get local address: {e}"))? .port(); - let (code_tx, code_rx) = tokio::sync::oneshot::channel::(); + let (code_tx, code_rx) = tokio::sync::oneshot::channel::>(); let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>(); let code_tx = std::sync::Arc::new(std::sync::Mutex::new(Some(code_tx))); @@ -393,8 +395,63 @@ pub async fn start_callback_server( axum::response::Html("State mismatch".to_string()), ); } + + if let Some(error) = params.error { + let desc = params + .error_description + .unwrap_or_else(|| error.clone()); + if let Some(tx) = code_tx.lock().unwrap().take() { + let _ = tx.send(Err(desc.clone())); + } + if let Some(tx) = shutdown_tx.lock().unwrap().take() { + let _ = tx.send(()); + } + return ( + axum::http::StatusCode::BAD_REQUEST, + axum::response::Html(format!( + r#" + + + +Authorization Failed + + + +
+
✗
+

Authorization Failed

+

{desc}

+
+ +"# + )), + ); + } + + let code = match params.code { + Some(c) => c, + None => { + if let Some(tx) = code_tx.lock().unwrap().take() { + let _ = tx.send(Err("No authorization code received".to_string())); + } + if let Some(tx) = shutdown_tx.lock().unwrap().take() { + let _ = tx.send(()); + } + return ( + axum::http::StatusCode::BAD_REQUEST, + axum::response::Html("No authorization code received".to_string()), + ); + } + }; + if let Some(tx) = code_tx.lock().unwrap().take() { - let _ = tx.send(params.code); + let _ = tx.send(Ok(code)); } if let Some(tx) = shutdown_tx.lock().unwrap().take() { let _ = tx.send(()); @@ -462,7 +519,8 @@ pub async fn run_browser_flow(issuer: &str, client_id: &str) -> Result