use std::error::Error;
use std::net::UdpSocket;
use std::sync::Arc;
use std::time::Duration;
#[cfg(feature = "bebop")]
use crate::generated::schema::ServerConnectInfo;
use crate::helpers::get_internal_websocket::{handle_websocket, TryConnectGuard};
use crate::helpers::get_outer_websocket::wrap_get_outer_websocket;
use crate::helpers::scan_manager::ScanManager;
use crate::helpers::{
common::make_ping_message,
connection_store::ConnectionStore,
get_internal_websocket::wrap_get_internal_websocket,
server_sender::{to_ws_url, SenderStatus, ServerSenderTrait},
traits::date_time::now,
};
use crate::{helpers::metrics::Metrics, log_debug, log_error, AtomicWebsocketType};
use tokio::sync::mpsc::Receiver;
use tokio_util::sync::CancellationToken;
use super::types::RwServerSender;
#[derive(Clone)]
pub struct ClientOptions {
pub use_ping: bool,
pub url: String,
pub retry_seconds: u64,
pub use_keep_ip: bool,
pub connect_timeout_seconds: u64,
pub atomic_websocket_type: AtomicWebsocketType,
#[cfg(feature = "rustls")]
pub use_tls: bool,
pub handler_buffer_size: usize,
pub status_buffer_size: usize,
pub per_connection_buffer_size: usize,
pub spillover_buffer_size: usize,
pub use_scan_discovery: bool,
pub scan_timeout_seconds: u64,
}
impl Default for ClientOptions {
fn default() -> Self {
Self {
use_ping: true,
url: "".into(),
retry_seconds: 30,
use_keep_ip: false,
connect_timeout_seconds: 3,
atomic_websocket_type: AtomicWebsocketType::Internal,
#[cfg(feature = "rustls")]
use_tls: true,
handler_buffer_size: 256,
status_buffer_size: 8,
per_connection_buffer_size: 8,
spillover_buffer_size: 1024,
use_scan_discovery: false,
scan_timeout_seconds: 60,
}
}
}
pub struct AtomicClient {
pub server_sender: RwServerSender,
pub options: ClientOptions,
pub(crate) cancel_token: CancellationToken,
}
impl AtomicClient {
pub async fn internal_initialize(&self, connection_store: Arc<dyn ConnectionStore>) {
self.regist_id(connection_store).await;
tokio::spawn(internal_ping_loop_cheker(
self.server_sender.clone(),
self.options.clone(),
self.cancel_token.clone(),
));
}
pub async fn outer_initialize(&self, connection_store: Arc<dyn ConnectionStore>) {
#[cfg(feature = "rustls")]
self.initial_rustls();
self.regist_id(connection_store).await;
tokio::spawn(outer_ping_loop_cheker(
self.server_sender.clone(),
self.options.clone(),
self.cancel_token.clone(),
));
}
pub async fn disconnect(&self) {
self.cancel_token.cancel();
self.server_sender.remove_ip().await;
}
pub async fn scan_and_connect(&self, port: &str, connection_store: Arc<dyn ConnectionStore>) -> bool {
if !self.server_sender.is_need_connect().await {
return false;
}
if self.server_sender.is_valid_server_ip().await {
self.server_sender.send_status(SenderStatus::Connected).await;
return true;
}
{
let mut guard = self.server_sender.write().await;
if guard.is_scanning {
return false;
}
guard.is_scanning = true;
}
if get_ip_address().is_empty() {
self.server_sender.write().await.is_scanning = false;
self.server_sender.send_status(SenderStatus::Disconnected).await;
return false;
}
self.server_sender.send_status(SenderStatus::Connecting).await;
let mut manager = ScanManager::new(port);
let scan_timeout = Duration::from_secs(self.options.scan_timeout_seconds.max(1));
let found = manager.run_with_timeout(scan_timeout).await;
self.server_sender.write().await.is_scanning = false;
match found {
Some((server_ip, ws_stream)) => {
let Some(connect_guard) =
TryConnectGuard::try_acquire(self.server_sender.clone()).await
else {
log_debug!(
"Scan found a server but a connection attempt is already in progress"
);
return false;
};
let server_sender = self.server_sender.clone();
let options = self.options.clone();
tokio::spawn(async move {
if let Err(error) = handle_websocket(
connection_store,
server_sender.clone(),
options,
server_ip,
ws_stream,
connect_guard,
)
.await
{
log_error!("Error handling websocket: {:?}", error);
}
});
true
}
None => {
self.server_sender.send_status(SenderStatus::Disconnected).await;
false
}
}
}
pub async fn get_outer_connect(
&self,
connection_store: Arc<dyn ConnectionStore>,
) -> Result<(), Box<dyn Error>> {
get_outer_connect(connection_store, self.server_sender.clone(), self.options.clone()).await
}
#[cfg(all(feature = "native-db", feature = "bebop"))]
pub async fn get_internal_connect(
&self,
input: Option<ServerConnectInfo<'_>>,
connection_store: Arc<dyn ConnectionStore>,
) -> Result<(), Box<dyn Error>> {
get_internal_connect(
input,
connection_store,
self.server_sender.clone(),
self.options.clone(),
)
.await
}
#[cfg(all(not(feature = "native-db"), feature = "bebop"))]
pub async fn get_internal_connect(
&self,
_input: Option<ServerConnectInfo<'_>>,
connection_store: Arc<dyn ConnectionStore>,
) -> Result<(), Box<dyn Error>> {
get_internal_connect(
None,
connection_store,
self.server_sender.clone(),
self.options.clone(),
)
.await
}
#[cfg(not(feature = "bebop"))]
pub async fn get_internal_connect(
&self,
_input: Option<()>,
connection_store: Arc<dyn ConnectionStore>,
) -> Result<(), Box<dyn Error>> {
get_internal_connect(
None,
connection_store,
self.server_sender.clone(),
self.options.clone(),
)
.await
}
#[cfg(feature = "rustls")]
pub fn initial_rustls(&self) {
use rustls::crypto::{ring, CryptoProvider};
if CryptoProvider::get_default().is_none() {
let provider = ring::default_provider();
if let Err(e) = provider.install_default() {
log_error!("Failed to install rustls crypto provider: {:?}", e);
}
}
}
pub async fn regist_id(&self, connection_store: Arc<dyn ConnectionStore>) {
connection_store.ensure_client_id().await;
}
pub async fn get_status_receiver(&self) -> Option<Receiver<SenderStatus>> {
self.server_sender.get_status_receiver().await
}
pub async fn get_handle_message_receiver(&self) -> Option<Receiver<Vec<u8>>> {
self.server_sender.get_handle_message_receiver().await
}
pub async fn metrics(&self) -> std::sync::Arc<Metrics> {
self.server_sender.read().await.metrics.clone()
}
}
async fn internal_ping_loop_cheker(
server_sender: RwServerSender,
options: ClientOptions,
cancel_token: CancellationToken,
) {
let retry_seconds = options.retry_seconds.max(1);
let use_keep_ip = options.use_keep_ip;
let max_retry_seconds = retry_seconds * 8;
let mut current_retry_seconds = retry_seconds;
loop {
tokio::select! {
_ = cancel_token.cancelled() => {
log_debug!("internal_ping_loop_cheker cancelled");
break;
}
_ = tokio::time::sleep(Duration::from_secs(current_retry_seconds)) => {}
}
let server_sender_read = server_sender.read().await;
if server_sender_read.server_received_times > 0
&& server_sender_read.server_received_times + (retry_seconds as i64 * 4)
< now().timestamp()
{
drop(server_sender_read);
server_sender.send_status(SenderStatus::Disconnected).await;
if !use_keep_ip {
server_sender.remove_ip_if_valid_server_ip("").await;
}
server_sender.send_status(SenderStatus::Reconnecting).await;
let (metrics, connection_store) = {
let guard = server_sender.read().await;
(guard.metrics.clone(), guard.connection_store.clone())
};
metrics.inc_reconnections();
let server_sender = server_sender.clone();
let options = options.clone();
tokio::spawn(async move {
if let Err(e) =
get_internal_connect(None, connection_store, server_sender, options).await
{
log_error!("Internal reconnection failed: {:?}", e);
}
});
current_retry_seconds = (current_retry_seconds * 2).min(max_retry_seconds);
}
else if server_sender_read.server_received_times > 0
&& server_sender_read.server_received_times + (retry_seconds as i64 * 2)
< now().timestamp()
{
if options.use_ping {
log_debug!("Try ping from loop checker");
let connection_store = server_sender_read.connection_store.clone();
drop(server_sender_read);
let id: String = connection_store.get_client_id().await;
server_sender.send(make_ping_message(&id)).await;
}
} else {
current_retry_seconds = retry_seconds;
}
log_debug!("loop server checker finish");
}
}
async fn outer_ping_loop_cheker(
server_sender: RwServerSender,
options: ClientOptions,
cancel_token: CancellationToken,
) {
let retry_seconds = options.retry_seconds.max(1);
let use_keep_ip = options.use_keep_ip;
let max_retry_seconds = retry_seconds * 8;
let mut current_retry_seconds = retry_seconds;
loop {
tokio::select! {
_ = cancel_token.cancelled() => {
log_debug!("outer_ping_loop_cheker cancelled");
break;
}
_ = tokio::time::sleep(Duration::from_secs(current_retry_seconds)) => {}
}
let server_sender_read = server_sender.read().await;
if server_sender_read.server_received_times > 0
&& server_sender_read.server_received_times + (retry_seconds as i64 * 4)
< now().timestamp()
{
drop(server_sender_read);
server_sender.send_status(SenderStatus::Disconnected).await;
if !use_keep_ip {
server_sender.remove_ip().await;
}
server_sender.send_status(SenderStatus::Reconnecting).await;
let (metrics, connection_store) = {
let guard = server_sender.read().await;
(guard.metrics.clone(), guard.connection_store.clone())
};
metrics.inc_reconnections();
let server_sender = server_sender.clone();
let options = options.clone();
tokio::spawn(async move {
if let Err(e) = get_outer_connect(connection_store, server_sender, options).await {
log_error!("External reconnection failed: {:?}", e);
}
});
current_retry_seconds = (current_retry_seconds * 2).min(max_retry_seconds);
}
else if server_sender_read.server_received_times > 0
&& server_sender_read.server_received_times + (retry_seconds as i64 * 2)
< now().timestamp()
{
log_debug!(
"send: {:?}, current: {:?}",
server_sender_read.server_received_times,
now().timestamp()
);
if options.use_ping {
log_debug!("Try ping from loop checker");
let connection_store = server_sender_read.connection_store.clone();
drop(server_sender_read);
let id: String = connection_store.get_client_id().await;
server_sender.send(make_ping_message(&id)).await;
}
} else {
current_retry_seconds = retry_seconds;
}
log_debug!("loop server checker finish");
}
}
pub async fn get_outer_connect(
connection_store: Arc<dyn ConnectionStore>,
server_sender: RwServerSender,
options: ClientOptions,
) -> Result<(), Box<dyn Error>> {
if server_sender.read().await.is_try_connect {
return Ok(());
}
if server_sender.is_valid_server_ip().await {
server_sender.send_status(SenderStatus::Connected).await;
return Ok(());
}
let server_connect_info = connection_store.get_server_connect_info().await;
log_debug!("server_connect_info: {:?}", server_connect_info);
if options.url.is_empty() && !server_sender.is_valid_server_ip().await {
server_sender.send_status(SenderStatus::Disconnected).await;
return Ok(());
}
server_sender.send_status(SenderStatus::Connecting).await;
tokio::spawn(wrap_get_outer_websocket(
connection_store,
server_sender,
options,
));
Ok(())
}
#[cfg(all(feature = "native-db", feature = "bebop"))]
pub async fn get_internal_connect(
input: Option<ServerConnectInfo<'_>>,
connection_store: Arc<dyn ConnectionStore>,
server_sender: RwServerSender,
options: ClientOptions,
) -> Result<(), Box<dyn Error>> {
if server_sender.read().await.is_try_connect {
return Ok(());
}
if server_sender.is_valid_server_ip().await {
server_sender.send_status(SenderStatus::Connected).await;
return Ok(());
}
let server_connect_info = connection_store.get_server_connect_info().await;
log_debug!("server_connect_info: {:?}", server_connect_info);
if let (Some(input_ref), None) = (input.as_ref(), server_connect_info.as_ref()) {
connection_store
.set_server_connect_info("", input_ref.port)
.await;
}
if input.is_none() && server_connect_info.is_none() {
server_sender.send_status(SenderStatus::Disconnected).await;
return Ok(());
}
let (connect_server_ip, connect_port): (String, String) = match input.as_ref() {
Some(info) => {
let server_ip = server_connect_info
.as_ref()
.map(|(ip, _)| ip.clone())
.unwrap_or_default();
(server_ip, info.port.to_owned())
}
None => {
let Some((server_ip, port)) = server_connect_info else {
server_sender.send_status(SenderStatus::Disconnected).await;
return Ok(());
};
(server_ip, port)
}
};
match connect_server_ip.as_str() {
"" => {
if !options.use_scan_discovery {
server_sender.send_status(SenderStatus::Disconnected).await;
return Ok(());
}
{
let mut guard = server_sender.write().await;
if guard.is_scanning {
return Ok(());
}
guard.is_scanning = true;
}
if get_ip_address().is_empty() {
server_sender.write().await.is_scanning = false;
server_sender.send_status(SenderStatus::Disconnected).await;
return Ok(());
}
server_sender.send_status(SenderStatus::Connecting).await;
let scan_timeout = Duration::from_secs(options.scan_timeout_seconds.max(1));
let found = ScanManager::new(&connect_port)
.run_with_timeout(scan_timeout)
.await;
server_sender.write().await.is_scanning = false;
match found {
Some((server_ip, ws_stream)) => {
match TryConnectGuard::try_acquire(server_sender.clone()).await {
Some(connect_guard) => {
tokio::spawn(async move {
if let Err(error) = handle_websocket(
connection_store,
server_sender.clone(),
options,
server_ip,
ws_stream,
connect_guard,
)
.await
{
log_error!("Error handling websocket: {:?}", error);
}
});
}
None => {
log_debug!(
"Scan found a server but a connection attempt is already in progress"
);
}
}
}
None => {
server_sender.send_status(SenderStatus::Disconnected).await;
}
}
}
_server_ip => {
let url = to_ws_url(_server_ip, &connect_port);
server_sender.send_status(SenderStatus::Connecting).await;
tokio::spawn(wrap_get_internal_websocket(
connection_store,
server_sender.clone(),
url,
options.clone(),
));
}
};
Ok(())
}
#[cfg(not(all(feature = "native-db", feature = "bebop")))]
pub async fn get_internal_connect(
_input: Option<()>,
connection_store: Arc<dyn ConnectionStore>,
server_sender: RwServerSender,
options: ClientOptions,
) -> Result<(), Box<dyn Error>> {
if server_sender.read().await.is_try_connect {
return Ok(());
}
if server_sender.is_valid_server_ip().await {
server_sender.send_status(SenderStatus::Connected).await;
return Ok(());
}
if !options.use_scan_discovery {
server_sender.send_status(SenderStatus::Disconnected).await;
return Ok(());
}
{
let mut guard = server_sender.write().await;
if guard.is_scanning {
return Ok(());
}
guard.is_scanning = true;
}
if get_ip_address().is_empty() {
server_sender.write().await.is_scanning = false;
server_sender.send_status(SenderStatus::Disconnected).await;
return Ok(());
}
server_sender.send_status(SenderStatus::Connecting).await;
let scan_timeout = Duration::from_secs(options.scan_timeout_seconds.max(1));
let found = ScanManager::new("9000").run_with_timeout(scan_timeout).await;
server_sender.write().await.is_scanning = false;
match found {
Some((server_ip, ws_stream)) => {
match TryConnectGuard::try_acquire(server_sender.clone()).await {
Some(connect_guard) => {
tokio::spawn(async move {
if let Err(error) = handle_websocket(
connection_store,
server_sender.clone(),
options,
server_ip,
ws_stream,
connect_guard,
)
.await
{
log_error!("Error handling websocket: {:?}", error);
}
});
}
None => {
log_debug!(
"Scan found a server but a connection attempt is already in progress"
);
}
}
}
None => {
server_sender.send_status(SenderStatus::Disconnected).await;
}
}
Ok(())
}
pub fn get_ip_address() -> String {
let socket = UdpSocket::bind("0.0.0.0:0");
let socket = match socket {
Ok(socket) => socket,
Err(_) => return "".into(),
};
match socket.connect("8.8.8.8:80") {
Ok(_) => {}
Err(_) => return "".into(),
};
let addr = match socket.local_addr() {
Ok(addr) => addr,
Err(_) => return "".into(),
};
addr.ip().to_string()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_client_options_default() {
let options = ClientOptions::default();
assert!(options.use_ping);
assert_eq!(options.url, "");
assert_eq!(options.retry_seconds, 30);
assert!(!options.use_keep_ip);
assert_eq!(options.connect_timeout_seconds, 3);
assert!(matches!(
options.atomic_websocket_type,
AtomicWebsocketType::Internal
));
}
#[cfg(feature = "rustls")]
#[test]
fn test_client_options_default_with_tls() {
let options = ClientOptions::default();
assert!(options.use_tls);
}
#[test]
fn test_client_options_clone() {
let options = ClientOptions {
use_ping: false,
url: "ws://example.com:9000".to_string(),
retry_seconds: 60,
use_keep_ip: true,
connect_timeout_seconds: 10,
atomic_websocket_type: AtomicWebsocketType::External,
#[cfg(feature = "rustls")]
use_tls: false,
..Default::default()
};
let cloned = options.clone();
assert!(!cloned.use_ping);
assert_eq!(cloned.url, "ws://example.com:9000");
assert_eq!(cloned.retry_seconds, 60);
assert!(cloned.use_keep_ip);
assert_eq!(cloned.connect_timeout_seconds, 10);
assert!(matches!(
cloned.atomic_websocket_type,
AtomicWebsocketType::External
));
}
#[test]
fn test_client_options_custom_values() {
let options = ClientOptions {
use_ping: false,
url: "192.168.1.100:9000".to_string(),
retry_seconds: 5,
use_keep_ip: true,
connect_timeout_seconds: 1,
atomic_websocket_type: AtomicWebsocketType::Internal,
#[cfg(feature = "rustls")]
use_tls: true,
..Default::default()
};
assert!(!options.use_ping);
assert_eq!(options.url, "192.168.1.100:9000");
assert_eq!(options.retry_seconds, 5);
assert!(options.use_keep_ip);
assert_eq!(options.connect_timeout_seconds, 1);
}
#[test]
fn test_get_ip_address_format() {
let ip = get_ip_address();
if !ip.is_empty() {
let parts: Vec<&str> = ip.split('.').collect();
assert_eq!(parts.len(), 4, "IP should have 4 octets");
for part in parts {
let num: Result<u8, _> = part.parse();
assert!(num.is_ok(), "Each octet should be a valid u8");
}
}
}
}