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_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
16pub struct Producer;
18pub struct Consumer;
20
21pub 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 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 pub async fn broadcast(&mut self, data: &[u8]) -> Result<()> {
110 let plaintext_len = self.server_read_write_len as usize; 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 for chunk in data.chunks_exact(plaintext_len) {
121 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 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 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, &self.session_key,
195 &mut self.scratch_buf,
196 )
197 .await?;
198 Ok(decrypted)
199 }
200}
201
202impl<R> HydraClient<R> {
203 pub fn get_server_read_write_length(&self) -> u64 {
208 self.server_read_write_len
209 }
210
211 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}