draft version of sans io tcp client
diff --git a/Cargo.lock b/Cargo.lock index 4db7dd2..4e51c95 100644 --- a/Cargo.lock +++ b/Cargo.lock
@@ -3804,14 +3804,17 @@ "iggy_common", "mockall", "num_cpus", + "parking_lot 0.12.4", "quinn", "reqwest", "reqwest-middleware", "reqwest-retry", + "rustc-hash 2.1.1", "rustls", "serde", "tokio", "tokio-rustls", + "tokio-util", "tracing", "trait-variant", "webpki-roots 1.0.2",
diff --git a/core/bench/src/utils/mod.rs b/core/bench/src/utils/mod.rs index cc6d5fc..baeb206 100644 --- a/core/bench/src/utils/mod.rs +++ b/core/bench/src/utils/mod.rs
@@ -70,7 +70,11 @@ .login_user(DEFAULT_ROOT_USERNAME, DEFAULT_ROOT_PASSWORD) .await?; - client.get_stats().await + let stats= client.get_stats().await?; + let _ = client.disconnect().await; + + drop(client); + Ok(stats) } pub async fn collect_server_logs_and_save_to_file( @@ -93,6 +97,9 @@ .await? .0; + // Disconnect the client to ensure proper cleanup + let _ = client.disconnect().await; + fs::write(output_dir.join("server_logs.zip"), snapshot).map_err(|e| { error!("Failed to write server logs to file: {:?}", e); IggyError::CannotWriteToFile
diff --git a/core/common/src/error/iggy_error.rs b/core/common/src/error/iggy_error.rs index cc8a284..b2209ce 100644 --- a/core/common/src/error/iggy_error.rs +++ b/core/common/src/error/iggy_error.rs
@@ -459,6 +459,20 @@ CannotReadIndexPosition = 10011, #[error("Cannot read index timestamp")] CannotReadIndexTimestamp = 10012, + #[error("Send queue is full")] + SendQueueFull = 10050, + #[error("Too many pending requests")] + TooManyPendingRequests = 10051, + #[error("IO Error")] + IoError = 10052, + #[error("Max number of retry has exceeded")] + MaxRetriesExceeded = 10053, + #[error("Connection timeout")] + ConnectionTimeout = 10054, + #[error("Incorrect connection state")] + IncorrectConnectionState = 10055, + #[error("Connection missed socket")] + ConnectionMissedSocket = 10056, } impl IggyError {
diff --git a/core/integration/src/tcp_client.rs b/core/integration/src/tcp_client.rs index 9df8d9b..ccdde06 100644 --- a/core/integration/src/tcp_client.rs +++ b/core/integration/src/tcp_client.rs
@@ -18,7 +18,11 @@ use crate::test_server::{ClientFactory, Transport}; use async_trait::async_trait; -use iggy::prelude::{Client, ClientWrapper, TcpClient, TcpClientConfig}; +use iggy::{ + connection::NewTokioTcpClient, + prelude::{Client, ClientWrapper, IggyClient, TcpClient, TcpClientConfig}, + runtime::TokioRuntime, +}; use std::sync::Arc; #[derive(Debug, Clone, Default)] @@ -43,12 +47,25 @@ tls_validate_certificate: self.tls_validate_certificate, ..TcpClientConfig::default() }; - let client = TcpClient::create(Arc::new(config)).unwrap_or_else(|e| { + + // let factory: StreamFactory<TokioCompat> = Arc::new(move |addr: SocketAddr| -> SFut<TokioCompat> { + // Box::pin(tokio_tcp(addr)) + // }); + + let tokio_rt = Arc::new(TokioRuntime {}); + let client = NewTokioTcpClient::create(Arc::new(config), tokio_rt).unwrap_or_else(|e| { panic!( - "Failed to create TcpClient, iggy-server has address {}, error: {:?}", + "Failed to create NewTcpClient, iggy-server has address {}, error: {:?}", self.server_addr, e ) }); + // let client = TcpClient::create(Arc::new(config)).unwrap_or_else(|e| { + // panic!( + // "Failed to create TcpClient, iggy-server has address {}, error: {:?}", + // self.server_addr, e + // ) + // }); + Client::connect(&client).await.unwrap_or_else(|e| { if self.tls_enabled { panic!( @@ -65,7 +82,7 @@ ) } }); - ClientWrapper::Tcp(client) + ClientWrapper::TcpTokio(client) } fn transport(&self) -> Transport {
diff --git a/core/sdk/Cargo.toml b/core/sdk/Cargo.toml index c6b2672..1350f0a 100644 --- a/core/sdk/Cargo.toml +++ b/core/sdk/Cargo.toml
@@ -49,14 +49,17 @@ iggy_binary_protocol = { workspace = true } iggy_common = { workspace = true } num_cpus = "1.17.0" +parking_lot = "0.12.4" quinn = { workspace = true } reqwest = { workspace = true } reqwest-middleware = { workspace = true } reqwest-retry = { workspace = true } +rustc-hash = "2.1.1" rustls = { workspace = true } serde = { workspace = true } tokio = { workspace = true } tokio-rustls = { workspace = true } +tokio-util.workspace = true tracing = { workspace = true } trait-variant = { workspace = true } webpki-roots = { workspace = true }
diff --git a/core/sdk/src/client_wrappers/binary_client.rs b/core/sdk/src/client_wrappers/binary_client.rs index 41d5744..fc764b3 100644 --- a/core/sdk/src/client_wrappers/binary_client.rs +++ b/core/sdk/src/client_wrappers/binary_client.rs
@@ -30,6 +30,7 @@ ClientWrapper::Http(client) => client.connect().await, ClientWrapper::Tcp(client) => client.connect().await, ClientWrapper::Quic(client) => client.connect().await, + ClientWrapper::TcpTokio(client) => client.connect().await, } } @@ -39,6 +40,7 @@ ClientWrapper::Http(client) => client.disconnect().await, ClientWrapper::Tcp(client) => client.disconnect().await, ClientWrapper::Quic(client) => client.disconnect().await, + ClientWrapper::TcpTokio(client) => client.disconnect().await, } } @@ -48,6 +50,7 @@ ClientWrapper::Http(client) => client.shutdown().await, ClientWrapper::Tcp(client) => client.shutdown().await, ClientWrapper::Quic(client) => client.shutdown().await, + ClientWrapper::TcpTokio(client) => client.shutdown().await, } } @@ -57,6 +60,7 @@ ClientWrapper::Http(client) => client.subscribe_events().await, ClientWrapper::Tcp(client) => client.subscribe_events().await, ClientWrapper::Quic(client) => client.subscribe_events().await, + ClientWrapper::TcpTokio(client) => client.subscribe_events().await, } } }
diff --git a/core/sdk/src/client_wrappers/binary_consumer_group_client.rs b/core/sdk/src/client_wrappers/binary_consumer_group_client.rs index 1811758..b82a5b3 100644 --- a/core/sdk/src/client_wrappers/binary_consumer_group_client.rs +++ b/core/sdk/src/client_wrappers/binary_consumer_group_client.rs
@@ -51,6 +51,11 @@ .get_consumer_group(stream_id, topic_id, group_id) .await } + ClientWrapper::TcpTokio(client) => { + client + .get_consumer_group(stream_id, topic_id, group_id) + .await + } } } @@ -64,6 +69,7 @@ ClientWrapper::Http(client) => client.get_consumer_groups(stream_id, topic_id).await, ClientWrapper::Tcp(client) => client.get_consumer_groups(stream_id, topic_id).await, ClientWrapper::Quic(client) => client.get_consumer_groups(stream_id, topic_id).await, + ClientWrapper::TcpTokio(client) => client.get_consumer_groups(stream_id, topic_id).await, } } @@ -95,6 +101,11 @@ .create_consumer_group(stream_id, topic_id, name, group_id) .await } + ClientWrapper::TcpTokio(client) => { + client + .create_consumer_group(stream_id, topic_id, name, group_id) + .await + } } } @@ -125,6 +136,11 @@ .delete_consumer_group(stream_id, topic_id, group_id) .await } + ClientWrapper::TcpTokio(client) => { + client + .delete_consumer_group(stream_id, topic_id, group_id) + .await + } } } @@ -155,6 +171,11 @@ .join_consumer_group(stream_id, topic_id, group_id) .await } + ClientWrapper::TcpTokio(client) => { + client + .join_consumer_group(stream_id, topic_id, group_id) + .await + } } } @@ -185,6 +206,11 @@ .leave_consumer_group(stream_id, topic_id, group_id) .await } + ClientWrapper::TcpTokio(client) => { + client + .leave_consumer_group(stream_id, topic_id, group_id) + .await + } } } } @@ -205,6 +231,9 @@ ClientWrapper::Quic(client) => { let _ = client.logout_user().await; } + ClientWrapper::TcpTokio(client) => { + let _ = client.logout_user().await; + } } } }
diff --git a/core/sdk/src/client_wrappers/binary_consumer_offset_client.rs b/core/sdk/src/client_wrappers/binary_consumer_offset_client.rs index 668f120..e48741e 100644 --- a/core/sdk/src/client_wrappers/binary_consumer_offset_client.rs +++ b/core/sdk/src/client_wrappers/binary_consumer_offset_client.rs
@@ -52,6 +52,11 @@ .store_consumer_offset(consumer, stream_id, topic_id, partition_id, offset) .await } + ClientWrapper::TcpTokio(client) => { + client + .store_consumer_offset(consumer, stream_id, topic_id, partition_id, offset) + .await + } } } @@ -83,6 +88,11 @@ .get_consumer_offset(consumer, stream_id, topic_id, partition_id) .await } + ClientWrapper::TcpTokio(client) => { + client + .get_consumer_offset(consumer, stream_id, topic_id, partition_id) + .await + } } } @@ -114,6 +124,11 @@ .delete_consumer_offset(consumer, stream_id, topic_id, partition_id) .await } + ClientWrapper::TcpTokio(client) => { + client + .delete_consumer_offset(consumer, stream_id, topic_id, partition_id) + .await + } } } }
diff --git a/core/sdk/src/client_wrappers/binary_message_client.rs b/core/sdk/src/client_wrappers/binary_message_client.rs index 01e0c16..3675317 100644 --- a/core/sdk/src/client_wrappers/binary_message_client.rs +++ b/core/sdk/src/client_wrappers/binary_message_client.rs
@@ -88,6 +88,19 @@ ) .await } + ClientWrapper::TcpTokio(client) => { + client + .poll_messages( + stream_id, + topic_id, + partition_id, + consumer, + strategy, + count, + auto_commit, + ) + .await + } } } @@ -119,6 +132,11 @@ .send_messages(stream_id, topic_id, partitioning, messages) .await } + ClientWrapper::TcpTokio(client) => { + client + .send_messages(stream_id, topic_id, partitioning, messages) + .await + } } } @@ -150,6 +168,11 @@ .flush_unsaved_buffer(stream_id, topic_id, partitioning_id, fsync) .await } + ClientWrapper::TcpTokio(client) => { + client + .flush_unsaved_buffer(stream_id, topic_id, partitioning_id, fsync) + .await + } } } }
diff --git a/core/sdk/src/client_wrappers/binary_partition_client.rs b/core/sdk/src/client_wrappers/binary_partition_client.rs index 26ccad1..4463d23 100644 --- a/core/sdk/src/client_wrappers/binary_partition_client.rs +++ b/core/sdk/src/client_wrappers/binary_partition_client.rs
@@ -50,6 +50,11 @@ .create_partitions(stream_id, topic_id, partitions_count) .await } + ClientWrapper::TcpTokio(client) => { + client + .create_partitions(stream_id, topic_id, partitions_count) + .await + } } } @@ -80,6 +85,11 @@ .delete_partitions(stream_id, topic_id, partitions_count) .await } + ClientWrapper::TcpTokio(client) => { + client + .delete_partitions(stream_id, topic_id, partitions_count) + .await + } } } }
diff --git a/core/sdk/src/client_wrappers/binary_personal_access_token_client.rs b/core/sdk/src/client_wrappers/binary_personal_access_token_client.rs index 58c6fbe..4a83a5a 100644 --- a/core/sdk/src/client_wrappers/binary_personal_access_token_client.rs +++ b/core/sdk/src/client_wrappers/binary_personal_access_token_client.rs
@@ -32,6 +32,7 @@ ClientWrapper::Http(client) => client.get_personal_access_tokens().await, ClientWrapper::Tcp(client) => client.get_personal_access_tokens().await, ClientWrapper::Quic(client) => client.get_personal_access_tokens().await, + ClientWrapper::TcpTokio(client) => client.get_personal_access_tokens().await, } } @@ -45,6 +46,7 @@ ClientWrapper::Http(client) => client.create_personal_access_token(name, expiry).await, ClientWrapper::Tcp(client) => client.create_personal_access_token(name, expiry).await, ClientWrapper::Quic(client) => client.create_personal_access_token(name, expiry).await, + ClientWrapper::TcpTokio(client) => client.create_personal_access_token(name, expiry).await, } } @@ -54,6 +56,7 @@ ClientWrapper::Http(client) => client.delete_personal_access_token(name).await, ClientWrapper::Tcp(client) => client.delete_personal_access_token(name).await, ClientWrapper::Quic(client) => client.delete_personal_access_token(name).await, + ClientWrapper::TcpTokio(client) => client.delete_personal_access_token(name).await, } } @@ -66,6 +69,7 @@ ClientWrapper::Http(client) => client.login_with_personal_access_token(token).await, ClientWrapper::Tcp(client) => client.login_with_personal_access_token(token).await, ClientWrapper::Quic(client) => client.login_with_personal_access_token(token).await, + ClientWrapper::TcpTokio(client) => client.login_with_personal_access_token(token).await, } } }
diff --git a/core/sdk/src/client_wrappers/binary_segment_client.rs b/core/sdk/src/client_wrappers/binary_segment_client.rs index 0ec082f..4af4fac 100644 --- a/core/sdk/src/client_wrappers/binary_segment_client.rs +++ b/core/sdk/src/client_wrappers/binary_segment_client.rs
@@ -51,6 +51,11 @@ .delete_segments(stream_id, topic_id, partition_id, segments_count) .await } + ClientWrapper::TcpTokio(client) => { + client + .delete_segments(stream_id, topic_id, partition_id, segments_count) + .await + } } } }
diff --git a/core/sdk/src/client_wrappers/binary_stream_client.rs b/core/sdk/src/client_wrappers/binary_stream_client.rs index 93d785d..443d998 100644 --- a/core/sdk/src/client_wrappers/binary_stream_client.rs +++ b/core/sdk/src/client_wrappers/binary_stream_client.rs
@@ -29,6 +29,7 @@ ClientWrapper::Http(client) => client.get_stream(stream_id).await, ClientWrapper::Tcp(client) => client.get_stream(stream_id).await, ClientWrapper::Quic(client) => client.get_stream(stream_id).await, + ClientWrapper::TcpTokio(client) => client.get_stream(stream_id).await, } } @@ -38,6 +39,7 @@ ClientWrapper::Http(client) => client.get_streams().await, ClientWrapper::Tcp(client) => client.get_streams().await, ClientWrapper::Quic(client) => client.get_streams().await, + ClientWrapper::TcpTokio(client) => client.get_streams().await, } } @@ -51,6 +53,7 @@ ClientWrapper::Http(client) => client.create_stream(name, stream_id).await, ClientWrapper::Tcp(client) => client.create_stream(name, stream_id).await, ClientWrapper::Quic(client) => client.create_stream(name, stream_id).await, + ClientWrapper::TcpTokio(client) => client.create_stream(name, stream_id).await, } } @@ -60,6 +63,7 @@ ClientWrapper::Http(client) => client.update_stream(stream_id, name).await, ClientWrapper::Tcp(client) => client.update_stream(stream_id, name).await, ClientWrapper::Quic(client) => client.update_stream(stream_id, name).await, + ClientWrapper::TcpTokio(client) => client.update_stream(stream_id, name).await, } } @@ -69,6 +73,7 @@ ClientWrapper::Http(client) => client.delete_stream(stream_id).await, ClientWrapper::Tcp(client) => client.delete_stream(stream_id).await, ClientWrapper::Quic(client) => client.delete_stream(stream_id).await, + ClientWrapper::TcpTokio(client) => client.delete_stream(stream_id).await, } } @@ -78,6 +83,7 @@ ClientWrapper::Http(client) => client.purge_stream(stream_id).await, ClientWrapper::Tcp(client) => client.purge_stream(stream_id).await, ClientWrapper::Quic(client) => client.purge_stream(stream_id).await, + ClientWrapper::TcpTokio(client) => client.purge_stream(stream_id).await, } } }
diff --git a/core/sdk/src/client_wrappers/binary_system_client.rs b/core/sdk/src/client_wrappers/binary_system_client.rs index 825dbd8..d4ab299 100644 --- a/core/sdk/src/client_wrappers/binary_system_client.rs +++ b/core/sdk/src/client_wrappers/binary_system_client.rs
@@ -32,6 +32,7 @@ ClientWrapper::Http(client) => client.get_stats().await, ClientWrapper::Tcp(client) => client.get_stats().await, ClientWrapper::Quic(client) => client.get_stats().await, + ClientWrapper::TcpTokio(client) => client.get_stats().await, } } @@ -41,6 +42,7 @@ ClientWrapper::Http(client) => client.get_me().await, ClientWrapper::Tcp(client) => client.get_me().await, ClientWrapper::Quic(client) => client.get_me().await, + ClientWrapper::TcpTokio(client) => client.get_me().await, } } @@ -50,6 +52,7 @@ ClientWrapper::Http(client) => client.get_client(client_id).await, ClientWrapper::Tcp(client) => client.get_client(client_id).await, ClientWrapper::Quic(client) => client.get_client(client_id).await, + ClientWrapper::TcpTokio(client) => client.get_client(client_id).await, } } @@ -59,6 +62,7 @@ ClientWrapper::Http(client) => client.get_clients().await, ClientWrapper::Tcp(client) => client.get_clients().await, ClientWrapper::Quic(client) => client.get_clients().await, + ClientWrapper::TcpTokio(client) => client.get_clients().await, } } @@ -68,6 +72,7 @@ ClientWrapper::Http(client) => client.ping().await, ClientWrapper::Tcp(client) => client.ping().await, ClientWrapper::Quic(client) => client.ping().await, + ClientWrapper::TcpTokio(client) => client.ping().await, } } @@ -77,6 +82,7 @@ ClientWrapper::Http(client) => client.heartbeat_interval().await, ClientWrapper::Tcp(client) => client.heartbeat_interval().await, ClientWrapper::Quic(client) => client.heartbeat_interval().await, + ClientWrapper::TcpTokio(client) => client.heartbeat_interval().await, } } @@ -90,6 +96,7 @@ ClientWrapper::Http(client) => client.snapshot(compression, snapshot_types).await, ClientWrapper::Tcp(client) => client.snapshot(compression, snapshot_types).await, ClientWrapper::Quic(client) => client.snapshot(compression, snapshot_types).await, + ClientWrapper::TcpTokio(client) => client.snapshot(compression, snapshot_types).await, } } }
diff --git a/core/sdk/src/client_wrappers/binary_topic_client.rs b/core/sdk/src/client_wrappers/binary_topic_client.rs index c9685bd..6b2fbfd 100644 --- a/core/sdk/src/client_wrappers/binary_topic_client.rs +++ b/core/sdk/src/client_wrappers/binary_topic_client.rs
@@ -35,6 +35,7 @@ ClientWrapper::Http(client) => client.get_topic(stream_id, topic_id).await, ClientWrapper::Tcp(client) => client.get_topic(stream_id, topic_id).await, ClientWrapper::Quic(client) => client.get_topic(stream_id, topic_id).await, + ClientWrapper::TcpTokio(client) => client.get_topic(stream_id, topic_id).await, } } @@ -44,6 +45,7 @@ ClientWrapper::Http(client) => client.get_topics(stream_id).await, ClientWrapper::Tcp(client) => client.get_topics(stream_id).await, ClientWrapper::Quic(client) => client.get_topics(stream_id).await, + ClientWrapper::TcpTokio(client) => client.get_topics(stream_id).await, } } @@ -115,6 +117,20 @@ ) .await } + ClientWrapper::TcpTokio(client) => { + client + .create_topic( + stream_id, + name, + partitions_count, + compression_algorithm, + replication_factor, + topic_id, + message_expiry, + max_topic_size, + ) + .await + } } } @@ -181,6 +197,19 @@ ) .await } + ClientWrapper::TcpTokio(client) => { + client + .update_topic( + stream_id, + topic_id, + name, + compression_algorithm, + replication_factor, + message_expiry, + max_topic_size, + ) + .await + } } } @@ -194,6 +223,7 @@ ClientWrapper::Http(client) => client.delete_topic(stream_id, topic_id).await, ClientWrapper::Tcp(client) => client.delete_topic(stream_id, topic_id).await, ClientWrapper::Quic(client) => client.delete_topic(stream_id, topic_id).await, + ClientWrapper::TcpTokio(client) => client.delete_topic(stream_id, topic_id).await, } } @@ -207,6 +237,7 @@ ClientWrapper::Http(client) => client.purge_topic(stream_id, topic_id).await, ClientWrapper::Tcp(client) => client.purge_topic(stream_id, topic_id).await, ClientWrapper::Quic(client) => client.purge_topic(stream_id, topic_id).await, + ClientWrapper::TcpTokio(client) => client.purge_topic(stream_id, topic_id).await, } } }
diff --git a/core/sdk/src/client_wrappers/binary_user_client.rs b/core/sdk/src/client_wrappers/binary_user_client.rs index f7289c2..a6f5cf9 100644 --- a/core/sdk/src/client_wrappers/binary_user_client.rs +++ b/core/sdk/src/client_wrappers/binary_user_client.rs
@@ -31,6 +31,7 @@ ClientWrapper::Http(client) => client.get_user(user_id).await, ClientWrapper::Tcp(client) => client.get_user(user_id).await, ClientWrapper::Quic(client) => client.get_user(user_id).await, + ClientWrapper::TcpTokio(client) => client.get_user(user_id).await, } } @@ -40,6 +41,7 @@ ClientWrapper::Http(client) => client.get_users().await, ClientWrapper::Tcp(client) => client.get_users().await, ClientWrapper::Quic(client) => client.get_users().await, + ClientWrapper::TcpTokio(client) => client.get_users().await, } } @@ -71,6 +73,11 @@ .create_user(username, password, status, permissions) .await } + ClientWrapper::TcpTokio(client) => { + client + .create_user(username, password, status, permissions) + .await + } } } @@ -80,6 +87,7 @@ ClientWrapper::Tcp(client) => client.delete_user(user_id).await, ClientWrapper::Quic(client) => client.delete_user(user_id).await, ClientWrapper::Iggy(client) => client.delete_user(user_id).await, + ClientWrapper::TcpTokio(client) => client.delete_user(user_id).await, } } @@ -94,6 +102,7 @@ ClientWrapper::Tcp(client) => client.update_user(user_id, username, status).await, ClientWrapper::Quic(client) => client.update_user(user_id, username, status).await, ClientWrapper::Iggy(client) => client.update_user(user_id, username, status).await, + ClientWrapper::TcpTokio(client) => client.update_user(user_id, username, status).await, } } @@ -107,6 +116,7 @@ ClientWrapper::Http(client) => client.update_permissions(user_id, permissions).await, ClientWrapper::Tcp(client) => client.update_permissions(user_id, permissions).await, ClientWrapper::Quic(client) => client.update_permissions(user_id, permissions).await, + ClientWrapper::TcpTokio(client) => client.update_permissions(user_id, permissions).await, } } @@ -137,6 +147,11 @@ .change_password(user_id, current_password, new_password) .await } + ClientWrapper::TcpTokio(client) => { + client + .change_password(user_id, current_password, new_password) + .await + } } } @@ -146,6 +161,7 @@ ClientWrapper::Http(client) => client.login_user(username, password).await, ClientWrapper::Tcp(client) => client.login_user(username, password).await, ClientWrapper::Quic(client) => client.login_user(username, password).await, + ClientWrapper::TcpTokio(client) => client.login_user(username, password).await, } } @@ -155,6 +171,7 @@ ClientWrapper::Http(client) => client.logout_user().await, ClientWrapper::Tcp(client) => client.logout_user().await, ClientWrapper::Quic(client) => client.logout_user().await, + ClientWrapper::TcpTokio(client) => client.logout_user().await, } } }
diff --git a/core/sdk/src/client_wrappers/client_wrapper.rs b/core/sdk/src/client_wrappers/client_wrapper.rs index d391bfd..ea5266c 100644 --- a/core/sdk/src/client_wrappers/client_wrapper.rs +++ b/core/sdk/src/client_wrappers/client_wrapper.rs
@@ -17,6 +17,7 @@ */ use crate::clients::client::IggyClient; +use crate::connection::{NewTcpClient, NewTokioTcpClient}; use crate::http::http_client::HttpClient; use crate::quic::quic_client::QuicClient; use crate::tcp::tcp_client::TcpClient; @@ -28,4 +29,5 @@ Http(HttpClient), Tcp(TcpClient), Quic(QuicClient), + TcpTokio(NewTokioTcpClient), }
diff --git a/core/sdk/src/connection/mod.rs b/core/sdk/src/connection/mod.rs new file mode 100644 index 0000000..599a636 --- /dev/null +++ b/core/sdk/src/connection/mod.rs
@@ -0,0 +1,738 @@ +use std::{ + collections::VecDeque, + fmt::Debug, + io::{self, IoSlice}, + mem::MaybeUninit, + net::SocketAddr, + ops::Deref, + pin::Pin, + sync::{ + Arc, + atomic::{AtomicU64, Ordering}, + }, + task::{Context, Poll, Waker}, +}; + +use crate::{ + connection::transport::{ClientConfig, TokioTcpTransport, Transport}, + protocol::{ControlAction, ProtocolCore, ProtocolCoreConfig, TxBuf}, + runtime::{Runtime, TokioRuntime}, +}; +use async_broadcast::{Receiver, Sender, broadcast}; +use async_trait::async_trait; +use bytes::{BufMut, Bytes, BytesMut}; +use futures::{AsyncRead, AsyncWrite, task::AtomicWaker}; +use iggy_binary_protocol::{BinaryClient, BinaryTransport, Client}; +use iggy_common::{ClientState, Command, DiagnosticEvent, IggyDuration, IggyError}; +use parking_lot::Mutex; +use rustc_hash::FxHashMap; +use tracing::{debug, error}; + +mod transport; + +pub type NewTokioTcpClient = NewTcpClient<TokioTcpTransport, TokioRuntime>; + +pub enum ClientCommand { + Connect(SocketAddr), + Disconnect, + Shutdown, +} + +#[derive(Debug)] +pub struct ConnectionInner<T: Transport, R: Runtime> { + pub(crate) state: Mutex<State<T, R>>, +} + +#[derive(Debug)] +pub struct ConnectionRef<T: Transport, R: Runtime>(Arc<ConnectionInner<T, R>>); + +impl<T: Transport, R: Runtime> ConnectionRef<T, R> { + fn new(core: ProtocolCore, cfg: Arc<T::Config>, rt: Arc<R>) -> Self { + Self(Arc::new(ConnectionInner { + state: Mutex::new(State { + rt, + inner: core, + driver: None, + stream: None, + current_send: None, + send_offset: 0, + cfg, + recv_buffer: BytesMut::with_capacity(16 * 1024), + wait_timer: None, + waiters: Arc::new(Waiters { + map: Mutex::new(FxHashMap::with_capacity_and_hasher(256, Default::default())), + next_id: AtomicU64::new(0), + }), + requests_to_wait: FxHashMap::with_capacity_and_hasher(256, Default::default()), + pending_commands: VecDeque::new(), + connect_waiters: Vec::new(), + pending_connect: None, + }), + })) + } +} + +impl<T: Transport, R: Runtime> ConnectionRef<T, R> { + fn state(&self) -> ClientState { + let state = self.0.state.lock(); + state.inner.state + } + + fn set_state(&self, client_state: ClientState) { + let mut state = self.0.state.lock(); + state.inner.state = client_state + } +} + +impl<T: Transport, R: Runtime> Deref for ConnectionRef<T, R> { + type Target = ConnectionInner<T, R>; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl<T: Transport, R: Runtime> Clone for ConnectionRef<T, R> { + fn clone(&self) -> Self { + Self(self.0.clone()) + } +} + +struct WaitEntry<T> { + waker: AtomicWaker, + result: Option<T>, +} + +struct Waiters<T> { + map: Mutex<FxHashMap<u64, WaitEntry<T>>>, + next_id: AtomicU64, +} + +impl<T> Waiters<T> { + fn alloc(&self) -> u64 { + let id = self.next_id.fetch_add(1, Ordering::Relaxed); + self.map.lock().insert( + id, + WaitEntry { + waker: AtomicWaker::new(), + result: None, + }, + ); + id + } + + fn complete(&self, id: u64, val: T) -> bool { + if let Some(entry) = self.map.lock().get_mut(&id) { + entry.result = Some(val); + entry.waker.wake(); + true + } else { + false + } + } +} + +struct WaitFuture<T> { + waiters: Arc<Waiters<T>>, + id: u64, +} + +impl<T> Future for WaitFuture<T> { + type Output = T; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<T> { + let mut map = self.waiters.map.lock(); + + let ready: Option<T> = { + if let Some(entry) = map.get_mut(&self.id) { + if let Some(val) = entry.result.take() { + Some(val) + } else { + entry.waker.register(cx.waker()); + entry.result.take() + } + } else { + None + } + }; + + if let Some(val) = ready { + map.remove(&self.id); + Poll::Ready(val) + } else { + Poll::Pending + } + } +} + +pub struct State<T: Transport, R: Runtime> { + rt: Arc<R>, + inner: ProtocolCore, + driver: Option<Waker>, + stream: Option<T::Stream>, + current_send: Option<TxBuf>, + send_offset: usize, + recv_buffer: BytesMut, + cfg: Arc<T::Config>, + + wait_timer: Option<Pin<Box<R::Sleep>>>, + + waiters: Arc<Waiters<Result<Bytes, IggyError>>>, + requests_to_wait: FxHashMap<u64, u64>, + pending_commands: VecDeque<(u64, ClientCommand)>, + connect_waiters: Vec<u64>, + pending_connect: Option<Pin<Box<dyn Future<Output = io::Result<T::Stream>> + Send>>>, +} + +impl<T: Transport, R: Runtime> Debug for State<T, R> { + // todo implement debug + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "test") + } +} + +impl<T: Transport, R: Runtime> State<T, R> { + fn take_waker(&mut self) -> Option<Waker> { + self.driver.take() + } + + fn complete_all_waiters_with_error(&mut self, error: IggyError) { + for id in self.connect_waiters.drain(..) { + let _ = self.waiters.complete(id, Err(error.clone())); + } + + for (_request_id, wait_id) in self.requests_to_wait.drain() { + let _ = self.waiters.complete(wait_id, Err(error.clone())); + } + + for (wait_id, _cmd) in self.pending_commands.drain(..) { + let _ = self.waiters.complete(wait_id, Err(error.clone())); + } + } + + fn enqueu_command( + &mut self, + command: ClientCommand, + ) -> (WaitFuture<Result<Bytes, IggyError>>, Option<Waker>) { + let id = self.waiters.alloc(); + self.pending_commands.push_back((id, command)); + let waker = self.take_waker(); + ( + WaitFuture { + waiters: self.waiters.clone(), + id, + }, + waker, + ) + } + + fn enqueue_message( + &mut self, + code: u32, + payload: Bytes, + ) -> (WaitFuture<Result<Bytes, IggyError>>, Option<Waker>) { + let wait_id = self.waiters.alloc(); + let waker = match self.inner.send(code, payload) { + Ok(protocol_id) => { + self.requests_to_wait.insert(protocol_id, wait_id); + self.take_waker() + } + Err(e) => { + self.waiters.complete(wait_id, Err(e)); + None + } + }; + ( + WaitFuture { + waiters: self.waiters.clone(), + id: wait_id, + }, + waker, + ) + } + + fn drive_client_commands(&mut self) -> io::Result<bool> { + let mut made_progress = false; + for (request_id, cmd) in self.pending_commands.drain(..) { + made_progress = true; + match cmd { + ClientCommand::Connect(server_address) => { + debug!( + "ConnectionDriver: Processing Connect command to {}", + server_address + ); + let current_state = self.inner.state; + + if matches!( + current_state, + ClientState::Connected + | ClientState::Authenticating + | ClientState::Authenticated + ) { + debug!( + "ConnectionDriver: Already connected (state: {:?}), completing waiter immediately", + current_state + ); + let _ = self.waiters.complete(request_id, Ok(Bytes::new())); + continue; + } + + if matches!(current_state, ClientState::Connecting) { + debug!("ConnectionDriver: Already connecting, adding to waiters"); + self.connect_waiters.push(request_id); + continue; + } + + self.connect_waiters.push(request_id); + self.inner.desire_connect(server_address).map_err(|e| { + error!("ConnectionDriver: desire_connect failed: {}", e.as_string()); + io::Error::new(io::ErrorKind::ConnectionAborted, e.as_string()) + })?; + debug!( + "ConnectionDriver: desire_connect successful, state: {:?}", + self.inner.state + ); + } + ClientCommand::Disconnect => { + self.inner.disconnect(); + self.waiters.complete(request_id, Ok(Bytes::new())); + } + ClientCommand::Shutdown => { + self.inner.shutdown(); + self.waiters.complete(request_id, Ok(Bytes::new())); + } + } + } + Ok(made_progress) + } + + fn drive_connect(&mut self, cx: &mut Context<'_>) -> io::Result<bool> { + if let Some(fut) = self.pending_connect.as_mut() { + match fut.as_mut().poll(cx) { + Poll::Pending => return Ok(false), + Poll::Ready(Ok(stream)) => { + self.stream = Some(stream); + self.pending_connect = None; + self.inner.on_connected().map_err(|e| { + io::Error::new(io::ErrorKind::ConnectionRefused, e.as_string()) + })?; + debug!( + "ConnectionDriver: Connection established, completing {} waiters", + self.connect_waiters.len() + ); + if !self.inner.should_wait_auth() { + for id in self.connect_waiters.drain(..) { + debug!("ConnectionDriver: Completing connect waiter {}", id); + let _ = self.waiters.complete(id, Ok(Bytes::new())); + } + } + return Ok(true); + } + Poll::Ready(Err(_e)) => { + self.pending_connect = None; + self.inner.disconnect(); + for id in self.connect_waiters.drain(..) { + let _ = self + .waiters + .complete(id, Err(IggyError::CannotEstablishConnection)); + } + return Ok(true); + } + } + } + Ok(false) + } + + fn drive_timer(&mut self, cx: &mut Context<'_>) -> bool { + if let Some(t) = &mut self.wait_timer { + if t.as_mut().poll(cx).is_pending() { + return false; + } + self.wait_timer = None; + return true; + } + false + } + + fn drive_transmit(&mut self, cx: &mut Context<'_>) -> io::Result<bool> { + if self.current_send.is_none() { + if let Some(tx) = self.inner.poll_transmit() { + self.current_send = Some(tx); + self.send_offset = 0; + } else { + return Ok(false); + } + } + + let stream = self + .stream + .as_mut() + .ok_or_else(|| io::Error::new(io::ErrorKind::NotConnected, "No stream"))?; + + let buf = self.current_send.as_ref().unwrap(); + let mut offset = self.send_offset; + + while offset < buf.total_len() { + let mut storage = [IoSlice::new(&[]), IoSlice::new(&[])]; + + let iov = if offset < 8 { + storage[0] = IoSlice::new(&buf.header[offset..]); + if !buf.payload.is_empty() { + storage[1] = IoSlice::new(&buf.payload); + &storage[..2] + } else { + &storage[..1] + } + } else { + let body_off = offset - 8; + storage[0] = IoSlice::new(&buf.payload[body_off..]); + &storage[..1] + }; + + let written = match Pin::new(&mut *stream).poll_write_vectored(cx, iov)? { + Poll::Ready(0) => { + return Err(io::Error::new( + io::ErrorKind::WriteZero, + "write returned 0 bytes", + )); + } + Poll::Ready(n) => n, + Poll::Pending => return Ok(false), + }; + + offset += written; + self.send_offset += written; + } + + match Pin::new(stream).poll_flush(cx)? { + Poll::Pending => return Ok(false), + Poll::Ready(()) => {} + } + + self.send_offset = 0; + self.current_send = None; + Ok(true) + } + + fn drive_receive(&mut self, cx: &mut Context<'_>) -> io::Result<bool> { + let mut progress = false; + + // todo add some const var + for _ in 0..16 { + if self.recv_buffer.spare_capacity_mut().is_empty() { + self.recv_buffer.reserve(8192); + } + + let spare: &mut [MaybeUninit<u8>] = self.recv_buffer.spare_capacity_mut(); + + let buf: &mut [u8] = unsafe { &mut *(spare as *mut [MaybeUninit<u8>] as *mut [u8]) }; + + let n = { + let stream = self + .stream + .as_mut() + .ok_or(io::Error::new(io::ErrorKind::NotConnected, "No stream"))?; + match Pin::new(&mut *stream).poll_read(cx, buf)? { + Poll::Pending => return Ok(progress), + Poll::Ready(0) => { + self.inner.disconnect(); + self.stream = None; + self.complete_all_waiters_with_error(IggyError::CannotEstablishConnection); + return Ok(true); + } + Poll::Ready(n) => n, + } + }; + + unsafe { + self.recv_buffer.advance_mut(n); + } + + self.inner + .process_incoming_with(&mut self.recv_buffer, |req_id, status, payload| { + if let Some(wait_id) = self.requests_to_wait.remove(&req_id) { + let res = if status == 0 { + Ok(payload) + } else { + Err(IggyError::from_code(status)) + }; + let _ = self.waiters.complete(wait_id, res); + } + }); + progress = true; + } + + if let Some(auth_res) = self.inner.take_auth_result() { + match auth_res { + Ok(()) => { + for id in self.connect_waiters.drain(..) { + let _ = self.waiters.complete(id, Ok(Bytes::new())); + } + } + Err(e) => { + for id in self.connect_waiters.drain(..) { + let _ = self.waiters.complete(id, Err(e.clone())); + } + } + } + progress = true; + } + + Ok(progress) + } +} + +struct ConnectionDriver<T: Transport, R: Runtime>(ConnectionRef<T, R>); + +impl<T: Transport, R: Runtime> Future for ConnectionDriver<T, R> { + type Output = Result<(), io::Error>; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> { + let st = &mut *self.0.state.lock(); + let mut keep_going = false; + + keep_going |= st.drive_client_commands()?; + + let order = st.inner.poll(); + match order { + ControlAction::Wait(dur) => { + if st.wait_timer.is_none() { + st.wait_timer = Some(Box::pin(st.rt.sleep(dur))); + } + keep_going |= st.drive_timer(cx); + } + ControlAction::Connect(server_adress) => { + if st.pending_connect.is_none() { + st.pending_connect = Some(T::connect(st.cfg.clone(), server_adress)); + } + keep_going |= st.drive_connect(cx)?; + } + ControlAction::Noop | ControlAction::Authenticate { .. } => {} + ControlAction::Error(e) => { + if matches!(e, IggyError::ClientShutdown) { + debug!("ConnectionDriver: Received ClientShutdown, terminating gracefully"); + return Poll::Ready(Ok(())); + } + error!("ConnectionDriver: Error - {e:?}"); + return Poll::Ready(Err(io::Error::new(io::ErrorKind::Other, format!("{e:?}")))); + } + } + + if st.stream.is_some() { + match st.drive_transmit(cx) { + Ok(progress) => keep_going |= progress, + Err(e) => { + debug!("Transmit error, disconnecting: {:?}", e); + st.inner.disconnect(); + st.stream = None; + st.complete_all_waiters_with_error(IggyError::CannotEstablishConnection); + keep_going = true; + } + } + + if st.stream.is_some() { + match st.drive_receive(cx) { + Ok(progress) => keep_going |= progress, + Err(e) => { + debug!("Receive error, disconnecting: {:?}", e); + st.inner.disconnect(); + st.stream = None; + st.complete_all_waiters_with_error(IggyError::CannotEstablishConnection); + keep_going = true; + } + } + } + } + + let current_state = st.inner.state; + if current_state == ClientState::Shutdown { + if st.stream.is_some() { + debug!("ConnectionDriver: Closing stream due to Shutdown"); + st.stream = None; + } + return Poll::Ready(Ok(())); + } + + if current_state == ClientState::Disconnected && st.stream.is_some() { + debug!("ConnectionDriver: Closing stream due to Disconnected but keeping driver alive"); + st.stream = None; + } + + st.driver = Some(cx.waker().clone()); + + if keep_going { + cx.waker().wake_by_ref(); + } + + Poll::Pending + } +} + +#[derive(Debug)] +pub struct NewTcpClient<T: Transport, R: Runtime> { + state: ConnectionRef<T, R>, + config: Arc<T::Config>, + events: (Sender<DiagnosticEvent>, Receiver<DiagnosticEvent>), + _driver_handle: R::Join, +} + +impl<T: Transport, R: Runtime + 'static> NewTcpClient<T, R> { + pub fn create(config: Arc<T::Config>, rt: Arc<R>) -> Result<Self, IggyError> { + let (tx, rx) = broadcast(1000); + + let proto_config = ProtocolCoreConfig { + auto_login: config.auto_login(), + reestablish_after: config.reconnection_reestablish_after(), + max_retries: config.reconnection_max_retries(), + }; + + let conn = ConnectionRef::new(ProtocolCore::new(proto_config), config.clone(), rt.clone()); + let driver = ConnectionDriver(conn.clone()); + let driver_handle = rt.spawn(async move { + if let Err(e) = driver.await { + error!("I/O error: {e}"); + } + }); + + Ok(Self { + state: conn, + config, + events: (tx, rx), + _driver_handle: driver_handle, + }) + } + + async fn send_raw(&self, code: u32, payload: Bytes) -> Result<Bytes, IggyError> { + let (wait_future, waker) = { + let mut state = self.state.0.state.lock(); + state.enqueue_message(code, payload) + }; + if let Some(waker) = waker { + waker.wake(); + } + wait_future.await + } +} + +#[async_trait] +impl<T, R> Client for NewTcpClient<T, R> +where + T: Transport + Debug, + R: Runtime + Debug + Send + Sync + 'static, + T::Config: ClientConfig + Debug + Send + Sync + 'static, + R::Join: Debug + Send + Sync + 'static, +{ + async fn connect(&self) -> Result<(), IggyError> { + let address = self.config.server_address(); + let (fut, waker) = { + let mut state = self.state.0.state.lock(); + state.enqueu_command(ClientCommand::Connect(address)) + }; + if let Some(waker) = waker { + waker.wake(); + } + + match fut.await { + Ok(_) => { + self.publish_event(DiagnosticEvent::Connected).await; + Ok(()) + } + Err(IggyError::CannotEstablishConnection) => { + self.publish_event(DiagnosticEvent::Disconnected).await; + Err(IggyError::CannotEstablishConnection) + } + Err(e) => { + error!("Got error: {e} on connect"); + Err(e) + } + } + } + + async fn disconnect(&self) -> Result<(), IggyError> { + let (fut, waker) = { + let mut state = self.state.0.state.lock(); + state.enqueu_command(ClientCommand::Disconnect) + }; + if let Some(waker) = waker { + waker.wake(); + } + fut.await?; + self.publish_event(DiagnosticEvent::Disconnected).await; + Ok(()) + } + + async fn shutdown(&self) -> Result<(), IggyError> { + let (fut, waker) = { + let mut state = self.state.0.state.lock(); + state.enqueu_command(ClientCommand::Shutdown) + }; + if let Some(waker) = waker { + waker.wake(); + } + fut.await?; + self.publish_event(DiagnosticEvent::Shutdown).await; + Ok(()) + } + + async fn subscribe_events(&self) -> Receiver<DiagnosticEvent> { + self.events.1.clone() + } +} + +#[async_trait] +impl<T, R> BinaryTransport for NewTcpClient<T, R> +where + T: Transport + Debug, + R: Runtime + Debug + Send + Sync + 'static, + T::Config: ClientConfig + Debug + Send + Sync + 'static, + R::Join: Debug + Send + Sync + 'static, +{ + async fn get_state(&self) -> ClientState { + self.state.state() + } + + async fn set_state(&self, state: ClientState) { + self.state.set_state(state); + } + + async fn publish_event(&self, event: DiagnosticEvent) { + if let Err(error) = self.events.0.broadcast(event).await { + error!("Failed to send a TCP diagnostic event: {error}"); + } + } + + async fn send_with_response<C: Command>(&self, command: &C) -> Result<Bytes, IggyError> { + command.validate()?; + self.send_raw_with_response(command.code(), command.to_bytes()) + .await + } + + async fn send_raw_with_response(&self, code: u32, payload: Bytes) -> Result<Bytes, IggyError> { + self.send_raw(code, payload).await + } + + fn get_heartbeat_interval(&self) -> IggyDuration { + self.config.heartbeat_interval() + } +} + +impl<T, R> BinaryClient for NewTcpClient<T, R> +where + T: Transport + Debug, + R: Runtime + Debug + Send + Sync + 'static, + T::Config: ClientConfig + Debug + Send + Sync + 'static, + R::Join: Debug + Send + Sync + 'static, +{ +} + +impl<T: Transport, R: Runtime> Drop for NewTcpClient<T, R> { + fn drop(&mut self) { + let mut state = self.state.0.state.lock(); + state.inner.disconnect(); + let waker = state.take_waker(); + drop(state); + if let Some(waker) = waker { + waker.wake(); + } + } +}
diff --git a/core/sdk/src/connection/transport.rs b/core/sdk/src/connection/transport.rs new file mode 100644 index 0000000..2ab573e --- /dev/null +++ b/core/sdk/src/connection/transport.rs
@@ -0,0 +1,73 @@ +use std::{io, net::SocketAddr, pin::Pin, str::FromStr, sync::Arc}; + +use futures::{AsyncRead, AsyncWrite}; +use iggy_common::{AutoLogin, IggyDuration, TcpClientConfig}; +use tokio::net::TcpSocket; +use tokio_util::compat::{Compat, TokioAsyncReadCompatExt}; + +pub trait ClientConfig { + fn server_address(&self) -> SocketAddr; + fn auto_login(&self) -> AutoLogin; + fn reconnection_reestablish_after(&self) -> IggyDuration; + fn reconnection_max_retries(&self) -> Option<u32>; + fn heartbeat_interval(&self) -> IggyDuration; +} + +impl ClientConfig for TcpClientConfig { + fn auto_login(&self) -> AutoLogin { + self.auto_login.clone() + } + + fn heartbeat_interval(&self) -> IggyDuration { + self.heartbeat_interval + } + + fn reconnection_max_retries(&self) -> Option<u32> { + self.reconnection.max_retries + } + + fn reconnection_reestablish_after(&self) -> IggyDuration { + self.reconnection.reestablish_after + } + + fn server_address(&self) -> SocketAddr { + SocketAddr::from_str(&self.server_address).unwrap() + } +} + +pub trait Transport: Send + Sync + 'static { + type Stream: AsyncRead + AsyncWrite + Unpin + Send + 'static; + type Config: ClientConfig + Clone + Send + Sync + 'static; + + fn connect( + cfg: Arc<Self::Config>, + server_address: SocketAddr, + ) -> Pin<Box<dyn Future<Output = io::Result<Self::Stream>> + Send>>; +} + +#[derive(Debug)] +pub struct TokioTcpTransport; + +impl Transport for TokioTcpTransport { + type Stream = Compat<tokio::net::TcpStream>; + type Config = TcpClientConfig; + + fn connect( + cfg: Arc<Self::Config>, + server_address: SocketAddr, + ) -> Pin<Box<dyn Future<Output = io::Result<Self::Stream>> + Send>> { + let nodelay = cfg.nodelay; + + Box::pin(async move { + let sock = match server_address { + std::net::SocketAddr::V4(_) => TcpSocket::new_v4()?, + _ => TcpSocket::new_v6()?, + }; + if nodelay { + sock.set_nodelay(true).ok(); + } + let s = sock.connect(server_address).await?; + Ok(s.compat()) + }) + } +}
diff --git a/core/sdk/src/lib.rs b/core/sdk/src/lib.rs index c315fd5..06322de 100644 --- a/core/sdk/src/lib.rs +++ b/core/sdk/src/lib.rs
@@ -20,9 +20,12 @@ pub mod client_provider; pub mod client_wrappers; pub mod clients; +pub mod connection; pub mod consumer_ext; pub mod http; pub mod prelude; +pub mod protocol; pub mod quic; +pub mod runtime; pub mod stream_builder; pub mod tcp;
diff --git a/core/sdk/src/protocol/mod.rs b/core/sdk/src/protocol/mod.rs new file mode 100644 index 0000000..95a8f69 --- /dev/null +++ b/core/sdk/src/protocol/mod.rs
@@ -0,0 +1,280 @@ +use std::{ + collections::VecDeque, + net::SocketAddr, +}; + +use bytes::{Buf, BufMut, Bytes, BytesMut}; +use iggy_common::{AutoLogin, ClientState, Credentials, IggyDuration, IggyError, IggyTimestamp}; +use tracing::{debug, info, warn}; + +const REQUEST_INITIAL_BYTES_LENGTH: usize = 4; +const REQUEST_HEADER_BYTES: usize = 8; +const RESPONSE_INITIAL_BYTES_LENGTH: usize = 8; +#[derive(Debug)] +pub struct ProtocolCoreConfig { + pub auto_login: AutoLogin, + pub reestablish_after: IggyDuration, + pub max_retries: Option<u32>, +} + +#[derive(Debug)] +pub enum ControlAction { + Connect(SocketAddr), + Wait(IggyDuration), + Authenticate { username: String, password: String }, + Noop, + Error(IggyError), +} + +pub struct TxBuf { + pub header: [u8; 8], + pub payload: Bytes, + pub request_id: u64, +} + +impl TxBuf { + #[inline] + pub fn total_len(&self) -> usize { + REQUEST_HEADER_BYTES + self.payload.len() + } +} + +#[derive(Debug)] +pub struct ProtocolCore { + pub state: ClientState, + config: ProtocolCoreConfig, + last_connect_attempt: Option<IggyTimestamp>, + pub retry_count: u32, + next_request_id: u64, + pending_sends: VecDeque<(u32, Bytes, u64)>, + sent_order: VecDeque<u64>, + auth_pending: bool, + auth_request_id: Option<u64>, + server_address: Option<SocketAddr>, + last_auth_result: Option<Result<(), IggyError>>, +} + +impl ProtocolCore { + pub fn new(config: ProtocolCoreConfig) -> Self { + Self { + state: ClientState::Disconnected, + config, + last_connect_attempt: None, + retry_count: 0, + next_request_id: 1, + pending_sends: VecDeque::new(), + sent_order: VecDeque::new(), + auth_pending: false, + auth_request_id: None, + server_address: None, + last_auth_result: None, + } + } + + pub fn poll_transmit(&mut self) -> Option<TxBuf> { + if let Some((code, payload, request_id)) = self.pending_sends.pop_front() { + let total_len = (payload.len() + REQUEST_INITIAL_BYTES_LENGTH) as u32; + self.sent_order.push_back(request_id); + + Some(TxBuf { + payload, + header: make_header(total_len, code), + request_id, + }) + } else { + None + } + } + + pub fn send(&mut self, code: u32, payload: Bytes) -> Result<u64, IggyError> { + match self.state { + ClientState::Shutdown => Err(IggyError::ClientShutdown), + ClientState::Disconnected | ClientState::Connecting => Err(IggyError::NotConnected), + ClientState::Connected | ClientState::Authenticating | ClientState::Authenticated => { + Ok(self.queue_send(code, payload)) + } + } + } + + fn queue_send(&mut self, code: u32, payload: Bytes) -> u64 { + let request_id = self.next_request_id; + self.next_request_id += 1; + self.pending_sends.push_back((code, payload, request_id)); + request_id + } + + pub fn process_incoming_with<F: FnMut(u64, u32, Bytes)>( + &mut self, + buf: &mut BytesMut, + mut f: F, + ) { + loop { + if buf.len() < RESPONSE_INITIAL_BYTES_LENGTH { + break; + } + let status = u32::from_le_bytes(buf[..4].try_into().unwrap()); + let length = u32::from_le_bytes(buf[4..8].try_into().unwrap()); + let total = RESPONSE_INITIAL_BYTES_LENGTH + length as usize; + if buf.len() < total { + break; + } + + buf.advance(RESPONSE_INITIAL_BYTES_LENGTH); + let payload = if length <= 1 { + Bytes::new() + } else { + buf.split_to(length as usize).freeze() + }; + if let Some(id) = self.on_response(status) { + f(id, status, payload); + } + } + } + + pub fn on_response(&mut self, status: u32) -> Option<u64> { + let request_id = self.sent_order.pop_front()?; + + if Some(request_id) == self.auth_request_id { + if status == 0 { + debug!("Authentication successful"); + self.state = ClientState::Authenticated; + self.auth_pending = false; + self.last_auth_result = Some(Ok(())); + } else { + warn!("Authentication failed with status: {}", status); + self.state = ClientState::Connected; + self.auth_pending = false; + self.last_auth_result = Some(Err(IggyError::Unauthenticated)); + } + self.auth_request_id = None; + } + + Some(request_id) + } + + pub fn poll(&mut self) -> ControlAction { + match self.state { + ClientState::Shutdown => ControlAction::Error(IggyError::ClientShutdown), + ClientState::Disconnected => ControlAction::Noop, + ClientState::Authenticated | ClientState::Authenticating | ClientState::Connected => { + ControlAction::Noop + } + ClientState::Connecting => { + let server_address = match self.server_address { + Some(addr) => addr, + None => return ControlAction::Error(IggyError::ConnectionMissedSocket), + }; + + if let Some(last) = self.last_connect_attempt { + let now = IggyTimestamp::now(); + let elapsed = now.as_micros().saturating_sub(last.as_micros()); + if elapsed < self.config.reestablish_after.as_micros() { + let remaining = + IggyDuration::from(self.config.reestablish_after.as_micros() - elapsed); + return ControlAction::Wait(remaining); + } + } + + if let Some(max_retries) = self.config.max_retries { + if self.retry_count >= max_retries { + return ControlAction::Error(IggyError::MaxRetriesExceeded); + } + } + + self.retry_count += 1; + self.last_connect_attempt = Some(IggyTimestamp::now()); + + return ControlAction::Connect(server_address); + } + } + } + + pub fn desire_connect(&mut self, server_address: SocketAddr) -> Result<(), IggyError> { + match self.state { + ClientState::Shutdown => return Err(IggyError::ClientShutdown), + ClientState::Connecting => return Ok(()), + ClientState::Connected | ClientState::Authenticating | ClientState::Authenticated => { + return Ok(()); + } + _ => { + self.state = ClientState::Connecting; + self.server_address = Some(server_address); + } + } + + Ok(()) + } + + pub fn on_connected(&mut self) -> Result<(), IggyError> { + debug!("Transport connected"); + if self.state != ClientState::Connecting { + return Err(IggyError::IncorrectConnectionState); + } + self.state = ClientState::Connected; + self.retry_count = 0; + + match &self.config.auto_login { + AutoLogin::Disabled => { + info!("Automatic sign-in is disabled."); + } + AutoLogin::Enabled(credentials) => { + if !self.auth_pending { + self.state = ClientState::Authenticating; + self.auth_pending = true; + + match credentials { + Credentials::UsernamePassword(username, password) => { + let auth_payload = encode_auth(&username, &password); + let auth_id = self.queue_send(0x0A, auth_payload); + self.auth_request_id = Some(auth_id); + } + _ => { + todo!("add PersonalAccessToken") + } + } + } + } + } + + Ok(()) + } + + pub fn disconnect(&mut self) { + debug!("Transport disconnected"); + self.state = ClientState::Disconnected; + self.auth_pending = false; + self.auth_request_id = None; + self.sent_order.clear(); + } + + pub fn shutdown(&mut self) { + self.state = ClientState::Shutdown; + self.auth_pending = false; + self.auth_request_id = None; + self.sent_order.clear(); + } + + pub fn should_wait_auth(&self) -> bool { + matches!(self.config.auto_login, AutoLogin::Enabled(_)) && self.auth_pending + } + + pub fn take_auth_result(&mut self) -> Option<Result<(), IggyError>> { + self.last_auth_result.take() + } +} + +fn encode_auth(username: &str, password: &str) -> Bytes { + let mut buf = BytesMut::new(); + buf.put_u32_le(username.len() as u32); + buf.put_slice(username.as_bytes()); + buf.put_u32_le(password.len() as u32); + buf.put_slice(password.as_bytes()); + buf.freeze() +} + +fn make_header(total_len: u32, code: u32) -> [u8; 8] { + let mut h = [0u8; 8]; + h[..4].copy_from_slice(&total_len.to_le_bytes()); + h[4..].copy_from_slice(&code.to_le_bytes()); + h +}
diff --git a/core/sdk/src/runtime/mod.rs b/core/sdk/src/runtime/mod.rs new file mode 100644 index 0000000..eb9f2a1 --- /dev/null +++ b/core/sdk/src/runtime/mod.rs
@@ -0,0 +1,27 @@ +use iggy_common::IggyDuration; +use std::fmt::Debug; +use tokio::{task::JoinHandle, time::Sleep}; + +pub trait Runtime: Send + Sync + Debug { + type Join: Send + 'static; + type Sleep: Future<Output = ()> + Send + 'static; + + fn spawn(&self, fut: impl Future<Output = ()> + Send + 'static) -> Self::Join; + fn sleep(&self, dur: IggyDuration) -> Self::Sleep; +} + +#[derive(Debug)] +pub struct TokioRuntime {} + +impl Runtime for TokioRuntime { + type Join = JoinHandle<()>; + type Sleep = Sleep; + + fn spawn(&self, fut: impl Future<Output = ()> + Send + 'static) -> Self::Join { + tokio::spawn(fut) + } + + fn sleep(&self, dur: IggyDuration) -> Sleep { + tokio::time::sleep(dur.get_duration()) + } +}