use std::sync::Arc;
use std::time::Duration;
use crate::client::http::HttpTransportAdapter;
use crate::client::stdio::StdioTransportAdapter;
use crate::transport::{
McpServerConnectionConfig, McpTransport, McpTransportError, TransportTypeId,
};
pub struct TransportFactory;
impl TransportFactory {
pub async fn create(
config: &McpServerConnectionConfig,
) -> Result<Arc<dyn McpTransport>, McpTransportError> {
match config.transport {
TransportTypeId::Stdio => Self::create_stdio(config).await,
TransportTypeId::Http => Self::create_http(config).await,
}
}
async fn create_stdio(
config: &McpServerConnectionConfig,
) -> Result<Arc<dyn McpTransport>, McpTransportError> {
let command = config.command.as_ref().ok_or_else(|| {
McpTransportError::TransportError("Stdio transport requires command".to_string())
})?;
let timeout = Duration::from_secs(config.timeout_secs);
let transport = StdioTransportAdapter::connect_with_env(
command,
&config.args,
config.env.clone(),
Some(config.config.clone()),
timeout,
)
.await?;
Ok(Arc::new(transport))
}
async fn create_http(
config: &McpServerConnectionConfig,
) -> Result<Arc<dyn McpTransport>, McpTransportError> {
let url = config.url.as_ref().ok_or_else(|| {
McpTransportError::TransportError("HTTP transport requires URL".to_string())
})?;
let timeout = Duration::from_secs(config.timeout_secs);
let transport = HttpTransportAdapter::with_timeout(url, timeout)?;
Ok(Arc::new(transport))
}
pub fn is_supported(transport_type: TransportTypeId) -> bool {
matches!(
transport_type,
TransportTypeId::Stdio | TransportTypeId::Http
)
}
pub fn supported_types() -> Vec<TransportTypeId> {
vec![TransportTypeId::Stdio, TransportTypeId::Http]
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_is_supported() {
assert!(TransportFactory::is_supported(TransportTypeId::Stdio));
assert!(TransportFactory::is_supported(TransportTypeId::Http));
}
#[test]
fn test_supported_types() {
let types = TransportFactory::supported_types();
assert!(types.contains(&TransportTypeId::Stdio));
assert!(types.contains(&TransportTypeId::Http));
}
}