use crate::client::config::ClientConfig;
use crate::client::transports::common::{ClientConnectionHelper, ClientMessageObserver};
use crate::client::transports::{Client, ClientCore};
use crate::common::config_types::TlsConfig;
use crate::common::error::{FlareError, Result};
use crate::common::generate_id;
use crate::common::platform::{sleep, timeout};
use crate::common::protocol::Frame;
use crate::transport::connection::Connection;
use crate::transport::events::{ArcObserver, ConnectionEvent};
use crate::transport::quic::QUICTransport;
use async_trait::async_trait;
use quinn::Endpoint;
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::net::lookup_host;
use tokio::sync::Mutex;
pub struct QUICClient {
config: ClientConfig,
connection: Option<Arc<Mutex<Box<dyn Connection>>>>,
connection_id: String,
core: ClientCore,
reconnect_attempts: u32,
endpoint: Option<Endpoint>,
_client_config: Option<quinn::ClientConfig>, }
impl QUICClient {
pub fn new(config: ClientConfig) -> Result<Self> {
let connection_id = config.connection_id.clone().unwrap_or_else(generate_id);
let core = ClientCore::new(&config);
let (endpoint, client_config) = Self::create_quic_endpoint_with_tls(&config.tls)?;
Ok(Self {
config,
connection: None,
connection_id,
core,
reconnect_attempts: 0,
endpoint: Some(endpoint),
_client_config: Some(client_config),
})
}
pub async fn connect_with_config(config: ClientConfig) -> Result<Self> {
let mut client = Self::new(config)?;
client.connect().await?;
Ok(client)
}
pub fn with_core(config: ClientConfig, core: ClientCore) -> Result<Self> {
Self::with_core_and_endpoint(config, core, None)
}
pub fn with_core_and_endpoint(
config: ClientConfig,
core: ClientCore,
endpoint_opt: Option<(Endpoint, quinn::ClientConfig)>,
) -> Result<Self> {
let connection_id = config.connection_id.clone().unwrap_or_else(generate_id);
let (endpoint, client_config) = endpoint_opt.unwrap_or_else(|| {
Self::create_quic_endpoint_with_tls(&config.tls)
.expect("Failed to create QUIC endpoint")
});
Ok(Self {
config,
connection: None,
connection_id,
core,
reconnect_attempts: 0,
endpoint: Some(endpoint),
_client_config: Some(client_config),
})
}
pub fn create_quic_endpoint() -> Result<(Endpoint, quinn::ClientConfig)> {
Self::create_quic_endpoint_with_tls(&TlsConfig::none())
}
pub fn create_quic_endpoint_with_tls(
tls: &TlsConfig,
) -> Result<(Endpoint, quinn::ClientConfig)> {
use crate::common::cert::create_client_config;
use crate::common::cert::create_client_config_with_tls;
use quinn::crypto::rustls::QuicClientConfig;
let rustls_config = if tls.requires_custom_client_tls() {
create_client_config_with_tls(tls)
} else {
create_client_config()
}
.map_err(|e| {
FlareError::protocol_error(format!("Failed to create client TLS config: {}", e))
})?;
let rustls_config_arc = Arc::new(rustls_config);
let quic_crypto_config = QuicClientConfig::try_from(rustls_config_arc).map_err(|e| {
FlareError::protocol_error(format!("Failed to create QUIC client config: {}", e))
})?;
let client_config = quinn::ClientConfig::new(Arc::new(quic_crypto_config));
let mut endpoint = Endpoint::client("[::]:0".parse().unwrap()).map_err(|e| {
FlareError::connection_failed(format!("Failed to create endpoint: {}", e))
})?;
endpoint.set_default_client_config(client_config.clone());
Ok((endpoint, client_config))
}
pub async fn establish_network_connection(
&mut self,
) -> Result<Arc<Mutex<Box<dyn Connection>>>> {
let connection = self.establish_quic_connection().await?;
let connection_arc = Arc::new(Mutex::new(connection));
self.connection = Some(Arc::clone(&connection_arc));
Ok(connection_arc)
}
async fn internal_connect(&mut self) -> Result<()> {
let connection_arc = if let Some(connection) = &self.connection {
connection.clone()
} else {
self.establish_network_connection().await?
};
self.setup_connection_with_observer(connection_arc.clone())
.await?;
self.core
.handle_connection_event(&ConnectionEvent::Connected);
self.reconnect_attempts = 0;
Ok(())
}
async fn establish_quic_connection(&self) -> Result<Box<dyn Connection>> {
let (server_addr, hostname) = self.parse_server_address().await?;
let endpoint = self
.endpoint
.as_ref()
.ok_or_else(|| FlareError::connection_failed("Endpoint not initialized".to_string()))?;
let connecting = endpoint
.connect(server_addr, &hostname)
.map_err(|e| FlareError::connection_failed(e.to_string()))?;
let quinn_connection = timeout(self.config.connect_timeout, connecting)
.await
.map_err(|_| FlareError::connection_timeout("Connection timeout".to_string()))?
.map_err(|e| FlareError::connection_failed(e.to_string()))?;
let (send, recv) = quinn_connection
.open_bi()
.await
.map_err(|e| FlareError::connection_failed(e.to_string()))?;
let transport = QUICTransport::new(send, recv);
Ok(Box::new(transport))
}
async fn parse_server_address(&self) -> Result<(SocketAddr, String)> {
let address_str = self
.config
.get_protocol_url(&crate::common::config_types::TransportProtocol::QUIC)
.replace("quic://", "");
let server_addr = match address_str.parse::<SocketAddr>() {
Ok(addr) => addr,
Err(_) => {
lookup_host(&address_str)
.await
.map_err(|e| FlareError::protocol_error(format!("DNS lookup failed: {}", e)))?
.next()
.ok_or_else(|| {
FlareError::protocol_error(format!("No address found for {}", address_str))
})?
}
};
let hostname =
if address_str.starts_with("localhost") || address_str.starts_with("127.0.0.1") {
"localhost".to_string()
} else {
address_str
.split(':')
.next()
.unwrap_or("localhost")
.to_string()
};
Ok((server_addr, hostname))
}
async fn setup_connection_with_observer(
&mut self,
connection: Arc<Mutex<Box<dyn Connection>>>,
) -> Result<()> {
let core_clone = Arc::new(self.core.clone());
let message_observer = Arc::new(ClientMessageObserver::new(core_clone));
ClientConnectionHelper::setup_connection_and_send_connect(
Arc::clone(&connection),
&mut self.core,
message_observer,
)
.await?;
self.connection = Some(connection);
Ok(())
}
async fn send_frame_internal(&self, frame: &Frame) -> Result<()> {
ClientConnectionHelper::send_frame_internal(&self.core, self.connection.as_ref(), frame)
.await
}
async fn try_reconnect(&mut self) -> Result<()> {
if let Some(max) = self.config.max_reconnect_attempts
&& self.reconnect_attempts >= max
{
return Err(FlareError::connection_failed(format!(
"Max reconnect attempts ({}) exceeded",
max
)));
}
self.core.state_manager.start_connecting();
self.reconnect_attempts += 1;
sleep(self.config.reconnect_interval).await;
if let Some(conn) = self.connection.take() {
let mut c = conn.lock().await;
let _ = c.close().await;
}
self.internal_connect().await
}
pub fn core(&self) -> &ClientCore {
&self.core
}
}
#[async_trait]
impl Client for QUICClient {
async fn connect(&mut self) -> Result<()> {
if !self.core.can_connect() {
return Err(FlareError::protocol_error(
"Cannot connect: connection state is not ready".to_string(),
));
}
self.core.state_manager.start_connecting();
match self.internal_connect().await {
Ok(()) => Ok(()),
Err(e) => {
self.core.state_manager.set_failed();
if ClientConnectionHelper::can_reconnect(self.config.max_reconnect_attempts) {
self.try_reconnect().await
} else {
Err(e)
}
}
}
}
fn set_disconnect_requested(&mut self, value: bool) {
self.core.set_disconnect_requested(value);
}
async fn disconnect(&mut self) -> Result<()> {
ClientConnectionHelper::disconnect_internal(self.connection.take(), &mut self.core).await
}
async fn send_frame(&mut self, frame: &Frame) -> Result<()> {
if !self.is_connected()
&& ClientConnectionHelper::can_reconnect(self.config.max_reconnect_attempts)
{
self.try_reconnect().await?;
}
self.send_frame_internal(frame).await
}
fn is_connected(&self) -> bool {
matches!(
self.core.state(),
crate::client::connection::ConnectionState::Connected
) && self.connection.is_some()
}
fn add_observer(&mut self, observer: ArcObserver) {
self.core.add_observer(observer);
}
fn remove_observer(&mut self, observer: ArcObserver) {
self.core.remove_observer(observer);
}
fn connection_id(&self) -> Option<String> {
Some(self.connection_id.clone())
}
}