use crate::{
connection::{Connection, Server}, error::TransportError,
SessionId,
};
use async_trait::async_trait;
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::RwLock;
pub type BoxFuture<T> = Pin<Box<dyn Future<Output = T> + Send>>;
#[async_trait]
pub trait ProtocolFactory: Send + Sync {
fn protocol_name(&self) -> &'static str;
fn supported_schemes(&self) -> Vec<&'static str>;
async fn create_connection(
&self,
uri: &str,
config: Option<Box<dyn std::any::Any + Send + Sync>>,
) -> Result<Box<dyn Connection>, TransportError>;
async fn create_server(
&self,
bind_addr: &str,
config: Option<Box<dyn std::any::Any + Send + Sync>>,
) -> Result<Box<dyn Server>, TransportError>;
fn parse_uri(
&self,
uri: &str,
) -> Result<(std::net::SocketAddr, HashMap<String, String>), TransportError> {
if let Ok(addr) = uri.parse::<std::net::SocketAddr>() {
Ok((addr, HashMap::new()))
} else {
Err(TransportError::config_error(
"uri",
format!("Invalid URI: {}", uri),
))
}
}
fn default_config(&self) -> Box<dyn std::any::Any + Send + Sync>;
}
pub struct ProtocolRegistry {
factories: RwLock<HashMap<String, Arc<dyn ProtocolFactory>>>,
schemes: RwLock<HashMap<String, Arc<dyn ProtocolFactory>>>,
}
impl ProtocolRegistry {
pub fn new() -> Self {
Self {
factories: RwLock::new(HashMap::new()),
schemes: RwLock::new(HashMap::new()),
}
}
pub async fn register<F>(&self, factory: F) -> Result<(), TransportError>
where
F: ProtocolFactory + 'static,
{
let factory = Arc::new(factory);
let protocol_name = factory.protocol_name().to_string();
let schemes = factory.supported_schemes();
{
let mut factories = self.factories.write().await;
if factories.contains_key(&protocol_name) {
return Err(TransportError::config_error(
"protocol_name",
format!("Protocol '{}' already registered", protocol_name),
));
}
factories.insert(protocol_name.clone(), factory.clone());
}
{
let mut schemes_map = self.schemes.write().await;
for scheme in schemes {
if schemes_map.contains_key(scheme) {
return Err(TransportError::config_error(
"scheme",
format!("Scheme '{}' already registered", scheme),
));
}
schemes_map.insert(scheme.to_string(), factory.clone());
}
}
tracing::info!("Registered protocol: {}", protocol_name);
Ok(())
}
pub async fn get_factory(&self, protocol: &str) -> Option<Arc<dyn ProtocolFactory>> {
self.factories.read().await.get(protocol).cloned()
}
pub async fn get_factory_by_scheme(&self, scheme: &str) -> Option<Arc<dyn ProtocolFactory>> {
self.schemes.read().await.get(scheme).cloned()
}
pub async fn create_connection(
&self,
uri: &str,
config: Option<Box<dyn std::any::Any + Send + Sync>>,
) -> Result<Box<dyn Connection>, TransportError> {
let scheme = self.extract_scheme(uri)?;
if let Some(factory) = self.get_factory_by_scheme(&scheme).await {
factory.create_connection(uri, config).await
} else {
Err(TransportError::protocol_error(
"unknown",
format!("Unsupported protocol scheme: {}", scheme),
))
}
}
pub async fn create_server(
&self,
bind_addr: &str,
protocol: &str,
config: Option<Box<dyn std::any::Any + Send + Sync>>,
) -> Result<Box<dyn Server>, TransportError> {
if let Some(factory) = self.get_factory(protocol).await {
factory.create_server(bind_addr, config).await
} else {
Err(TransportError::protocol_error(
"unknown",
format!("Unsupported protocol: {}", protocol),
))
}
}
pub async fn list_protocols(&self) -> Vec<String> {
self.factories.read().await.keys().cloned().collect()
}
pub async fn list_schemes(&self) -> Vec<String> {
self.schemes.read().await.keys().cloned().collect()
}
fn extract_scheme(&self, uri: &str) -> Result<String, TransportError> {
if let Some(pos) = uri.find("://") {
Ok(uri[..pos].to_string())
} else {
if uri.parse::<std::net::SocketAddr>().is_ok() {
Ok("tcp".to_string()) } else {
Err(TransportError::config_error(
"uri",
format!("Invalid URI format: {}", uri),
))
}
}
}
}
impl Default for ProtocolRegistry {
fn default() -> Self {
Self::new()
}
}
pub trait ProtocolSet: Send + Sync {
fn create_connection(&self, uri: &str) -> BoxFuture<Result<SessionId, TransportError>>;
fn create_server(&self, bind_addr: &str) -> BoxFuture<Result<Box<dyn Server>, TransportError>>;
}
#[allow(dead_code)]
pub struct StandardProtocols {
_phantom: std::marker::PhantomData<()>,
}
#[allow(dead_code)]
pub struct PluginManager {
registry: ProtocolRegistry,
}
#[allow(dead_code)]
impl PluginManager {
pub fn load_from_dylib(&mut self, _path: &std::path::Path) -> Result<(), TransportError> {
Err(TransportError::config_error(
"plugin",
"Plugin loading not yet implemented",
))
}
}