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_read_write_len,
6    read_status_code, write_encrypt_packet, write_join_header,
7};
8use anyhow::{Result, bail};
9use bytes::BytesMut;
10use colored::Colorize;
11use std::marker::PhantomData;
12use tokio::io::{AsyncWriteExt, BufReader, BufWriter};
13use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
14use tokio::net::{TcpStream, ToSocketAddrs};
15
16/// Unit structs to represent the `Role::Producer` at the type level for better safety and clarity.
17pub struct Producer;
18/// Unit structs to represent the `Role::Consumer` at the type level for better safety and clarity.
19pub struct Consumer;
20
21/// **HydraClient** - `Cross-lang-friendly` client connects to the [`HydraServer`](crate::server::HydraServer) as a producer or consumer,
22/// performs handshake, and sends/receives encrypted packet `(AES-GCM)`. It maintains an internal memory buffer (12 Mb) for zero-copy crypto and buffering.
23///
24/// ```no_run
25/// use hydra_sync::client::{HydraClient, Producer, Consumer};
26///
27/// #[tokio::main]
28/// async fn main() {
29///     let addr = "127.0.0.1:8000";
30///     let session_id = [0xFFu8; 64];
31///     let session_key = [0xAAu8; 32];
32///
33///     let mut producer = HydraClient::<Producer>::connect(addr, session_id, session_key).await.unwrap();
34///     producer.broadcast(b"I luv you >.<").await.unwrap(); // sends to all consumer
35///
36///     let mut consumer = HydraClient::<Consumer>::connect(addr, session_id, session_key).await.unwrap();
37///     consumer.recv().await.unwrap(); // recv whatever next frame on ring buf
38/// }
39/// ```
40///
41pub struct HydraClient<R> {
42    buf_reader: BufReader<OwnedReadHalf>,
43    buf_writer: BufWriter<OwnedWriteHalf>,
44    scratch_buf: BytesMut,
45    session_key: [u8; 32],
46    server_read_write_len: u64,
47    closed: bool,
48    _role: PhantomData<R>,
49}
50
51impl HydraClient<Producer> {
52    /// Connects to the server, performs handshake, sends encrypted `JoinHeader`, decrypts/reads the `StatusCode` response,
53    /// and returns a `HydraClient<Producer>` instance if successful.
54    pub async fn connect<A: ToSocketAddrs>(
55        addr: A,
56        session_id: [u8; 64],
57        session_key: [u8; 32],
58    ) -> Result<Self> {
59        let stream = TcpStream::connect(addr).await?;
60        stream.set_nodelay(true)?;
61        let (reader, writer) = stream.into_split();
62        let mut writer = BufWriter::with_capacity(BUFFER_SIZE, writer);
63        let mut reader = BufReader::with_capacity(BUFFER_SIZE, reader);
64        let transport_key = perform_client_handshake(&mut reader, &mut writer).await?;
65        let mut scratch_buf = BytesMut::with_capacity(1024 * 1024 * 12);
66        write_join_header(
67            &mut writer,
68            Role::Producer,
69            session_id,
70            &transport_key,
71            &mut scratch_buf,
72        )
73        .await?;
74
75        let status_code = read_status_code(&mut reader, &transport_key, &mut scratch_buf).await?;
76        match status_code {
77            StatusCode::Success => {}
78            StatusCode::ErrSessionAlreadyOccupied => {
79                bail!("ERR: Session already exists, cannot join as producer");
80            }
81            StatusCode::ErrInternalServerError => {
82                bail!("ERR: Internal server error, cannot join as producer");
83            }
84            _ => {
85                unreachable!()
86            }
87        }
88
89        let server_packet_write_len =
90            read_server_read_write_len(&mut reader, &transport_key, &mut scratch_buf).await?;
91
92        Ok(Self {
93            buf_reader: reader,
94            buf_writer: writer,
95            scratch_buf,
96            session_key,
97            server_read_write_len: server_packet_write_len,
98            closed: false,
99            _role: PhantomData,
100        })
101    }
102
103    /// Broadcasts the given data as `fixed encrypted packets` (AES-GCM256) to all connected consumers **(Zero-copy)
104    /// and the [HydraServer](crate::server::HydraServer) does not decrypt it.**.
105    ///
106    /// **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.**
107    ///
108    /// > This may block until the packets are written to the server ring buffer, this behavior is configured by [`ChannelOverflowStrategy`](crate::ChannelOverflowStrategy).
109    pub async fn broadcast(&mut self, data: &[u8]) -> Result<()> {
110        let plaintext_len = self.server_read_write_len as usize; // get actual data
111        if plaintext_len == 0 || !data.len().is_multiple_of(plaintext_len) {
112            bail!(
113                "Data length {} is not an exact multiple of plaintext length {}",
114                data.len(),
115                plaintext_len
116            );
117        }
118
119        // write all chunks sequentially
120        for chunk in data.chunks_exact(plaintext_len) {
121            // this writes data.len() + crypto overhead (28 bytes)
122            write_encrypt_packet(
123                &mut self.buf_writer,
124                chunk,
125                &self.session_key,
126                &mut self.scratch_buf,
127            )
128            .await?;
129        }
130        Ok(())
131    }
132}
133
134impl HydraClient<Consumer> {
135    /// Connects to the server, performs handshake, sends encrypted `JoinHeader`, decrypts/reads the `StatusCode` response,
136    /// and returns a `HydraClient<Consumer>` instance if successful.
137    pub async fn connect<A: ToSocketAddrs>(
138        addr: A,
139        session_id: [u8; 64],
140        session_key: [u8; 32],
141    ) -> Result<Self> {
142        let stream = TcpStream::connect(addr).await?;
143        stream.set_nodelay(true)?;
144        let (reader, writer) = stream.into_split();
145        let mut writer = BufWriter::with_capacity(BUFFER_SIZE, writer);
146        let mut reader = BufReader::with_capacity(BUFFER_SIZE, reader);
147
148        let transport_key = perform_client_handshake(&mut reader, &mut writer).await?;
149        let mut scratch_buf = BytesMut::with_capacity(1024 * 1024 * 12);
150        write_join_header(
151            &mut writer,
152            Role::Consumer,
153            session_id,
154            &transport_key,
155            &mut scratch_buf,
156        )
157        .await?;
158
159        let status_code = read_status_code(&mut reader, &transport_key, &mut scratch_buf).await?;
160        match status_code {
161            StatusCode::Success => {}
162            StatusCode::ErrSessionNotFound => {
163                bail!("ERR: Session not found, cannot join as consumer");
164            }
165            StatusCode::ErrInternalServerError => {
166                bail!("ERR: Internal server error, cannot join as consumer");
167            }
168            _ => {
169                unreachable!()
170            }
171        }
172
173        let server_read_capacity =
174            read_server_read_write_len(&mut reader, &transport_key, &mut scratch_buf).await?;
175
176        Ok(Self {
177            buf_reader: reader,
178            buf_writer: writer,
179            scratch_buf,
180            session_key,
181            server_read_write_len: server_read_capacity,
182            closed: false,
183            _role: PhantomData,
184        })
185    }
186
187    /// Receives the next `fixed encrypted packet` from the producer, decrypts it, and returns the plaintext data as a byte slice.
188    /// The returned slice is valid until the next call to `recv` or `broadcast`, which reuse the internal scratch buffer.
189    /// > This may block until the next packet is available on server consumer's queue.
190    pub async fn recv(&mut self) -> Result<&[u8]> {
191        let decrypted = read_decrypt_packet(
192            &mut self.buf_reader,
193            self.server_read_write_len as usize + NONCE_LEN + TAG_LEN, // actual packet length (payload + AES-GCM nonce & tag)
194            &self.session_key,
195            &mut self.scratch_buf,
196        )
197        .await?;
198        Ok(decrypted)
199    }
200}
201
202impl<R> HydraClient<R> {
203    /// Returns the number of raw payload bytes per packet that server considers for `read/write` operations.
204    /// > NOTE: This is the plaintext data length of a single packet in bytes, the `AES-GCM` nonce and tag are handled internally.
205    /// > [`Broadcast`](HydraClient<Producer>::broadcast) already checks against `server_read_write_len` to ensure match the API correctly.
206    /// > So you might wanna check your packet size against `server_read_write_len`, no crypto arithmetic needed on your side.
207    pub fn get_server_read_write_length(&self) -> u64 {
208        self.server_read_write_len
209    }
210
211    /// Closes the client connection gracefully by flushing and shutting down the writer (proper FIN).
212    pub async fn close(&mut self) -> Result<()> {
213        self.buf_writer.flush().await?;
214        self.buf_writer.shutdown().await?;
215        self.closed = true;
216        Ok(())
217    }
218}
219
220impl<D> Drop for HydraClient<D> {
221    fn drop(&mut self) {
222        if !self.closed {
223            eprintln!(
224                "{}",
225                "Warning: HydraClient dropped without calling close()".yellow()
226            );
227        }
228    }
229}