prism_mcp_rs/transport/
http_auth.rs1use async_trait::async_trait;
7use std::sync::Arc;
8
9use crate::auth::{AuthConfig, AuthorizationClient};
10use crate::core::error::{McpError, McpResult};
11use crate::protocol::types::{JsonRpcNotification, JsonRpcRequest, JsonRpcResponse};
12use crate::transport::{http::HttpClientTransport, Transport};
13
14pub struct AuthorizedHttpTransport {
16 inner: HttpClientTransport,
18 auth_client: Arc<AuthorizationClient>,
20 auth_enabled: bool,
22}
23
24impl AuthorizedHttpTransport {
25 pub async fn new(
27 base_url: String,
28 sse_url: Option<String>,
29 auth_config: AuthConfig,
30 ) -> McpResult<Self> {
31 let auth_enabled = auth_config.enabled;
32 let inner = HttpClientTransport::new(&base_url, sse_url.as_ref()).await?;
33 let auth_client = Arc::new(AuthorizationClient::new(auth_config, base_url.clone()));
34
35 Ok(Self {
36 inner,
37 auth_client,
38 auth_enabled,
39 })
40 }
41
42 async fn send_with_auth(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
44 if !self.auth_enabled {
46 return self.inner.send_request(request).await;
47 }
48
49 let token = self.auth_client.get_token().await?;
51
52 crate::auth::client::add_auth_header(&mut self.inner.headers, &token);
54
55 match self.inner.send_request(request.clone()).await {
57 Ok(response) => Ok(response),
58 Err(McpError::Http(msg)) if msg.contains("401") || msg.contains("403") => {
59 let fresh_token = self.auth_client.token_manager().refresh_token().await?;
61
62 crate::auth::client::add_auth_header(&mut self.inner.headers, &fresh_token);
64
65 self.inner.send_request(request).await
67 }
68 Err(e) => Err(e),
69 }
70 }
71
72 pub async fn handle_unauthorized(&self, www_authenticate: &str) -> McpResult<String> {
74 self.auth_client.handle_unauthorized(www_authenticate).await
75 }
76
77 pub async fn handle_callback(&self, callback_url: &str) -> McpResult<String> {
79 self.auth_client.handle_callback(callback_url).await
80 }
81
82 pub async fn is_authenticated(&self) -> bool {
84 if !self.auth_enabled {
85 return true; }
87 self.auth_client.is_authenticated().await
88 }
89
90 pub async fn get_authorization_url(&self) -> McpResult<String> {
92 if !self.auth_enabled {
93 return Err(McpError::Auth("Authorization is not enabled".to_string()));
94 }
95
96 let context = self.auth_client.token_manager().get_context().await;
99
100 if let Some(auth_metadata) = context.auth_server_metadata {
101 self.auth_client
102 .start_authorization_flow(&auth_metadata)
103 .await
104 } else {
105 Err(McpError::Auth(
106 "Authorization server not discovered yet".to_string(),
107 ))
108 }
109 }
110
111 pub async fn logout(&self) {
113 self.auth_client.logout().await;
114 }
115}
116
117#[async_trait]
118impl Transport for AuthorizedHttpTransport {
119 async fn send_request(&mut self, request: JsonRpcRequest) -> McpResult<JsonRpcResponse> {
120 self.send_with_auth(request).await
121 }
122
123 async fn send_notification(&mut self, notification: JsonRpcNotification) -> McpResult<()> {
124 if self.auth_enabled {
125 let token = self.auth_client.get_token().await?;
127 crate::auth::client::add_auth_header(&mut self.inner.headers, &token);
128 }
129
130 self.inner.send_notification(notification).await
131 }
132
133 async fn receive_notification(&mut self) -> McpResult<Option<JsonRpcNotification>> {
134 self.inner.receive_notification().await
135 }
136
137 async fn close(&mut self) -> McpResult<()> {
138 self.inner.close().await
139 }
140}
141
142pub struct AuthorizedHttpTransportBuilder {
144 base_url: String,
145 sse_url: Option<String>,
146 auth_config: AuthConfig,
147}
148
149impl AuthorizedHttpTransportBuilder {
150 pub fn new(base_url: String) -> Self {
152 Self {
153 base_url,
154 sse_url: None,
155 auth_config: AuthConfig::default(),
156 }
157 }
158
159 pub fn with_sse(mut self, url: String) -> Self {
161 self.sse_url = Some(url);
162 self
163 }
164
165 pub fn with_auth(mut self, enabled: bool) -> Self {
167 self.auth_config.enabled = enabled;
168 self
169 }
170
171 pub fn with_client_credentials(
173 mut self,
174 client_id: String,
175 client_secret: Option<String>,
176 ) -> Self {
177 self.auth_config.client_id = Some(client_id);
178 self.auth_config.client_secret = client_secret;
179 self
180 }
181
182 pub fn with_redirect_uri(mut self, uri: String) -> Self {
184 self.auth_config.redirect_uri = uri;
185 self
186 }
187
188 pub fn with_scopes(mut self, scopes: Vec<String>) -> Self {
190 self.auth_config.scopes = scopes;
191 self
192 }
193
194 pub fn with_dynamic_registration(mut self, enabled: bool) -> Self {
196 self.auth_config.enable_dynamic_registration = enabled;
197 self
198 }
199
200 pub async fn build(self) -> McpResult<AuthorizedHttpTransport> {
202 AuthorizedHttpTransport::new(self.base_url, self.sse_url, self.auth_config).await
203 }
204}
205
206#[cfg(test)]
207mod tests {
208 use super::*;
209
210 #[tokio::test]
211 async fn test_authorized_transport_creation() {
212 let transport = AuthorizedHttpTransportBuilder::new("https://mcp.example.com".to_string())
213 .with_auth(false) .build()
215 .await;
216
217 assert!(transport.is_ok());
218 let transport = transport.unwrap();
219 assert!(transport.is_authenticated().await);
220 }
221
222 #[tokio::test]
223 async fn test_auth_disabled_passthrough() {
224 let transport = AuthorizedHttpTransportBuilder::new("https://mcp.example.com".to_string())
225 .with_auth(false)
226 .build()
227 .await
228 .unwrap();
229
230 assert!(transport.is_authenticated().await);
232
233 transport.logout().await;
235 }
236}