Skip to main content

hydra_sync/
client.rs

1//! Producer/Consumer client for the Hydra protocol, connects to the server, performs handshake, and sends/receives encrypted packets (AES-GCM).
2use crate::BUFFER_SIZE;
3use crate::crypto::{NONCE_LEN, TAG_LEN};
4use crate::protocol::{
5    Role, StatusCode, perform_client_handshake, read_decrypt_packet, read_server_metrics,
6    read_server_read_write_len, read_status_code, write_encrypt_packet, write_join_header,
7};
8
9use anyhow::{Result, bail};
10use bytes::BytesMut;
11use colored::Colorize;
12use std::marker::PhantomData;
13use tokio::io::{AsyncWriteExt, BufReader, BufWriter};
14use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
15use tokio::net::{TcpStream, ToSocketAddrs};
16
17/// Unit structs to represent the `Role::Producer` at the type level for better safety and clarity.
18pub struct Producer;
19/// Unit structs to represent the `Role::Consumer` at the type level for better safety and clarity.
20pub struct Consumer;
21/// Unit structs to represent the `Role::Observer` at the type level for better safety and clarity.
22pub struct Observer;
23
24/// **HydraClient** - `Cross-lang-friendly` client connects to the [`HydraServer`](crate::server::HydraServer) as a producer or consumer,
25/// performs handshake, and sends/receives encrypted packet `(AES-GCM)`. It maintains an internal memory buffer (12 Mb) for zero-copy crypto and buffering.
26///
27/// ```no_run
28/// use hydra_sync::client::{HydraClient, Producer, Consumer};
29///
30/// #[tokio::main]
31/// async fn main() {
32///     let addr = "127.0.0.1:8000";
33///     let session_uuid = [0xFFu8; 64]; // unique spmc session in server map
34///     let session_key = [0xAAu8; 32]; // aes-gcm key for encrypting packets
35///
36///     let mut producer = HydraClient::<Producer>::connect(addr, session_uuid, session_key).await.unwrap();
37///     producer.broadcast(b"I luv you >.<").await.unwrap(); // sends to all consumer
38///
39///     let mut consumer = HydraClient::<Consumer>::connect(addr, session_uuid, session_key).await.unwrap();
40///     consumer.recv().await.unwrap(); // recv whatever next frame on ring buf
41/// }
42/// ```
43///
44pub struct HydraClient<R> {
45    buf_reader: BufReader<OwnedReadHalf>,
46    buf_writer: BufWriter<OwnedWriteHalf>,
47    scratch_buf: BytesMut,
48    session_key: [u8; 32],
49    transport_key: [u8; 32],
50    server_read_write_len: u64,
51    closed: bool,
52    _role: PhantomData<R>,
53}
54
55impl HydraClient<Producer> {
56    /// Connects to the server, performs handshake, sends encrypted `JoinHeader`, decrypts/reads the `StatusCode` response,
57    /// and returns a `HydraClient<Producer>` instance if successful.
58    pub async fn connect<A: ToSocketAddrs>(
59        addr: A,
60        session_uuid: [u8; 64],
61        session_key: [u8; 32],
62    ) -> Result<Self> {
63        let stream = TcpStream::connect(addr).await?;
64        stream.set_nodelay(true)?;
65        let (reader, writer) = stream.into_split();
66        let mut writer = BufWriter::with_capacity(BUFFER_SIZE, writer);
67        let mut reader = BufReader::with_capacity(BUFFER_SIZE, reader);
68        let transport_key = perform_client_handshake(&mut reader, &mut writer).await?;
69        let mut scratch_buf = BytesMut::with_capacity(1024 * 1024 * 12);
70        write_join_header(
71            &mut writer,
72            Role::Producer,
73            session_uuid,
74            &transport_key,
75            &mut scratch_buf,
76        )
77        .await?;
78
79        let status_code = read_status_code(&mut reader, &transport_key, &mut scratch_buf).await?;
80        match status_code {
81            StatusCode::Success => {}
82            StatusCode::ErrSessionAlreadyOccupied => {
83                bail!("ERR: Session already exists, cannot join as producer");
84            }
85            StatusCode::ErrInternalServerError => {
86                bail!("ERR: Internal server error, cannot join as producer");
87            }
88            _ => {
89                unreachable!()
90            }
91        }
92
93        let server_packet_write_len =
94            read_server_read_write_len(&mut reader, &transport_key, &mut scratch_buf).await?;
95
96        Ok(Self {
97            buf_reader: reader,
98            buf_writer: writer,
99            scratch_buf,
100            session_key,
101            transport_key,
102            server_read_write_len: server_packet_write_len,
103            closed: false,
104            _role: PhantomData,
105        })
106    }
107
108    /// Broadcasts the given data as `fixed encrypted packets` (AES-GCM256) to all connected consumers **(Zero-copy)
109    /// and the [HydraServer](crate::server::HydraServer) does not decrypt it.**.
110    ///
111    /// **The caller must ensure `data.len()` is an exact multiple of the server's payload capacity `(server_read_write_len i.e. raw bytes per packet, AES-GCM (28 bytes) overhead is handled internally)`, or this will return an error.**
112    ///
113    /// > This may block until the packets are written to the server ring buffer, this behavior is configured by [`ChannelOverflowStrategy`](crate::ChannelOverflowStrategy).
114    pub async fn broadcast(&mut self, data: &[u8]) -> Result<()> {
115        let plaintext_len = self.server_read_write_len as usize; // get actual data
116        if plaintext_len == 0 || !data.len().is_multiple_of(plaintext_len) {
117            bail!(
118                "Data length {} is not an exact multiple of plaintext length {}",
119                data.len(),
120                plaintext_len
121            );
122        }
123
124        // write all chunks sequentially
125        for chunk in data.chunks_exact(plaintext_len) {
126            // this writes data.len() + crypto overhead (28 bytes)
127            write_encrypt_packet(
128                &mut self.buf_writer,
129                chunk,
130                &self.session_key,
131                &mut self.scratch_buf,
132            )
133            .await?;
134        }
135        Ok(())
136    }
137}
138
139impl HydraClient<Consumer> {
140    /// Connects to the server, performs handshake, sends encrypted `JoinHeader`, decrypts/reads the `StatusCode` response,
141    /// and returns a `HydraClient<Consumer>` instance if successful.
142    pub async fn connect<A: ToSocketAddrs>(
143        addr: A,
144        session_uuid: [u8; 64],
145        session_key: [u8; 32],
146    ) -> Result<Self> {
147        let stream = TcpStream::connect(addr).await?;
148        stream.set_nodelay(true)?;
149        let (reader, writer) = stream.into_split();
150        let mut writer = BufWriter::with_capacity(BUFFER_SIZE, writer);
151        let mut reader = BufReader::with_capacity(BUFFER_SIZE, reader);
152
153        let transport_key = perform_client_handshake(&mut reader, &mut writer).await?;
154        let mut scratch_buf = BytesMut::with_capacity(1024 * 1024 * 12);
155        write_join_header(
156            &mut writer,
157            Role::Consumer,
158            session_uuid,
159            &transport_key,
160            &mut scratch_buf,
161        )
162        .await?;
163
164        let status_code = read_status_code(&mut reader, &transport_key, &mut scratch_buf).await?;
165        match status_code {
166            StatusCode::Success => {}
167            StatusCode::ErrSessionNotFound => {
168                bail!("ERR: Session not found, cannot join as consumer");
169            }
170            StatusCode::ErrInternalServerError => {
171                bail!("ERR: Internal server error, cannot join as consumer");
172            }
173            _ => {
174                unreachable!()
175            }
176        }
177
178        let server_read_capacity =
179            read_server_read_write_len(&mut reader, &transport_key, &mut scratch_buf).await?;
180
181        Ok(Self {
182            buf_reader: reader,
183            buf_writer: writer,
184            scratch_buf,
185            session_key,
186            transport_key,
187            server_read_write_len: server_read_capacity,
188            closed: false,
189            _role: PhantomData,
190        })
191    }
192
193    /// Receives the next `fixed encrypted packet` from the producer, decrypts it, and returns the plaintext data as a byte slice.
194    /// The returned slice is valid until the next call to `recv` or `broadcast`, which reuse the internal scratch buffer.
195    /// > This may block until the next packet is available on server consumer's queue.
196    pub async fn recv(&mut self) -> Result<&[u8]> {
197        let decrypted = read_decrypt_packet(
198            &mut self.buf_reader,
199            self.server_read_write_len as usize + NONCE_LEN + TAG_LEN, // actual packet length (payload + AES-GCM nonce & tag)
200            &self.session_key,
201            &mut self.scratch_buf,
202        )
203        .await?;
204        Ok(decrypted)
205    }
206}
207
208/// Represents server metrics that can be sent to clients or used for monitoring.
209#[repr(C)]
210pub struct ServerMetrics {
211    pub uptime_hrs: f64,
212    pub total_sessions: u64,
213    pub active_sessions: u64,
214    pub total_network_bandwidth: u64,
215    // TODO: add more fields
216}
217
218impl HydraClient<Observer> {
219    /// Connects to the server, performs handshake, sends encrypted `JoinHeader`, decrypts/reads the `StatusCode` response,
220    /// and returns a `HydraClient<Observer>` instance if successful.
221    pub async fn connect<A: ToSocketAddrs>(addr: A, observe_token: [u8; 64]) -> Result<Self> {
222        let stream = TcpStream::connect(addr).await?;
223        stream.set_nodelay(true)?;
224        let (reader, writer) = stream.into_split();
225        let mut writer = BufWriter::with_capacity(BUFFER_SIZE, writer);
226        let mut reader = BufReader::with_capacity(BUFFER_SIZE, reader);
227
228        let transport_key = perform_client_handshake(&mut reader, &mut writer).await?;
229        let mut scratch_buf = BytesMut::with_capacity(1024 * 1024 * 12);
230        write_join_header(
231            &mut writer,
232            Role::Observer,
233            observe_token,
234            &transport_key,
235            &mut scratch_buf,
236        )
237        .await?;
238
239        let status_code = read_status_code(&mut reader, &transport_key, &mut scratch_buf).await?;
240        match status_code {
241            StatusCode::Success => {}
242            StatusCode::ErrInvalidToken => {
243                bail!("ERR: Invalid observe token, cannot join as observer");
244            }
245            StatusCode::ErrInternalServerError => {
246                bail!("ERR: Internal server error, cannot join as observer");
247            }
248            _ => {
249                unreachable!()
250            }
251        }
252
253        let server_read_capacity =
254            read_server_read_write_len(&mut reader, &transport_key, &mut scratch_buf).await?;
255
256        Ok(Self {
257            buf_reader: reader,
258            buf_writer: writer,
259            scratch_buf,
260            session_key: [0u8; 32], // observers don't need session_key, they just recv
261            transport_key,
262            server_read_write_len: server_read_capacity,
263            closed: false,
264            _role: PhantomData,
265        })
266    }
267
268    /// Pulls the latest [`ServerMetrics`] as server sends it.
269    /// > This has roundtrip latency, server sends metrics only when requested by observer client (sends 1 dumb byte signal).
270    pub async fn observe(&mut self) -> Result<ServerMetrics> {
271        // this is fine, its just signal, server does not 'read' it
272        self.buf_writer.write_all(&[0u8]).await?;
273        self.buf_writer.flush().await?;
274
275        let metrics = read_server_metrics(
276            &mut self.buf_reader,
277            &self.transport_key,
278            &mut self.scratch_buf,
279        )
280        .await?;
281        Ok(metrics)
282    }
283}
284
285impl<R> HydraClient<R> {
286    /// Returns the number of raw payload bytes per packet that server considers for `read/write` operations.
287    /// > NOTE: This is the plaintext data length of a single packet in bytes, the `AES-GCM` nonce and tag are handled internally.
288    /// > [`Broadcast`](HydraClient<Producer>::broadcast) already checks against `server_read_write_len` to ensure match the API correctly.
289    /// > So you might wanna check your packet size against `server_read_write_len`, no crypto arithmetic needed on your side.
290    pub fn get_server_read_write_length(&self) -> u64 {
291        self.server_read_write_len
292    }
293
294    /// Closes the client connection gracefully by flushing and shutting down the writer (proper FIN).
295    pub async fn close(&mut self) -> Result<()> {
296        self.buf_writer.flush().await?;
297        self.buf_writer.shutdown().await?;
298        self.closed = true;
299        Ok(())
300    }
301}
302
303impl<D> Drop for HydraClient<D> {
304    fn drop(&mut self) {
305        if !self.closed {
306            eprintln!(
307                "{}",
308                "Warning: HydraClient dropped without calling close()".yellow()
309            );
310        }
311    }
312}