1use std::collections::HashMap;
7use std::sync::atomic::AtomicU64;
8use std::sync::Arc;
9use std::time::{Duration, Instant};
10
11use serde::{Deserialize, Serialize};
12use serde_json::Value;
13use tokio::sync::Mutex;
14
15use crate::core::error::{McpError, McpResult};
16use crate::protocol::types::{JsonRpcRequest, JsonRpcResponse};
17use crate::transport::http::HttpClientTransport;
18use crate::transport::traits::{Transport, TransportConfig};
19
20#[derive(Debug, Clone, Serialize, Deserialize)]
26pub struct ServerInfo {
27 pub name: String,
29 pub version: String,
31 pub protocol_version: String,
33 pub capabilities: HashMap<String, Value>,
35 pub metadata: HashMap<String, Value>,
37}
38
39#[derive(Debug, Clone, Default)]
41pub struct ConnectionStats {
42 pub requests_sent: u64,
44 pub responses_received: u64,
46 pub request_failures: u64,
48 pub notifications_sent: u64,
50 pub notifications_received: u64,
52 pub uptime: Duration,
54 pub connected_at: Option<Instant>,
56 pub last_success_at: Option<Instant>,
58 pub last_error_at: Option<Instant>,
60 pub avg_response_time: Duration,
62 pub reconnect_attempts: u64,
64}
65
66#[derive(Debug, Clone)]
68pub struct HttpEndpoints {
69 pub mcp: String,
71 pub notify: String,
73 pub events: Option<String>,
75 pub health: String,
77}
78
79#[derive(Debug, Clone)]
81pub struct RetryConfig {
82 pub max_attempts: u32,
84 pub initial_delay: Duration,
86 pub max_delay: Duration,
88 pub backoff_multiplier: f64,
90 pub retry_on_timeout: bool,
92 pub retry_on_connection: bool,
94}
95
96impl Default for RetryConfig {
97 fn default() -> Self {
98 Self {
99 max_attempts: 3,
100 initial_delay: Duration::from_millis(100),
101 max_delay: Duration::from_secs(10),
102 backoff_multiplier: 2.0,
103 retry_on_timeout: true,
104 retry_on_connection: true,
105 }
106 }
107}
108
109#[derive(Debug, Clone, Default)]
111pub struct RetryPolicy {
112 pub default: RetryConfig,
114 pub method_specific: HashMap<String, RetryConfig>,
116}
117
118#[derive(Debug, Clone, Default)]
120pub struct TransportMetrics {
121 pub connection_stats: ConnectionStats,
123 pub performance: PerformanceMetrics,
125 pub errors: ErrorMetrics,
127}
128
129#[derive(Debug, Clone, Default)]
131pub struct PerformanceMetrics {
132 pub avg_latency: Duration,
134 pub p95_latency: Duration,
136 pub p99_latency: Duration,
138 pub requests_per_second: f64,
140 pub throughput_bps: f64,
142}
143
144#[derive(Debug, Clone, Default)]
146pub struct ErrorMetrics {
147 pub total_errors: u64,
149 pub timeout_errors: u64,
151 pub connection_errors: u64,
153 pub protocol_errors: u64,
155 pub http_errors: HashMap<u16, u64>,
157}
158
159#[allow(dead_code)]
165pub struct HttpClientTransportExtensions {
166 stats: Arc<Mutex<ConnectionStats>>,
168 request_counter: Arc<AtomicU64>,
170 retry_policy: Arc<Mutex<RetryPolicy>>,
172 request_logging: Arc<Mutex<bool>>,
174 last_error: Arc<Mutex<Option<McpError>>>,
176 response_times: Arc<Mutex<Vec<Duration>>>,
178}
179
180impl Default for HttpClientTransportExtensions {
181 fn default() -> Self {
182 Self {
183 stats: Arc::new(Mutex::new(ConnectionStats::default())),
184 request_counter: Arc::new(AtomicU64::new(0)),
185 retry_policy: Arc::new(Mutex::new(RetryPolicy::default())),
186 request_logging: Arc::new(Mutex::new(false)),
187 last_error: Arc::new(Mutex::new(None)),
188 response_times: Arc::new(Mutex::new(Vec::new())),
189 }
190 }
191}
192
193impl HttpClientTransport {
195 pub async fn ping(&mut self) -> McpResult<Duration> {
201 let start = Instant::now();
202
203 let health_request = JsonRpcRequest {
204 jsonrpc: "2.0".to_string(),
205 method: "ping".to_string(),
206 params: Some(Value::Object(serde_json::Map::new())),
207 id: Value::from(self.next_request_id().await),
208 };
209
210 match self.send_request(health_request).await {
211 Ok(_) => {
212 let duration = start.elapsed();
213 Ok(duration)
214 }
215 Err(_e) => {
216 let url = format!("{}/health", self.base_url);
218 let _response = self
219 .client
220 .get(&url)
221 .send()
222 .await
223 .map_err(|e| McpError::Http(format!("Health check failed: {e}")))?
224 .error_for_status()
225 .map_err(|e| McpError::Http(format!("Health check failed: {e}")))?;
226
227 let duration = start.elapsed();
228 Ok(duration)
229 }
230 }
231 }
232
233 pub async fn get_server_info(&mut self) -> McpResult<ServerInfo> {
235 let request = JsonRpcRequest {
236 jsonrpc: "2.0".to_string(),
237 method: "initialize".to_string(),
238 params: Some(serde_json::json!({
239 "protocolVersion": "2025-11-25",
240 "capabilities": {},
241 "clientInfo": {
242 "name": "prism-mcp-rs",
243 "version": env!("CARGO_PKG_VERSION")
244 }
245 })),
246 id: Value::from(self.next_request_id().await),
247 };
248
249 let response = self.send_request(request).await?;
250
251 if let Some(result) = response.result {
252 let server_info = ServerInfo {
253 name: result
254 .get("serverInfo")
255 .and_then(|info| info.get("name"))
256 .and_then(|name| name.as_str())
257 .unwrap_or("Unknown")
258 .to_string(),
259 version: result
260 .get("serverInfo")
261 .and_then(|info| info.get("version"))
262 .and_then(|version| version.as_str())
263 .unwrap_or("Unknown")
264 .to_string(),
265 protocol_version: result
266 .get("protocolVersion")
267 .and_then(|version| version.as_str())
268 .unwrap_or("Unknown")
269 .to_string(),
270 capabilities: result
271 .get("capabilities")
272 .and_then(|caps| caps.as_object())
273 .map(|obj| obj.iter().map(|(k, v)| (k.clone(), v.clone())).collect())
274 .unwrap_or_default(),
275 metadata: result
276 .as_object()
277 .map(|obj| obj.iter().map(|(k, v)| (k.clone(), v.clone())).collect())
278 .unwrap_or_default(),
279 };
280 Ok(server_info)
281 } else {
282 Err(McpError::Protocol(
283 "Invalid server info response".to_string(),
284 ))
285 }
286 }
287
288 pub async fn get_connection_stats(&self) -> ConnectionStats {
290 ConnectionStats {
293 requests_sent: 0, responses_received: 0,
295 request_failures: 0,
296 notifications_sent: 0,
297 notifications_received: 0,
298 uptime: Duration::from_secs(0),
299 connected_at: Some(Instant::now()),
300 last_success_at: None,
301 last_error_at: None,
302 avg_response_time: Duration::from_millis(0),
303 reconnect_attempts: 0,
304 }
305 }
306
307 pub fn is_healthy(&self) -> bool {
309 self.is_connected()
310 }
311
312 pub async fn call_method<T: Serialize, R: for<'de> Deserialize<'de>>(
318 &mut self,
319 method: &str,
320 params: T,
321 ) -> McpResult<R> {
322 let request =
323 JsonRpcRequest {
324 jsonrpc: "2.0".to_string(),
325 method: method.to_string(),
326 params: Some(serde_json::to_value(params).map_err(|e| {
327 McpError::Protocol(format!("Failed to serialize parameters: {e}"))
328 })?),
329 id: Value::from(self.next_request_id().await),
330 };
331
332 let response = self.send_request(request).await?;
333
334 if let Some(result) = response.result {
335 serde_json::from_value(result)
336 .map_err(|e| McpError::Protocol(format!("Failed to deserialize response: {e}")))
337 } else {
338 Err(McpError::Protocol("Missing result in response".to_string()))
339 }
340 }
341
342 pub async fn call_method_simple(&mut self, method: &str) -> McpResult<Value> {
344 let request = JsonRpcRequest {
345 jsonrpc: "2.0".to_string(),
346 method: method.to_string(),
347 params: None,
348 id: Value::from(self.next_request_id().await),
349 };
350
351 let response = self.send_request(request).await?;
352 response
353 .result
354 .ok_or_else(|| McpError::Protocol("Missing result in response".to_string()))
355 }
356
357 pub async fn batch_requests(
359 &mut self,
360 requests: Vec<JsonRpcRequest>,
361 ) -> McpResult<Vec<JsonRpcResponse>> {
362 let mut responses = Vec::with_capacity(requests.len());
365
366 for request in requests {
367 let response = self.send_request(request).await?;
368 responses.push(response);
369 }
370
371 Ok(responses)
372 }
373
374 pub async fn reconnect(&mut self) -> McpResult<()> {
380 self.close().await?;
382
383 let new_transport =
385 Self::with_config(&self.base_url, self.sse_url.as_ref(), self.config.clone()).await?;
386
387 *self = new_transport;
389
390 Ok(())
391 }
392
393 pub async fn test_connection(&self) -> McpResult<bool> {
395 let url = format!("{}/health", self.base_url);
396 match self.client.get(&url).send().await {
397 Ok(response) => Ok(response.status().is_success()),
398 Err(_) => Ok(false),
399 }
400 }
401
402 pub fn update_headers(&mut self, new_headers: HashMap<String, String>) {
408 for (key, value) in new_headers {
409 if let (Ok(header_name), Ok(header_value)) = (
410 key.parse::<axum::http::HeaderName>(),
411 value.parse::<axum::http::HeaderValue>(),
412 ) {
413 self.headers.insert(header_name, header_value);
414 }
415 }
416 }
417
418 pub fn set_timeout(&mut self, timeout_ms: u64) {
420 self.config.read_timeout_ms = Some(timeout_ms);
421 self.config.write_timeout_ms = Some(timeout_ms);
422 }
423
424 pub fn get_config(&self) -> &TransportConfig {
426 &self.config
427 }
428
429 pub fn get_base_url(&self) -> &str {
435 &self.base_url
436 }
437
438 pub fn get_sse_url(&self) -> Option<&str> {
440 self.sse_url.as_deref()
441 }
442
443 pub fn get_endpoints(&self) -> HttpEndpoints {
445 HttpEndpoints {
446 mcp: format!("{}/mcp", self.base_url),
447 notify: format!("{}/mcp/notify", self.base_url),
448 events: self.sse_url.clone(),
449 health: format!("{}/health", self.base_url),
450 }
451 }
452
453 pub async fn call_with_retry<T: Serialize + Clone, R: for<'de> Deserialize<'de>>(
459 &mut self,
460 method: &str,
461 params: T,
462 retry_config: RetryConfig,
463 ) -> McpResult<R> {
464 let mut last_error = None;
465 let mut delay = retry_config.initial_delay;
466
467 for attempt in 0..=retry_config.max_attempts {
468 match self.call_method(method, params.clone()).await {
469 Ok(result) => return Ok(result),
470 Err(e) => {
471 last_error = Some(e.clone());
472
473 if attempt == retry_config.max_attempts {
475 break;
476 }
477
478 let should_retry = match &e {
480 McpError::Timeout(_) => retry_config.retry_on_timeout,
481 McpError::Connection(_) => retry_config.retry_on_connection,
482 McpError::Http(_) => retry_config.retry_on_connection,
483 _ => false,
484 };
485
486 if !should_retry {
487 break;
488 }
489
490 tokio::time::sleep(delay).await;
492
493 delay = std::cmp::min(
495 Duration::from_millis(
496 (delay.as_millis() as f64 * retry_config.backoff_multiplier) as u64,
497 ),
498 retry_config.max_delay,
499 );
500 }
501 }
502 }
503
504 Err(last_error
505 .unwrap_or_else(|| McpError::Protocol("Retry failed without error".to_string())))
506 }
507
508 pub fn set_retry_policy(&mut self, _policy: RetryPolicy) {
510 }
513
514 pub fn enable_request_logging(&mut self, _enabled: bool) {
520 }
523
524 pub fn get_last_error(&self) -> Option<&McpError> {
526 None
529 }
530
531 pub async fn export_metrics(&self) -> McpResult<TransportMetrics> {
533 Ok(TransportMetrics {
534 connection_stats: self.get_connection_stats().await,
535 performance: PerformanceMetrics::default(),
536 errors: ErrorMetrics::default(),
537 })
538 }
539}
540
541pub struct HttpClientTransportBuilder {
547 base_url: Option<String>,
548 sse_url: Option<String>,
549 config: TransportConfig,
550}
551
552impl HttpClientTransport {
553 pub fn builder() -> HttpClientTransportBuilder {
555 HttpClientTransportBuilder {
556 base_url: None,
557 sse_url: None,
558 config: TransportConfig::default(),
559 }
560 }
561}
562
563impl HttpClientTransportBuilder {
564 pub fn base_url<S: Into<String>>(mut self, url: S) -> Self {
566 self.base_url = Some(url.into());
567 self
568 }
569
570 pub fn sse_url<S: Into<String>>(mut self, url: S) -> Self {
572 self.sse_url = Some(url.into());
573 self
574 }
575
576 pub fn timeout(mut self, ms: u64) -> Self {
578 self.config.read_timeout_ms = Some(ms);
579 self.config.write_timeout_ms = Some(ms);
580 self
581 }
582
583 pub fn header<S: Into<String>>(mut self, key: S, value: S) -> Self {
585 self.config.headers.insert(key.into(), value.into());
586 self
587 }
588
589 pub fn compression(mut self, enabled: bool) -> Self {
591 self.config.compression = enabled;
592 self
593 }
594
595 pub fn connect_timeout(mut self, ms: u64) -> Self {
597 self.config.connect_timeout_ms = Some(ms);
598 self
599 }
600
601 pub fn max_message_size(mut self, size: usize) -> Self {
603 self.config.max_message_size = Some(size);
604 self
605 }
606
607 pub async fn build(self) -> McpResult<HttpClientTransport> {
609 let base_url = self
610 .base_url
611 .ok_or_else(|| McpError::protocol("Base URL is required"))?;
612
613 HttpClientTransport::with_config(base_url, self.sse_url.clone(), self.config).await
614 }
615}
616
617#[cfg(test)]
622mod tests {
623 use super::*;
624
625 #[tokio::test]
626 async fn test_builder_pattern() {
627 let result = HttpClientTransport::builder()
628 .base_url("http://localhost:3000")
629 .sse_url("http://localhost:3000/events")
630 .timeout(30_000)
631 .header("Authorization", "Bearer token")
632 .compression(true)
633 .build()
634 .await;
635
636 assert!(result.is_ok());
637 let transport = result.unwrap();
638 assert_eq!(transport.get_base_url(), "http://localhost:3000");
639 assert_eq!(
640 transport.get_sse_url(),
641 Some("http://localhost:3000/events")
642 );
643 }
644
645 #[test]
646 fn test_retry_config_default() {
647 let config = RetryConfig::default();
648 assert_eq!(config.max_attempts, 3);
649 assert_eq!(config.backoff_multiplier, 2.0);
650 assert!(config.retry_on_timeout);
651 assert!(config.retry_on_connection);
652 }
653
654 #[test]
655 fn test_http_endpoints() {
656 let base_url = "http://localhost:3000";
657 let endpoints = HttpEndpoints {
658 mcp: format!("{base_url}/mcp"),
659 notify: format!("{base_url}/mcp/notify"),
660 events: Some("http://localhost:3000/events".to_string()),
661 health: format!("{base_url}/health"),
662 };
663
664 assert_eq!(endpoints.mcp, "http://localhost:3000/mcp");
665 assert_eq!(endpoints.notify, "http://localhost:3000/mcp/notify");
666 assert_eq!(endpoints.health, "http://localhost:3000/health");
667 }
668
669 #[test]
670 fn test_connection_stats_default() {
671 let stats = ConnectionStats::default();
672 assert_eq!(stats.requests_sent, 0);
673 assert_eq!(stats.responses_received, 0);
674 assert_eq!(stats.uptime, Duration::from_secs(0));
675 }
676}