chore: properly identify when to try oauth (#4918)
This commit is contained in:
@@ -4,9 +4,12 @@ use chrono::{DateTime, Utc};
|
|||||||
use futures::stream::{FuturesUnordered, StreamExt};
|
use futures::stream::{FuturesUnordered, StreamExt};
|
||||||
use futures::{future, FutureExt};
|
use futures::{future, FutureExt};
|
||||||
use rmcp::service::ClientInitializeError;
|
use rmcp::service::ClientInitializeError;
|
||||||
use rmcp::transport::streamable_http_client::StreamableHttpClientTransportConfig;
|
use rmcp::transport::streamable_http_client::{
|
||||||
|
AuthRequiredError, StreamableHttpClientTransportConfig, StreamableHttpError,
|
||||||
|
};
|
||||||
use rmcp::transport::{
|
use rmcp::transport::{
|
||||||
ConfigureCommandExt, SseClientTransport, StreamableHttpClientTransport, TokioChildProcess,
|
ConfigureCommandExt, DynamicTransportError, SseClientTransport, StreamableHttpClientTransport,
|
||||||
|
TokioChildProcess,
|
||||||
};
|
};
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::process::Stdio;
|
use std::process::Stdio;
|
||||||
@@ -205,6 +208,28 @@ async fn child_process_client(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn extract_auth_error(
|
||||||
|
res: &Result<McpClient, ClientInitializeError>,
|
||||||
|
) -> Option<&AuthRequiredError> {
|
||||||
|
match res {
|
||||||
|
Ok(_) => None,
|
||||||
|
Err(err) => match err {
|
||||||
|
ClientInitializeError::TransportError {
|
||||||
|
error: DynamicTransportError { error, .. },
|
||||||
|
..
|
||||||
|
} => error
|
||||||
|
.downcast_ref::<StreamableHttpError<reqwest::Error>>()
|
||||||
|
.and_then(|auth_error| match auth_error {
|
||||||
|
StreamableHttpError::AuthRequired(auth_required_error) => {
|
||||||
|
Some(auth_required_error)
|
||||||
|
}
|
||||||
|
_ => None,
|
||||||
|
}),
|
||||||
|
_ => None,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl ExtensionManager {
|
impl ExtensionManager {
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self {
|
Self {
|
||||||
@@ -340,15 +365,10 @@ impl ExtensionManager {
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
let client = if let Err(e) = client_res {
|
let client = if let Some(_auth_error) = extract_auth_error(&client_res) {
|
||||||
// make an attempt at oauth, but failing that, return the original error,
|
let am = oauth_flow(uri, name)
|
||||||
// because this might not have been an auth error at all.
|
.await
|
||||||
// TODO: when rmcp supports it, we should trigger this flow on 401s with
|
.map_err(|_| ExtensionError::SetupError("auth error".to_string()))?;
|
||||||
// WWW-Authenticate headers, not just any init error
|
|
||||||
let am = match oauth_flow(uri, name).await {
|
|
||||||
Ok(am) => am,
|
|
||||||
Err(_) => return Err(e.into()),
|
|
||||||
};
|
|
||||||
let client = AuthClient::new(reqwest::Client::default(), am);
|
let client = AuthClient::new(reqwest::Client::default(), am);
|
||||||
let transport = StreamableHttpClientTransport::with_client(
|
let transport = StreamableHttpClientTransport::with_client(
|
||||||
client,
|
client,
|
||||||
|
|||||||
Reference in New Issue
Block a user