Skip to main content

prism_mcp_rs/transport/
http_convenience.rs

1//! Production-grade convenience methods for HttpClientTransport
2//!
3//! Module extends the basic HttpClientTransport with high-level convenience
4//! methods expected in a production SDK.
5
6use 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// ============================================================================
21// Additional Types for Convenience Methods
22// ============================================================================
23
24/// Server information returned by get_server_info
25#[derive(Debug, Clone, Serialize, Deserialize)]
26pub struct ServerInfo {
27    /// Server name
28    pub name: String,
29    /// Server version
30    pub version: String,
31    /// Supported protocol version
32    pub protocol_version: String,
33    /// Server capabilities
34    pub capabilities: HashMap<String, Value>,
35    /// Additional server metadata
36    pub metadata: HashMap<String, Value>,
37}
38
39/// Connection statistics for monitoring
40#[derive(Debug, Clone, Default)]
41pub struct ConnectionStats {
42    /// Number of requests sent
43    pub requests_sent: u64,
44    /// Number of successful responses
45    pub responses_received: u64,
46    /// Number of failed requests
47    pub request_failures: u64,
48    /// Number of notifications sent
49    pub notifications_sent: u64,
50    /// Number of notifications received
51    pub notifications_received: u64,
52    /// Total connection uptime
53    pub uptime: Duration,
54    /// Connection start time
55    pub connected_at: Option<Instant>,
56    /// Last successful request time
57    pub last_success_at: Option<Instant>,
58    /// Last error time
59    pub last_error_at: Option<Instant>,
60    /// Average response time
61    pub avg_response_time: Duration,
62    /// Number of reconnection attempts
63    pub reconnect_attempts: u64,
64}
65
66/// HTTP endpoint URLs
67#[derive(Debug, Clone)]
68pub struct HttpEndpoints {
69    /// Main MCP endpoint
70    pub mcp: String,
71    /// Notification endpoint
72    pub notify: String,
73    /// SSE events endpoint
74    pub events: Option<String>,
75    /// Health check endpoint
76    pub health: String,
77}
78
79/// Retry configuration for resilient requests
80#[derive(Debug, Clone)]
81pub struct RetryConfig {
82    /// Maximum number of retry attempts
83    pub max_attempts: u32,
84    /// Initial delay between retries
85    pub initial_delay: Duration,
86    /// Maximum delay between retries
87    pub max_delay: Duration,
88    /// Exponential backoff multiplier
89    pub backoff_multiplier: f64,
90    /// Whether to retry on timeout errors
91    pub retry_on_timeout: bool,
92    /// Whether to retry on connection errors
93    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/// Retry policy for automatic retries
110#[derive(Debug, Clone, Default)]
111pub struct RetryPolicy {
112    /// Default retry configuration
113    pub default: RetryConfig,
114    /// Method-specific retry configurations
115    pub method_specific: HashMap<String, RetryConfig>,
116}
117
118/// Transport metrics for observability
119#[derive(Debug, Clone, Default)]
120pub struct TransportMetrics {
121    /// Connection statistics
122    pub connection_stats: ConnectionStats,
123    /// Performance metrics
124    pub performance: PerformanceMetrics,
125    /// Error metrics
126    pub errors: ErrorMetrics,
127}
128
129/// Performance metrics
130#[derive(Debug, Clone, Default)]
131pub struct PerformanceMetrics {
132    /// Average request latency
133    pub avg_latency: Duration,
134    /// 95th percentile latency
135    pub p95_latency: Duration,
136    /// 99th percentile latency
137    pub p99_latency: Duration,
138    /// Requests per second
139    pub requests_per_second: f64,
140    /// Throughput in bytes per second
141    pub throughput_bps: f64,
142}
143
144/// Error metrics
145#[derive(Debug, Clone, Default)]
146pub struct ErrorMetrics {
147    /// Total error count
148    pub total_errors: u64,
149    /// Timeout errors
150    pub timeout_errors: u64,
151    /// Connection errors
152    pub connection_errors: u64,
153    /// Protocol errors
154    pub protocol_errors: u64,
155    /// HTTP errors by status code
156    pub http_errors: HashMap<u16, u64>,
157}
158
159// ============================================================================
160// improved HttpClientTransport with Convenience Methods
161// ============================================================================
162
163/// Extended functionality for HttpClientTransport
164#[allow(dead_code)]
165pub struct HttpClientTransportExtensions {
166    /// Connection statistics
167    stats: Arc<Mutex<ConnectionStats>>,
168    /// Request ID counter
169    request_counter: Arc<AtomicU64>,
170    /// Retry policy
171    retry_policy: Arc<Mutex<RetryPolicy>>,
172    /// Request logging enabled
173    request_logging: Arc<Mutex<bool>>,
174    /// Last error
175    last_error: Arc<Mutex<Option<McpError>>>,
176    /// Response time tracker
177    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
193/// Production-grade convenience methods for HttpClientTransport
194impl HttpClientTransport {
195    // ============================================================================
196    // 1. Connection Health & Status
197    // ============================================================================
198
199    /// Send a health check request and measure response time
200    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                // Try the health endpoint as fallback
217                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    /// Retrieve server capabilities and version information
234    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    /// Get connection statistics for monitoring
289    pub async fn get_connection_stats(&self) -> ConnectionStats {
290        // This would require extending HttpClientTransport with statistics tracking
291        // For now, return basic stats
292        ConnectionStats {
293            requests_sent: 0, // Would be tracked in actual implementation
294            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    /// Quick health check based on recent activity
308    pub fn is_healthy(&self) -> bool {
309        self.is_connected()
310    }
311
312    // ============================================================================
313    // 2. Type-Safe Request Builders & Shortcuts
314    // ============================================================================
315
316    /// Type-safe method calling with automatic serialization
317    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    /// Simple method call without parameters
343    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    /// Send multiple requests efficiently
358    pub async fn batch_requests(
359        &mut self,
360        requests: Vec<JsonRpcRequest>,
361    ) -> McpResult<Vec<JsonRpcResponse>> {
362        // For HTTP transport, we send requests sequentially
363        // A more complete implementation could use HTTP/2 multiplexing
364        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    // ============================================================================
375    // 3. Connection Management
376    // ============================================================================
377
378    /// Reconnect with same configuration
379    pub async fn reconnect(&mut self) -> McpResult<()> {
380        // Close current connection
381        self.close().await?;
382
383        // Create new transport with same configuration
384        let new_transport =
385            Self::with_config(&self.base_url, self.sse_url.as_ref(), self.config.clone()).await?;
386
387        // Replace current transport state
388        *self = new_transport;
389
390        Ok(())
391    }
392
393    /// Test if connection is working without modifying state
394    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    // ============================================================================
403    // 4. Configuration Updates
404    // ============================================================================
405
406    /// Update headers without recreating transport
407    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    /// Update timeout configuration
419    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    /// Access current configuration
425    pub fn get_config(&self) -> &TransportConfig {
426        &self.config
427    }
428
429    // ============================================================================
430    // 5. URL and Endpoint Management
431    // ============================================================================
432
433    /// Get current base URL
434    pub fn get_base_url(&self) -> &str {
435        &self.base_url
436    }
437
438    /// Get current SSE URL
439    pub fn get_sse_url(&self) -> Option<&str> {
440        self.sse_url.as_deref()
441    }
442
443    /// Get all endpoint URLs
444    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    // ============================================================================
454    // 6. Retry and Resilience
455    // ============================================================================
456
457    /// Method call with automatic retry logic
458    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                    // Don't retry on the last attempt
474                    if attempt == retry_config.max_attempts {
475                        break;
476                    }
477
478                    // Check if we should retry this error
479                    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                    // Wait before retry
491                    tokio::time::sleep(delay).await;
492
493                    // Exponential backoff
494                    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    /// Set retry policy for automatic retries (would require state extension)
509    pub fn set_retry_policy(&mut self, _policy: RetryPolicy) {
510        // Implementation would require extending HttpClientTransport with retry state
511        // For now, this is a placeholder
512    }
513
514    // ============================================================================
515    // 7. Debugging and Observability
516    // ============================================================================
517
518    /// Enable/disable request/response logging (placeholder)
519    pub fn enable_request_logging(&mut self, _enabled: bool) {
520        // Implementation would require extending HttpClientTransport with logging state
521        // For now, this is a placeholder
522    }
523
524    /// Get the last error that occurred (placeholder)
525    pub fn get_last_error(&self) -> Option<&McpError> {
526        // Implementation would require extending HttpClientTransport with error tracking
527        // For now, this is a placeholder
528        None
529    }
530
531    /// Export detailed metrics for monitoring (placeholder)
532    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
541// ============================================================================
542// 6. Builder Pattern Support
543// ============================================================================
544
545/// Fluent builder for HttpClientTransport configuration
546pub struct HttpClientTransportBuilder {
547    base_url: Option<String>,
548    sse_url: Option<String>,
549    config: TransportConfig,
550}
551
552impl HttpClientTransport {
553    /// Create a new builder for HttpClientTransport
554    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    /// Set the base URL
565    pub fn base_url<S: Into<String>>(mut self, url: S) -> Self {
566        self.base_url = Some(url.into());
567        self
568    }
569
570    /// Set the SSE URL for notifications
571    pub fn sse_url<S: Into<String>>(mut self, url: S) -> Self {
572        self.sse_url = Some(url.into());
573        self
574    }
575
576    /// Set request timeout
577    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    /// Add a custom header
584    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    /// Enable or disable compression
590    pub fn compression(mut self, enabled: bool) -> Self {
591        self.config.compression = enabled;
592        self
593    }
594
595    /// Set connection timeout
596    pub fn connect_timeout(mut self, ms: u64) -> Self {
597        self.config.connect_timeout_ms = Some(ms);
598        self
599    }
600
601    /// Set maximum message size
602    pub fn max_message_size(mut self, size: usize) -> Self {
603        self.config.max_message_size = Some(size);
604        self
605    }
606
607    /// Build the HttpClientTransport
608    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// ============================================================================
618// Tests
619// ============================================================================
620
621#[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}