use crate::{Error, Result};
use pforge_config::TransportType;
#[cfg(feature = "sse")]
#[allow(deprecated)]
use pmcp::shared::{OptimizedSseConfig, OptimizedSseTransport};
use pmcp::shared::{StdioTransport, Transport};
#[cfg(feature = "websocket")]
use pmcp::shared::{WebSocketConfig, WebSocketTransport};
#[cfg(any(feature = "sse", feature = "websocket"))]
use std::time::Duration;
pub fn create_transport(transport_type: &TransportType) -> Result<Box<dyn Transport>> {
match transport_type {
TransportType::Stdio => {
let transport = StdioTransport::new();
Ok(Box::new(transport))
}
TransportType::Sse => create_sse_transport(),
TransportType::WebSocket => create_websocket_transport(),
}
}
#[allow(clippy::unnecessary_wraps)]
#[cfg(feature = "sse")]
fn create_sse_transport() -> Result<Box<dyn Transport>> {
let config = OptimizedSseConfig {
url: "http://localhost:8080/sse".to_string(),
connection_timeout: Duration::from_secs(30),
keepalive_interval: Duration::from_secs(15),
max_reconnects: 5,
reconnect_delay: Duration::from_secs(1),
buffer_size: 100,
flush_interval: Duration::from_millis(100),
enable_pooling: true,
max_connections: 10,
enable_compression: false,
};
#[allow(deprecated)]
let transport = OptimizedSseTransport::new(config);
Ok(Box::new(transport))
}
#[cfg(not(feature = "sse"))]
fn create_sse_transport() -> Result<Box<dyn Transport>> {
Err(Error::feature_disabled("sse", "transport `sse`"))
}
#[cfg(feature = "websocket")]
fn create_websocket_transport() -> Result<Box<dyn Transport>> {
let url = "ws://localhost:8080/ws"
.parse()
.map_err(|e| Error::Handler(format!("Invalid WebSocket URL: {}", e)))?;
let config = WebSocketConfig {
url,
auto_reconnect: true,
reconnect_delay: Duration::from_secs(1),
max_reconnect_delay: Duration::from_secs(30),
max_reconnect_attempts: Some(5),
ping_interval: Some(Duration::from_secs(30)),
request_timeout: Duration::from_secs(10),
};
let transport = WebSocketTransport::new(config);
Ok(Box::new(transport))
}
#[cfg(not(feature = "websocket"))]
fn create_websocket_transport() -> Result<Box<dyn Transport>> {
Err(Error::feature_disabled(
"websocket",
"transport `websocket`",
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_create_stdio_transport() {
let transport = create_transport(&TransportType::Stdio);
assert!(transport.is_ok());
let t = transport.unwrap();
assert_eq!(t.transport_type(), "stdio");
}
#[cfg(feature = "sse")]
#[tokio::test]
async fn test_create_sse_transport() {
let transport = create_transport(&TransportType::Sse);
assert!(transport.is_ok());
}
#[cfg(not(feature = "sse"))]
#[test]
fn test_sse_without_feature_errors_and_names_the_feature() {
let msg = create_transport(&TransportType::Sse)
.expect_err("sse must fail when compiled out, not fall back to stdio")
.to_string();
assert!(
msg.contains("sse") && msg.contains("--features"),
"the error must tell the operator how to fix it, got: {msg}"
);
}
#[cfg(feature = "websocket")]
#[test]
fn test_create_websocket_transport() {
let transport = create_transport(&TransportType::WebSocket);
assert!(transport.is_ok());
}
#[cfg(not(feature = "websocket"))]
#[test]
fn test_websocket_without_feature_errors_and_names_the_feature() {
let msg = create_transport(&TransportType::WebSocket)
.expect_err("websocket must fail when compiled out, not fall back to stdio")
.to_string();
assert!(
msg.contains("websocket") && msg.contains("--features"),
"the error must tell the operator how to fix it, got: {msg}"
);
}
#[test]
fn test_stdio_works_in_every_feature_configuration() {
assert!(create_transport(&TransportType::Stdio).is_ok());
}
}