use anyhow::Result;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
#[cfg(feature = "enhanced")]
pub mod enhanced;
pub mod http;
pub mod proxy;
pub mod stdio;
pub mod websocket;
pub use http::HttpTransport;
pub use proxy::ProxyTransport;
pub use stdio::StdioTransport;
pub use websocket::WebSocketTransport;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TransportMessage {
pub id: String,
pub payload: serde_json::Value,
pub metadata: TransportMetadata,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct TransportMetadata {
pub client_id: Option<String>,
pub source_address: Option<String>,
pub headers: std::collections::HashMap<String, String>,
pub timestamp: Option<chrono::DateTime<chrono::Utc>>,
pub trace_id: Option<String>,
}
#[derive(Debug, Clone)]
pub struct ConnectionInfo {
pub id: String,
pub transport_type: TransportType,
pub client_id: Option<String>,
pub remote_addr: Option<String>,
pub connected_at: chrono::DateTime<chrono::Utc>,
pub security_info: Option<SecurityInfo>,
}
#[derive(Debug, Clone)]
pub struct SecurityInfo {
pub encrypted: bool,
pub tls_version: Option<String>,
pub cipher_suite: Option<String>,
pub client_cert: Option<CertificateInfo>,
}
#[derive(Debug, Clone)]
pub struct CertificateInfo {
pub subject: String,
pub issuer: String,
pub serial: String,
pub expires_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum TransportType {
Stdio,
Http,
WebSocket,
Grpc,
Custom(u32),
}
#[async_trait]
pub trait Transport: Send + Sync {
fn transport_type(&self) -> TransportType;
async fn start(&mut self) -> Result<()>;
async fn stop(&mut self) -> Result<()>;
async fn accept(&mut self) -> Result<Box<dyn TransportConnection>>;
async fn connect(&mut self, address: &str) -> Result<Box<dyn TransportConnection>>;
fn is_running(&self) -> bool;
fn get_stats(&self) -> TransportStats;
async fn set_option(&mut self, key: &str, value: serde_json::Value) -> Result<()>;
}
#[async_trait]
pub trait TransportConnection: Send + Sync {
fn connection_info(&self) -> &ConnectionInfo;
async fn send(&mut self, message: TransportMessage) -> Result<()>;
async fn receive(&mut self) -> Result<Option<TransportMessage>>;
async fn close(&mut self) -> Result<()>;
fn is_connected(&self) -> bool;
fn get_stats(&self) -> ConnectionStats;
async fn set_option(&mut self, key: &str, value: serde_json::Value) -> Result<()>;
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct TransportStats {
pub connections_accepted: u64,
pub active_connections: u64,
pub messages_sent: u64,
pub messages_received: u64,
pub bytes_sent: u64,
pub bytes_received: u64,
pub errors: u64,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ConnectionStats {
pub messages_sent: u64,
pub messages_received: u64,
pub bytes_sent: u64,
pub bytes_received: u64,
pub errors: u64,
pub rtt_us: Option<u64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TransportConfig {
pub multiplexing: bool,
pub transports: Vec<TransportTypeConfig>,
pub timeouts: TimeoutConfig,
pub security: SecurityConfig,
pub buffers: BufferConfig,
}
impl Default for TransportConfig {
fn default() -> Self {
Self {
multiplexing: false,
transports: vec![TransportTypeConfig {
transport_type: TransportType::Stdio,
enabled: true,
config: serde_json::json!({}),
}],
timeouts: TimeoutConfig::default(),
security: SecurityConfig::default(),
buffers: BufferConfig::default(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TransportTypeConfig {
pub transport_type: TransportType,
pub enabled: bool,
pub config: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TimeoutConfig {
pub connect_ms: u64,
pub read_ms: u64,
pub write_ms: u64,
pub keepalive_ms: Option<u64>,
}
impl Default for TimeoutConfig {
fn default() -> Self {
Self {
connect_ms: 5000,
read_ms: 30000,
write_ms: 30000,
keepalive_ms: Some(60000),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SecurityConfig {
pub require_tls: bool,
pub min_tls_version: Option<String>,
pub cipher_suites: Option<Vec<String>>,
pub client_auth: ClientAuthConfig,
pub encryption: bool,
}
impl Default for SecurityConfig {
fn default() -> Self {
Self {
require_tls: true,
min_tls_version: Some("1.2".to_string()),
cipher_suites: None,
client_auth: ClientAuthConfig::default(),
encryption: true,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClientAuthConfig {
pub required: bool,
pub ca_certs: Vec<String>,
pub verify_depth: u32,
}
impl Default for ClientAuthConfig {
fn default() -> Self {
Self {
required: false,
ca_certs: Vec::new(),
verify_depth: 3,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BufferConfig {
pub recv_buffer_size: usize,
pub send_buffer_size: usize,
pub message_queue_size: usize,
}
impl Default for BufferConfig {
fn default() -> Self {
Self {
recv_buffer_size: 65536,
send_buffer_size: 65536,
message_queue_size: 1000,
}
}
}
pub struct TransportManager {
#[allow(dead_code)]
config: TransportConfig,
transports: Vec<Box<dyn Transport>>,
connections: Arc<tokio::sync::RwLock<Vec<Box<dyn TransportConnection>>>>,
#[allow(dead_code)]
message_handler: Arc<dyn MessageHandler>,
}
#[async_trait]
pub trait MessageHandler: Send + Sync {
async fn handle_message(
&self,
message: TransportMessage,
connection: &dyn TransportConnection,
) -> Result<Option<TransportMessage>>;
async fn on_connect(&self, connection: &dyn TransportConnection) -> Result<()>;
async fn on_disconnect(&self, connection: &dyn TransportConnection) -> Result<()>;
}
impl TransportManager {
pub fn new(config: TransportConfig, handler: Arc<dyn MessageHandler>) -> Result<Self> {
Ok(Self {
config,
transports: Vec::new(),
connections: Arc::new(tokio::sync::RwLock::new(Vec::new())),
message_handler: handler,
})
}
pub fn add_transport(&mut self, transport: Box<dyn Transport>) -> Result<()> {
self.transports.push(transport);
Ok(())
}
pub async fn start(&mut self) -> Result<()> {
for transport in &mut self.transports {
transport.start().await?;
}
Ok(())
}
pub async fn stop(&mut self) -> Result<()> {
let mut connections = self.connections.write().await;
for conn in connections.iter_mut() {
let _ = conn.close().await;
}
connections.clear();
for transport in &mut self.transports {
transport.stop().await?;
}
Ok(())
}
pub fn get_stats(&self) -> Vec<(TransportType, TransportStats)> {
self.transports
.iter()
.map(|t| (t.transport_type(), t.get_stats()))
.collect()
}
}
pub trait TransportFactory: Send + Sync {
fn create(&self, config: &TransportTypeConfig) -> Result<Box<dyn Transport>>;
}
pub struct DefaultTransportFactory;
impl TransportFactory for DefaultTransportFactory {
fn create(&self, config: &TransportTypeConfig) -> Result<Box<dyn Transport>> {
match config.transport_type {
TransportType::Stdio => Ok(Box::new(StdioTransport::new(config.config.clone())?)),
TransportType::Http => Ok(Box::new(HttpTransport::new(config.config.clone())?)),
TransportType::WebSocket => {
Ok(Box::new(WebSocketTransport::new(config.config.clone())?))
},
#[cfg(feature = "enhanced")]
TransportType::Grpc => Ok(Box::new(enhanced::GrpcTransport::new(
config.config.clone(),
)?)),
_ => Err(anyhow::anyhow!(
"Unsupported transport type: {:?}",
config.transport_type
)),
}
}
}
pub fn create_transport(_config: &crate::config::Config) -> Arc<dyn Transport> {
let transport_config = TransportTypeConfig {
transport_type: TransportType::Stdio,
enabled: true,
config: serde_json::json!({}),
};
let factory = DefaultTransportFactory;
factory
.create(&transport_config)
.unwrap_or_else(|_| Box::new(StdioTransport::new(serde_json::json!({})).unwrap()))
.into()
}
pub struct TransportMessageBuilder {
message: TransportMessage,
}
impl TransportMessageBuilder {
pub fn new(payload: serde_json::Value) -> Self {
Self {
message: TransportMessage {
id: uuid::Uuid::new_v4().to_string(),
payload,
metadata: TransportMetadata::default(),
},
}
}
pub fn with_client_id(mut self, client_id: String) -> Self {
self.message.metadata.client_id = Some(client_id);
self
}
pub fn with_trace_id(mut self, trace_id: String) -> Self {
self.message.metadata.trace_id = Some(trace_id);
self
}
pub fn with_header(mut self, key: String, value: String) -> Self {
self.message.metadata.headers.insert(key, value);
self
}
pub fn build(mut self) -> TransportMessage {
self.message.metadata.timestamp = Some(chrono::Utc::now());
self.message
}
}