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#[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
54pub 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 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
85pub 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 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)]
117fn 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#[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#[repr(C)]
156pub struct JoinHeader {
157 pub role: u8,
158 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#[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#[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#[inline]
305pub 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]
318pub 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
329impl 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
381pub 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 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
413pub 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 if scratch_buf.len() < ciphertext_len {
426 scratch_buf.resize(ciphertext_len, 0);
427 }
428
429 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}