Skip to main content

hydra_sync/
protocol.rs

1use crate::client::ServerMetrics;
2use crate::crypto::{NONCE_LEN, TAG_LEN, decrypt_into, encrypt_into, generate_x25519_keypair};
3use anyhow::{Result, bail};
4use bytes::BytesMut;
5use rand::Rng;
6use sha3::{Digest, Sha3_256};
7use std::fmt::Display;
8use tokio::io::{AsyncReadExt, AsyncWriteExt};
9use x25519_dalek::PublicKey;
10
11/// Represents the role of a client in the Hydra protocol, either as an Observer, Producer, or Consumer.
12#[repr(u8)]
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub enum Role {
15    Producer,
16    Consumer,
17    Observer,
18}
19
20impl Role {
21    #[inline(always)]
22    pub fn to_u8(self) -> u8 {
23        match self {
24            Role::Producer => 0x00,
25            Role::Consumer => 0x01,
26            Role::Observer => 0x02,
27        }
28    }
29    #[inline(always)]
30    pub fn from_u8(val: u8) -> Result<Self> {
31        match val {
32            0x00 => Ok(Self::Producer),
33            0x01 => Ok(Self::Consumer),
34            0x02 => Ok(Self::Observer),
35            _ => bail!("Unknown role: {:#04x}", val),
36        }
37    }
38}
39
40impl Display for Role {
41    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
42        match self {
43            Role::Observer => write!(f, "Observer"),
44            Role::Producer => write!(f, "Producer"),
45            Role::Consumer => write!(f, "Consumer"),
46        }
47    }
48}
49
50const HANDSHAKE_SALT_LEN: usize = 128;
51const HANDSHAKE_INFO: &[u8] =
52    concat!("hydra-sync transport key/v", env!("CARGO_PKG_VERSION")).as_bytes();
53
54/// Perform X25519 key exchange handshake on client side and return the derived shared secret key
55pub async fn perform_client_handshake<R: AsyncReadExt + Unpin, W: AsyncWriteExt + Unpin>(
56    reader: &mut R,
57    writer: &mut W,
58) -> Result<[u8; 32]> {
59    let (secret, client_pub) = generate_x25519_keypair()?;
60
61    let mut client_nonce = [0u8; HANDSHAKE_SALT_LEN];
62    rand::rng().fill_bytes(&mut client_nonce);
63
64    // write nonce + 32 bytes key
65    writer.write_all(&client_nonce).await?;
66    writer.write_all(client_pub.as_bytes()).await?;
67    writer.flush().await?;
68
69    let mut server_nonce = [0u8; HANDSHAKE_SALT_LEN];
70    reader.read_exact(&mut server_nonce).await?;
71    let mut server_pub_bytes = [0u8; 32];
72    reader.read_exact(&mut server_pub_bytes).await?;
73    let server_pub = PublicKey::from(server_pub_bytes);
74
75    let shared_secret = secret.diffie_hellman(&server_pub);
76    derive_encryption_key(
77        shared_secret.as_bytes(),
78        client_pub.as_bytes(),
79        &client_nonce,
80        server_pub.as_bytes(),
81        &server_nonce,
82    )
83}
84
85/// Perform X25519 key exchange handshake on server side and return the derived shared secret key
86pub async fn perform_server_handshake<R: AsyncReadExt + Unpin, W: AsyncWriteExt + Unpin>(
87    reader: &mut R,
88    writer: &mut W,
89) -> Result<[u8; 32]> {
90    let (secret, server_pub) = generate_x25519_keypair()?;
91
92    // read client nonce + 32 bytes key
93    let mut client_nonce = [0u8; HANDSHAKE_SALT_LEN];
94    reader.read_exact(&mut client_nonce).await?;
95    let mut client_pub_bytes = [0u8; 32];
96    reader.read_exact(&mut client_pub_bytes).await?;
97    let client_pub = PublicKey::from(client_pub_bytes);
98
99    let mut server_nonce = [0u8; HANDSHAKE_SALT_LEN];
100    rand::rng().fill_bytes(&mut server_nonce);
101
102    writer.write_all(&server_nonce).await?;
103    writer.write_all(server_pub.as_bytes()).await?;
104    writer.flush().await?;
105
106    let shared_secret = secret.diffie_hellman(&client_pub);
107    derive_encryption_key(
108        shared_secret.as_bytes(),
109        client_pub.as_bytes(),
110        &client_nonce,
111        server_pub.as_bytes(),
112        &server_nonce,
113    )
114}
115
116#[inline(always)]
117/// Derives encryption key by hashing the concatenation of client/server public keys, nonce, handshake info, and shared secret using SHA3-256
118fn derive_encryption_key(
119    shared_secret: &[u8],
120    client_pub: &[u8; 32],
121    client_nonce: &[u8; HANDSHAKE_SALT_LEN],
122    server_pub: &[u8; 32],
123    server_nonce: &[u8; HANDSHAKE_SALT_LEN],
124) -> Result<[u8; 32]> {
125    let mut transcript = vec![0u8; 32 + HANDSHAKE_SALT_LEN + 32 + HANDSHAKE_SALT_LEN + 32];
126    transcript.extend_from_slice(client_pub);
127    transcript.extend_from_slice(client_nonce);
128    transcript.extend_from_slice(server_pub);
129    transcript.extend_from_slice(server_nonce);
130    transcript.extend_from_slice(shared_secret);
131    transcript.extend_from_slice(HANDSHAKE_INFO);
132    Ok(Sha3_256::digest(transcript).into())
133}
134
135/// Asserts that the size of a type is equal to `size`.
136#[macro_export]
137macro_rules! assert_size {
138    ($ty:ty, $expected:expr $(,)?) => {
139        const _: () = assert!(
140            ::core::mem::size_of::<$ty>() == $expected,
141            concat!(
142                "size of `",
143                stringify!($ty),
144                "` is not ",
145                stringify!($expected),
146                " bytes"
147            ),
148        );
149    };
150}
151
152// ---
153
154/// Header sent by the client to the server during the join/role determine process.
155#[repr(C)]
156pub struct JoinHeader {
157    pub role: u8,
158    // for Role::Producer/Consumer, this is session_uuid
159    // for Role::Observer, this is observe_token
160    pub uuid_or_token: [u8; 64],
161}
162
163impl JoinHeader {
164    #[inline(always)]
165    pub fn as_bytes(&self) -> &[u8] {
166        assert_size!(JoinHeader, 65);
167        unsafe { std::slice::from_raw_parts(self as *const _ as *const u8, size_of::<Self>()) }
168    }
169
170    #[inline(always)]
171    pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
172        if bytes.len() < size_of::<Self>() {
173            bail!(
174                "Invalid JoinHeader length: expected {}, got {}",
175                size_of::<Self>(),
176                bytes.len()
177            );
178        }
179        let session_uuid: [u8; 64] = bytes[1..65].try_into()?;
180        Ok(Self {
181            role: bytes[0],
182            uuid_or_token: session_uuid,
183        })
184    }
185}
186
187pub async fn read_join_header<R: AsyncReadExt + Unpin>(
188    reader: &mut R,
189    transport_key: &[u8; 32],
190    decrypt_buf: &mut BytesMut,
191) -> Result<JoinHeader> {
192    let read_len = NONCE_LEN + size_of::<JoinHeader>() + TAG_LEN;
193    let decrypted_buf = read_decrypt_packet(reader, read_len, transport_key, decrypt_buf).await?;
194    JoinHeader::from_bytes(decrypted_buf)
195}
196
197pub async fn write_join_header<W: AsyncWriteExt + Unpin>(
198    writer: &mut W,
199    role: Role,
200    uuid_or_token: [u8; 64],
201    transport_key: &[u8; 32],
202    encrypt_buf: &mut BytesMut,
203) -> Result<()> {
204    let join_header = JoinHeader {
205        role: role.to_u8(),
206        uuid_or_token,
207    };
208    write_encrypt_packet(writer, join_header.as_bytes(), transport_key, encrypt_buf).await
209}
210
211// ---
212
213// #[repr(C)]
214// pub struct PacketHeader {
215//     pub seq: u128,
216//     pub timestamp: u128,
217// }
218// pub const PACKET_HEADER_SIZE_OVERHEAD: usize = size_of::<PacketHeader>();
219// impl PacketHeader {
220//     #[inline(always)]
221//     pub fn as_bytes(&self) -> &[u8] {
222//         assert_size!(PacketHeader, 32);
223//         unsafe { std::slice::from_raw_parts(self as *const _ as *const u8, size_of::<Self>()) }
224//     }
225//
226//     #[inline(always)]
227//     pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
228//         if bytes.len() < size_of::<Self>() {
229//             bail!(
230//                 "Invalid PacketHeader length: expected {}, got {}",
231//                 size_of::<Self>(),
232//                 bytes.len()
233//             );
234//         }
235//         let seq = u128::from_le_bytes(bytes[0..16].try_into()?);
236//         let timestamp = u128::from_le_bytes(bytes[16..32].try_into()?);
237//         Ok(Self { seq, timestamp })
238//     }
239// }
240
241#[repr(u8)]
242#[derive(Debug)]
243pub enum StatusCode {
244    Success,
245    ErrInvalidRole,
246    ErrSessionNotFound,
247    ErrSessionAlreadyOccupied,
248    ErrInvalidToken,
249    ErrInternalServerError,
250}
251
252impl StatusCode {
253    #[inline(always)]
254    pub fn to_u8(&self) -> u8 {
255        match self {
256            StatusCode::Success => 0x00,
257            StatusCode::ErrInvalidRole => 0x01,
258            StatusCode::ErrSessionNotFound => 0x02,
259            StatusCode::ErrSessionAlreadyOccupied => 0x03,
260            StatusCode::ErrInvalidToken => 0x04,
261            StatusCode::ErrInternalServerError => 0x05,
262        }
263    }
264    #[inline(always)]
265    pub fn from_u8(val: u8) -> Result<Self> {
266        match val {
267            0x00 => Ok(StatusCode::Success),
268            0x01 => Ok(StatusCode::ErrInvalidRole),
269            0x02 => Ok(StatusCode::ErrSessionNotFound),
270            0x03 => Ok(StatusCode::ErrSessionAlreadyOccupied),
271            0x04 => Ok(StatusCode::ErrInvalidToken),
272            0x05 => Ok(StatusCode::ErrInternalServerError),
273            _ => bail!("Unknown Status code: {:#04x}", val),
274        }
275    }
276}
277
278// ---
279
280#[inline]
281pub async fn read_status_code<R: AsyncReadExt + Unpin>(
282    reader: &mut R,
283    session_key: &[u8; 32],
284    scratch_buf: &mut BytesMut,
285) -> Result<StatusCode> {
286    let read_len = NONCE_LEN + 1 + TAG_LEN;
287    let status_code_bytes = read_decrypt_packet(reader, read_len, session_key, scratch_buf).await?;
288    StatusCode::from_u8(status_code_bytes[0])
289}
290
291#[inline]
292pub async fn write_status_code<W: AsyncWriteExt + Unpin>(
293    writer: &mut W,
294    status_code: StatusCode,
295    session_key: &[u8; 32],
296    scratch_buf: &mut BytesMut,
297) -> Result<()> {
298    let status_code_bytes = [status_code.to_u8()];
299    write_encrypt_packet(writer, &status_code_bytes, session_key, scratch_buf).await
300}
301
302// ---
303
304#[inline]
305/// Decrypt & read the server's read/write length from the provided reader.
306pub async fn read_server_read_write_len<R: AsyncReadExt + Unpin>(
307    reader: &mut R,
308    session_key: &[u8; 32],
309    scratch_buf: &mut BytesMut,
310) -> Result<u64> {
311    let read_len = NONCE_LEN + size_of::<u64>() + TAG_LEN;
312    let capacity_bytes = read_decrypt_packet(reader, read_len, session_key, scratch_buf).await?;
313    let read_write_capacity = u64::from_le_bytes(capacity_bytes.try_into()?);
314    Ok(read_write_capacity)
315}
316
317#[inline]
318/// Encrypt & write the server's read/write length to the provided writer.
319pub async fn write_server_read_write_len<W: AsyncWriteExt + Unpin>(
320    writer: &mut W,
321    read_write_capacity: u64,
322    session_key: &[u8; 32],
323    scratch_buf: &mut BytesMut,
324) -> Result<()> {
325    let capacity_bytes: [u8; 8] = read_write_capacity.to_le_bytes();
326    write_encrypt_packet(writer, &capacity_bytes, session_key, scratch_buf).await
327}
328
329// ---
330
331impl ServerMetrics {
332    #[inline(always)]
333    pub fn as_bytes(&self) -> &[u8] {
334        assert_size!(ServerMetrics, 32);
335        unsafe { std::slice::from_raw_parts(self as *const _ as *const u8, size_of::<Self>()) }
336    }
337
338    #[inline(always)]
339    pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
340        if bytes.len() < size_of::<Self>() {
341            bail!(
342                "Invalid ServerMetrics length: expected {}, got {}",
343                size_of::<Self>(),
344                bytes.len()
345            );
346        }
347        let uptime_hrs = f64::from_le_bytes(bytes[0..8].try_into()?);
348        let total_sessions = u64::from_le_bytes(bytes[8..16].try_into()?);
349        let active_sessions = u64::from_le_bytes(bytes[16..24].try_into()?);
350        let total_network_bandwidth = u64::from_le_bytes(bytes[24..32].try_into()?);
351        Ok(Self {
352            uptime_hrs,
353            total_sessions,
354            active_sessions,
355            total_network_bandwidth,
356        })
357    }
358}
359
360#[inline]
361pub async fn read_server_metrics<R: AsyncReadExt + Unpin>(
362    reader: &mut R,
363    transport_key: &[u8; 32],
364    scratch_buf: &mut BytesMut,
365) -> Result<ServerMetrics> {
366    let read_len = NONCE_LEN + size_of::<ServerMetrics>() + TAG_LEN;
367    let metrics_bytes = read_decrypt_packet(reader, read_len, transport_key, scratch_buf).await?;
368    ServerMetrics::from_bytes(metrics_bytes)
369}
370
371#[inline]
372pub async fn write_server_metrics<W: AsyncWriteExt + Unpin>(
373    writer: &mut W,
374    metrics: &ServerMetrics,
375    session_key: &[u8; 32],
376    scratch_buf: &mut BytesMut,
377) -> Result<()> {
378    write_encrypt_packet(writer, metrics.as_bytes(), session_key, scratch_buf).await
379}
380
381// ---
382
383/// Decrypt & read an encrypted packet from the provided reader, and returns a reference to the decrypted data `(plaintext)`.
384/// The `read_len` is the total `length of the encrypted packet`, which must be > `NONCE_LEN` + `TAG_LEN`.
385pub async fn read_decrypt_packet<'a, R: AsyncReadExt + Unpin>(
386    reader: &mut R,
387    read_len: usize,
388    encrypt_key: &[u8; 32],
389    scratch_buf: &'a mut BytesMut,
390) -> Result<&'a [u8]> {
391    if read_len < NONCE_LEN + TAG_LEN {
392        bail!("Data length too short for decryption: {}", read_len);
393    }
394
395    let plaintext_len = read_len - NONCE_LEN - TAG_LEN;
396    let required_len = read_len + plaintext_len;
397
398    // should accommodate both encrypted & decrypted chunk
399    if scratch_buf.len() < required_len {
400        scratch_buf.resize(required_len, 0);
401    }
402
403    let (encrypted_chunk, decrypted_chunk) = scratch_buf.split_at_mut(read_len);
404    let encrypted_buf = &mut encrypted_chunk[..read_len];
405    reader.read_exact(encrypted_buf).await?;
406
407    let decrypted_buf = &mut decrypted_chunk[..plaintext_len];
408    decrypt_into(encrypted_buf, decrypted_buf, encrypt_key)?;
409
410    Ok(decrypted_buf)
411}
412
413/// Encrypts the given data and writes it to the provided writer.
414/// The encrypted packet's len is `NONCE_LEN` + `data.len()` + `TAG_LEN`.
415pub async fn write_encrypt_packet<W: AsyncWriteExt + Unpin>(
416    writer: &mut W,
417    data: &[u8],
418    encrypt_key: &[u8; 32],
419    scratch_buf: &mut BytesMut,
420) -> Result<()> {
421    let plaintext_len = data.len();
422    let ciphertext_len = NONCE_LEN + plaintext_len + TAG_LEN;
423
424    // alloc if less
425    if scratch_buf.len() < ciphertext_len {
426        scratch_buf.resize(ciphertext_len, 0);
427    }
428
429    // encrypt that son of a bitch
430    encrypt_into(data, &mut scratch_buf[..ciphertext_len], encrypt_key)?;
431    writer.write_all(&scratch_buf[..ciphertext_len]).await?;
432    writer.flush().await?;
433    Ok(())
434}