1use 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
17pub struct Producer;
19pub struct Consumer;
21pub struct Observer;
23
24pub 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 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 pub async fn broadcast(&mut self, data: &[u8]) -> Result<()> {
115 let plaintext_len = self.server_read_write_len as usize; 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 for chunk in data.chunks_exact(plaintext_len) {
126 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 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 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, &self.session_key,
201 &mut self.scratch_buf,
202 )
203 .await?;
204 Ok(decrypted)
205 }
206}
207
208#[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 }
217
218impl HydraClient<Observer> {
219 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], transport_key,
262 server_read_write_len: server_read_capacity,
263 closed: false,
264 _role: PhantomData,
265 })
266 }
267
268 pub async fn observe(&mut self) -> Result<ServerMetrics> {
271 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 pub fn get_server_read_write_length(&self) -> u64 {
291 self.server_read_write_len
292 }
293
294 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}