1use std::collections::HashMap;
16use std::fmt::Debug;
17use std::net::{IpAddr, SocketAddr};
18use std::path::PathBuf;
19use std::result;
20use std::sync::Arc;
21use std::time::Duration;
22
23use argon2::{Argon2, PasswordVerifier, password_hash::PasswordHashString};
24use bloop_protocol::Capabilities;
25use bloop_protocol::codec::{Encode, EncodeError};
26use bloop_protocol::frame::{FrameError, RawMessage, read_frame, write_frame};
27use bloop_protocol::message::{
28 Authentication, AuthenticationAccepted, Bloop, ClientHandshake, ClientMessage, ErrorResponse,
29 Pong, PreloadCheck, RetrieveAudio, ServerHandshake, ServerMessage,
30};
31use bloop_protocol::set::{MessageSet, MessageSetError, NoExtension, Payload, encode_message};
32use rustls::ServerConfig;
33use rustls::pki_types::{
34 CertificateDer, PrivateKeyDer,
35 pem::{self, PemObject},
36};
37use thiserror::Error;
38use tokio::io::{self, AsyncRead, AsyncWrite, BufReader, BufWriter};
39use tokio::net::{TcpListener, TcpStream};
40use tokio::sync::{RwLock, broadcast, mpsc, oneshot};
41use tokio::time::timeout;
42#[cfg(feature = "tokio-graceful-shutdown")]
43use tokio_graceful_shutdown::{FutureExt, IntoSubsystem, SubsystemHandle};
44use tokio_rustls::TlsAcceptor;
45use tracing::{info, instrument, warn};
46
47use crate::engine::EngineRequest;
48use crate::event::Event;
49
50pub type ClientRegistry = HashMap<String, PasswordHashString>;
54
55pub const DEFAULT_MAX_PAYLOAD_LEN: u32 = 64 * 1024;
61
62#[derive(Error, Debug)]
63pub enum Error {
64 #[error(transparent)]
65 Io(#[from] io::Error),
66
67 #[error(transparent)]
68 Frame(#[from] FrameError),
69
70 #[error(transparent)]
71 Encode(#[from] EncodeError),
72
73 #[error(transparent)]
74 Oneshot(#[from] oneshot::error::RecvError),
75
76 #[error(
77 "client sent unexpected message with opcode 0x{:02x} and {} payload bytes",
78 .0.message_type,
79 .0.payload.len()
80 )]
81 UnexpectedMessage(RawMessage),
82
83 #[error("client sent malformed message")]
84 MalformedMessage(#[source] MessageSetError),
85
86 #[error("client requested an unsupported version range: {0} - {1}")]
87 UnsupportedVersion(u8, u8),
88
89 #[error("client provided invalid credentials")]
90 InvalidCredentials,
91}
92
93pub type Result<T> = result::Result<T, Error>;
94
95#[derive(Debug)]
97pub enum CustomOutcome<Res> {
98 Response(Res),
100
101 Error(ErrorResponse),
103}
104
105#[derive(Debug)]
111pub struct CustomRequestMessage<Req, Res> {
112 pub client_id: String,
113 pub request: Req,
114 pub response: oneshot::Sender<CustomOutcome<Res>>,
115}
116
117pub struct NetworkListener<Req = NoExtension, Res = NoExtension> {
150 clients: Arc<RwLock<ClientRegistry>>,
151 addr: SocketAddr,
152 tls_acceptor: TlsAcceptor,
153 engine_tx: mpsc::Sender<(EngineRequest, oneshot::Sender<ServerMessage>)>,
154 event_tx: broadcast::Sender<Event>,
155 custom_req_tx: Option<mpsc::Sender<CustomRequestMessage<Req, Res>>>,
156 max_payload_len: u32,
157}
158
159impl<Req, Res> Debug for NetworkListener<Req, Res> {
160 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
161 f.debug_struct("NetworkListener")
162 .field("clients", &self.clients)
163 .field("addr", &self.addr)
164 .field("engine_tx", &self.engine_tx)
165 .field("event_tx", &self.event_tx)
166 .field("custom_req_tx", &self.custom_req_tx)
167 .field("max_payload_len", &self.max_payload_len)
168 .finish()
169 }
170}
171
172impl<Req, Res> NetworkListener<Req, Res>
173where
174 Req: MessageSet + Send + 'static,
175 Res: MessageSet + Send + 'static,
176{
177 pub async fn listen(&self) -> Result<()> {
184 let listener = TcpListener::bind(self.addr).await?;
185 let mut con_counter: usize = 0;
186
187 loop {
188 let (stream, peer_addr) = listener.accept().await?;
189 let conn_id = con_counter;
190 con_counter += 1;
191 let event_tx = self.event_tx.clone();
192
193 self.handle_stream(stream, peer_addr, self.clients.clone(), conn_id, event_tx);
194 }
195 }
196
197 #[instrument(skip(self, stream, peer_addr, clients, event_tx))]
198 fn handle_stream(
199 &self,
200 stream: TcpStream,
201 peer_addr: SocketAddr,
202 clients: Arc<RwLock<ClientRegistry>>,
203 conn_id: usize,
204 event_tx: broadcast::Sender<Event>,
205 ) {
206 let acceptor = self.tls_acceptor.clone();
207 let engine_tx = self.engine_tx.clone();
208 let custom_req_tx = self.custom_req_tx.clone();
209 let max_payload_len = self.max_payload_len;
210
211 tokio::spawn(async move {
212 info!("new connection from {}", peer_addr);
213
214 let stream = match acceptor.accept(stream).await {
215 Ok(stream) => stream,
216 Err(error) => {
217 warn!("failed to accept stream: {}", error);
218 return;
219 }
220 };
221
222 let (reader, writer) = io::split(stream);
223 let mut reader = BufReader::new(reader);
224 let mut writer = BufWriter::new(writer);
225
226 let (client_id, local_ip, _version) = match timeout(
227 Duration::from_secs(2),
228 authenticate::<_, _, Req>(&mut reader, &mut writer, clients, max_payload_len),
229 )
230 .await
231 {
232 Ok(Ok(result)) => result,
233 Ok(Err(error)) => {
234 if let Some(response) = error_response(&error) {
235 warn!("client error: {}", error);
236 let _ = write_payload(&mut writer, &response).await;
237 } else {
238 warn!("client error: connection died: {:?}", error);
239 }
240
241 return;
242 }
243 Err(_) => {
244 warn!("client error: authentication timed out");
245 return;
246 }
247 };
248
249 let _ = event_tx.send(Event::ClientConnect {
250 client_id: client_id.clone(),
251 conn_id,
252 local_ip,
253 });
254
255 match handle_connection(
256 &mut reader,
257 &mut writer,
258 &client_id,
259 engine_tx,
260 custom_req_tx,
261 max_payload_len,
262 )
263 .await
264 {
265 Ok(()) => {
266 let _ = event_tx.send(Event::ClientDisconnect { client_id, conn_id });
267 }
268 Err(error) => {
269 if let Some(response) = error_response(&error) {
270 warn!("client error: {}", error);
271 let _ = write_payload(&mut writer, &response).await;
272 } else if matches!(error, Error::Oneshot(_) | Error::Encode(_)) {
273 warn!("internal error while serving client: {:?}", error);
274 } else {
275 warn!("client error: connection died: {:?}", error);
276 }
277
278 let _ = event_tx.send(Event::ClientConnectionLoss { client_id, conn_id });
279 }
280 }
281 });
282 }
283}
284
285#[cfg(feature = "tokio-graceful-shutdown")]
286impl<Req, Res> IntoSubsystem<Error> for NetworkListener<Req, Res>
287where
288 Req: MessageSet + Send + 'static,
289 Res: MessageSet + Send + 'static,
290{
291 async fn run(self, subsys: &mut SubsystemHandle) -> Result<()> {
292 if let Ok(result) = self.listen().cancel_on_shutdown(subsys).await {
293 result?
294 }
295
296 Ok(())
297 }
298}
299
300fn error_response(error: &Error) -> Option<ErrorResponse> {
305 match error {
306 Error::UnexpectedMessage(_) => Some(ErrorResponse::UnexpectedMessage),
307 Error::MalformedMessage(_) => Some(ErrorResponse::MalformedMessage),
308 Error::Frame(FrameError::PayloadTooLarge { .. }) => Some(ErrorResponse::MalformedMessage),
309 Error::UnsupportedVersion(_, _) => Some(ErrorResponse::UnsupportedVersionRange),
310 Error::InvalidCredentials => Some(ErrorResponse::InvalidCredentials),
311 _ => None,
312 }
313}
314
315async fn read_message<S, Req>(stream: &mut S, max_payload_len: u32) -> Result<ClientMessage<Req>>
316where
317 S: AsyncRead + Unpin,
318 Req: MessageSet,
319{
320 let raw = read_frame(stream, max_payload_len).await?;
321
322 ClientMessage::decode(&raw).map_err(|error| match error {
323 MessageSetError::UnknownOpcode(_) => Error::UnexpectedMessage(raw),
324 error => Error::MalformedMessage(error),
325 })
326}
327
328async fn write_message<S, M>(stream: &mut S, message: M) -> Result<()>
329where
330 S: AsyncWrite + Unpin,
331 M: MessageSet,
332{
333 write_frame(stream, &message.encode()?).await?;
334 Ok(())
335}
336
337async fn write_payload<S, M>(stream: &mut S, message: &M) -> Result<()>
338where
339 S: AsyncWrite + Unpin,
340 M: Payload + Encode,
341{
342 write_frame(stream, &encode_message(message)?).await?;
343 Ok(())
344}
345
346#[derive(Debug, Error)]
347pub enum BuilderError {
348 #[error("missing field: {0}")]
349 MissingField(&'static str),
350
351 #[error(transparent)]
352 AddrParse(#[from] std::net::AddrParseError),
353
354 #[error("failed to read PEM file at {path}: {source}")]
355 Pem {
356 path: PathBuf,
357 #[source]
358 source: pem::Error,
359 },
360
361 #[error(transparent)]
362 Rustls(#[from] rustls::Error),
363}
364
365pub type BuilderResult<T> = result::Result<T, BuilderError>;
366
367#[derive(Debug)]
398pub struct NetworkListenerBuilder<Req = NoExtension, Res = NoExtension> {
399 address: Option<String>,
400 cert_path: Option<PathBuf>,
401 key_path: Option<PathBuf>,
402 clients: Option<Arc<RwLock<ClientRegistry>>>,
403 engine_tx: Option<mpsc::Sender<(EngineRequest, oneshot::Sender<ServerMessage>)>>,
404 event_tx: Option<broadcast::Sender<Event>>,
405 custom_req_tx: Option<mpsc::Sender<CustomRequestMessage<Req, Res>>>,
406 max_payload_len: u32,
407}
408
409impl NetworkListenerBuilder {
410 pub fn new() -> Self {
412 Self {
413 address: None,
414 cert_path: None,
415 key_path: None,
416 clients: None,
417 engine_tx: None,
418 event_tx: None,
419 custom_req_tx: None,
420 max_payload_len: DEFAULT_MAX_PAYLOAD_LEN,
421 }
422 }
423}
424
425impl Default for NetworkListenerBuilder {
426 fn default() -> Self {
427 Self::new()
428 }
429}
430
431impl<Req, Res> NetworkListenerBuilder<Req, Res> {
432 pub fn custom_req_tx<Req2, Res2>(
435 self,
436 custom_req_tx: mpsc::Sender<CustomRequestMessage<Req2, Res2>>,
437 ) -> NetworkListenerBuilder<Req2, Res2> {
438 NetworkListenerBuilder {
439 clients: self.clients,
440 address: self.address,
441 cert_path: self.cert_path,
442 key_path: self.key_path,
443 engine_tx: self.engine_tx,
444 event_tx: self.event_tx,
445 custom_req_tx: Some(custom_req_tx),
446 max_payload_len: self.max_payload_len,
447 }
448 }
449
450 pub fn address(mut self, address: impl Into<String>) -> Self {
451 self.address = Some(address.into());
452 self
453 }
454
455 pub fn cert_path(mut self, path: impl Into<PathBuf>) -> Self {
456 self.cert_path = Some(path.into());
457 self
458 }
459
460 pub fn key_path(mut self, path: impl Into<PathBuf>) -> Self {
461 self.key_path = Some(path.into());
462 self
463 }
464
465 pub fn clients(mut self, clients: Arc<RwLock<ClientRegistry>>) -> Self {
466 self.clients = Some(clients);
467 self
468 }
469
470 pub fn engine_tx(
471 mut self,
472 tx: mpsc::Sender<(EngineRequest, oneshot::Sender<ServerMessage>)>,
473 ) -> Self {
474 self.engine_tx = Some(tx);
475 self
476 }
477
478 pub fn event_tx(mut self, tx: broadcast::Sender<Event>) -> Self {
479 self.event_tx = Some(tx);
480 self
481 }
482
483 pub fn max_payload_len(mut self, max_payload_len: u32) -> Self {
485 self.max_payload_len = max_payload_len;
486 self
487 }
488
489 pub fn build(self) -> BuilderResult<NetworkListener<Req, Res>> {
493 let addr: SocketAddr = self
494 .address
495 .ok_or_else(|| BuilderError::MissingField("address"))?
496 .parse()?;
497
498 let cert_path = self
499 .cert_path
500 .ok_or_else(|| BuilderError::MissingField("cert_path"))?;
501 let key_path = self
502 .key_path
503 .ok_or_else(|| BuilderError::MissingField("key_path"))?;
504
505 let certs = CertificateDer::pem_file_iter(&cert_path)
506 .map_err(|err| BuilderError::Pem {
507 path: cert_path.clone(),
508 source: err,
509 })?
510 .collect::<result::Result<Vec<_>, _>>()
511 .map_err(|err| BuilderError::Pem {
512 path: cert_path,
513 source: err,
514 })?;
515 let key = PrivateKeyDer::from_pem_file(&key_path).map_err(|err| BuilderError::Pem {
516 path: key_path,
517 source: err,
518 })?;
519
520 let config = ServerConfig::builder()
521 .with_no_client_auth()
522 .with_single_cert(certs, key)?;
523 let tls_acceptor = TlsAcceptor::from(Arc::new(config));
524
525 Ok(NetworkListener {
526 clients: self
527 .clients
528 .ok_or_else(|| BuilderError::MissingField("clients"))?,
529 addr,
530 tls_acceptor,
531 engine_tx: self
532 .engine_tx
533 .ok_or_else(|| BuilderError::MissingField("engine_tx"))?,
534 event_tx: self
535 .event_tx
536 .ok_or_else(|| BuilderError::MissingField("event_tx"))?,
537 custom_req_tx: self.custom_req_tx,
538 max_payload_len: self.max_payload_len,
539 })
540 }
541}
542
543async fn handle_connection<R, W, Req, Res>(
548 reader: &mut R,
549 writer: &mut W,
550 client_id: &str,
551 engine_tx: mpsc::Sender<(EngineRequest, oneshot::Sender<ServerMessage>)>,
552 custom_req_tx: Option<mpsc::Sender<CustomRequestMessage<Req, Res>>>,
553 max_payload_len: u32,
554) -> Result<()>
555where
556 R: AsyncRead + Unpin,
557 W: AsyncWrite + Unpin,
558 Req: MessageSet,
559 Res: MessageSet,
560{
561 loop {
562 let message = match timeout(
563 Duration::from_secs(30),
564 read_message::<_, Req>(reader, max_payload_len),
565 )
566 .await
567 {
568 Ok(Ok(message)) => message,
569 Ok(Err(error)) => return Err(error),
570 Err(_) => return Ok(()),
571 };
572
573 let engine_request = match message {
574 ClientMessage::Bloop(Bloop { nfc_uid }) => EngineRequest::Bloop {
575 nfc_uid,
576 client_id: client_id.to_string(),
577 },
578 ClientMessage::RetrieveAudio(RetrieveAudio { achievement_id }) => {
579 EngineRequest::RetrieveAudio { id: achievement_id }
580 }
581 ClientMessage::PreloadCheck(PreloadCheck {
582 audio_manifest_hash,
583 }) => EngineRequest::PreloadCheck {
584 manifest_hash: audio_manifest_hash,
585 },
586 ClientMessage::Ping(_) => {
587 write_payload(writer, &Pong).await?;
588 continue;
589 }
590 ClientMessage::Quit(_) => break,
591 ClientMessage::Custom(request) => {
592 let Some(sender) = custom_req_tx.as_ref() else {
593 return Err(Error::UnexpectedMessage(request.encode()?));
594 };
595
596 let (resp_tx, resp_rx) = oneshot::channel();
597
598 let _ = sender
599 .send(CustomRequestMessage {
600 client_id: client_id.to_string(),
601 request,
602 response: resp_tx,
603 })
604 .await;
605
606 match resp_rx.await? {
607 CustomOutcome::Response(response) => {
608 write_message(writer, response).await?;
609 }
610 CustomOutcome::Error(error) => {
611 write_payload(writer, &error).await?;
612 }
613 }
614
615 continue;
616 }
617 message => return Err(Error::UnexpectedMessage(message.encode()?)),
618 };
619
620 let (resp_tx, resp_rx) = oneshot::channel();
621 let _ = engine_tx.send((engine_request, resp_tx)).await;
622 let response = resp_rx.await?;
623
624 write_message(writer, response).await?;
625 }
626
627 Ok(())
628}
629
630async fn authenticate<R, W, Req>(
635 reader: &mut R,
636 writer: &mut W,
637 clients: Arc<RwLock<ClientRegistry>>,
638 max_payload_len: u32,
639) -> Result<(String, IpAddr, u8)>
640where
641 R: AsyncRead + Unpin,
642 W: AsyncWrite + Unpin,
643 Req: MessageSet,
644{
645 let (min_version, max_version) = match read_message::<_, Req>(reader, max_payload_len).await? {
646 ClientMessage::Handshake(ClientHandshake {
647 min_version,
648 max_version,
649 }) => (min_version, max_version),
650 message => return Err(Error::UnexpectedMessage(message.encode()?)),
651 };
652
653 if min_version > 3 || max_version < 3 {
654 return Err(Error::UnsupportedVersion(min_version, max_version));
655 }
656
657 write_payload(
658 writer,
659 &ServerHandshake {
660 accepted_version: 3,
661 capabilities: Capabilities::PreloadCheck,
662 },
663 )
664 .await?;
665
666 let (client_id, client_secret, ip_address) =
667 match read_message::<_, Req>(reader, max_payload_len).await? {
668 ClientMessage::Authentication(Authentication {
669 client_id,
670 client_secret,
671 ip_address,
672 }) => (client_id, client_secret, ip_address),
673 message => return Err(Error::UnexpectedMessage(message.encode()?)),
674 };
675
676 let clients = clients.read().await;
677 let Some(secret_hash) = clients.get(&client_id) else {
678 return Err(Error::InvalidCredentials);
679 };
680
681 if Argon2::default()
682 .verify_password(client_secret.as_bytes(), &secret_hash.password_hash())
683 .is_err()
684 {
685 return Err(Error::InvalidCredentials);
686 }
687
688 write_payload(writer, &AuthenticationAccepted).await?;
689
690 Ok((client_id.to_string(), ip_address, 3))
691}
692
693#[cfg(test)]
694mod tests {
695 use super::*;
696 use bloop_protocol::DataHash;
697 use bloop_protocol::message::{AchievementRecord, AudioData, BloopAccepted, PreloadMatch};
698 use std::fs;
699 use tempfile::tempdir;
700 use uuid::Uuid;
701
702 #[tokio::test]
703 async fn builder_fails_with_missing_fields() {
704 let builder = NetworkListenerBuilder::new();
705 let result = builder.build();
706 assert!(matches!(result, Err(BuilderError::MissingField(_))));
707 }
708
709 #[tokio::test]
710 async fn builder_fails_with_invalid_address() {
711 let builder = NetworkListenerBuilder::new()
712 .address("invalid-addr")
713 .cert_path("cert.pem")
714 .key_path("key.pem")
715 .clients(Arc::new(RwLock::new(Default::default())))
716 .engine_tx(dummy_engine_tx())
717 .event_tx(dummy_event_tx());
718
719 let result = builder.build();
720 assert!(matches!(result, Err(BuilderError::AddrParse(_))));
721 }
722
723 #[tokio::test]
724 async fn builder_fails_on_invalid_pem_files() {
725 let dir = tempdir().unwrap();
726 let cert_path = dir.path().join("cert.pem");
727 let key_path = dir.path().join("key.pem");
728 fs::write(&cert_path, b"invalid-cert").unwrap();
729 fs::write(&key_path, b"invalid-key").unwrap();
730
731 let builder = NetworkListenerBuilder::new()
732 .address("127.0.0.1:12345")
733 .cert_path(&cert_path)
734 .key_path(&key_path)
735 .clients(Arc::new(RwLock::new(Default::default())))
736 .engine_tx(dummy_engine_tx())
737 .event_tx(dummy_event_tx());
738
739 let result = builder.build();
740 assert!(matches!(result, Err(BuilderError::Pem { .. })));
741 }
742
743 #[tokio::test]
744 async fn builder_succeeds_with_valid_dummy_pem() {
745 let dir = tempdir().unwrap();
746 let cert_path = dir.path().join("cert.pem");
747 let key_path = dir.path().join("key.pem");
748
749 let cert_data = include_bytes!("../examples/cert.pem");
750 let key_data = include_bytes!("../examples/key.pem");
751
752 fs::write(&cert_path, cert_data).unwrap();
753 fs::write(&key_path, key_data).unwrap();
754
755 let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
756 let builder = NetworkListenerBuilder::new()
757 .address("127.0.0.1:12345")
758 .cert_path(&cert_path)
759 .key_path(&key_path)
760 .clients(Arc::new(RwLock::new(Default::default())))
761 .engine_tx(dummy_engine_tx())
762 .event_tx(dummy_event_tx());
763
764 let result = builder.build();
765 assert!(result.is_ok());
766 }
767
768 #[tokio::test]
769 async fn authentication_fails_with_wrong_client_id() {
770 let clients = Arc::new(RwLock::new(Default::default()));
771
772 let client_handshake = build_handshake(3, 3);
773 let authentication = build_authentication("unknown-client", "password", "127.0.0.1");
774
775 let mut reader = tokio_test::io::Builder::new()
776 .read(&client_handshake)
777 .read(&authentication)
778 .build();
779 let mut writer = tokio_test::io::Builder::new()
780 .write(&frame_bytes(&ServerHandshake {
781 accepted_version: 3,
782 capabilities: Capabilities::PreloadCheck,
783 }))
784 .build();
785
786 let result = authenticate::<_, _, NoExtension>(
787 &mut reader,
788 &mut writer,
789 clients,
790 DEFAULT_MAX_PAYLOAD_LEN,
791 )
792 .await;
793
794 assert!(matches!(result, Err(Error::InvalidCredentials)));
795 }
796
797 #[tokio::test]
798 async fn authentication_succeeds_with_correct_credentials() {
799 let clients = test_clients().await;
800
801 let client_handshake = build_handshake(3, 3);
802 let authentication = build_authentication("client", "secret", "127.0.0.1");
803
804 let mut reader = tokio_test::io::Builder::new()
805 .read(&client_handshake)
806 .read(&authentication)
807 .build();
808 let mut writer = tokio_test::io::Builder::new()
809 .write(&frame_bytes(&ServerHandshake {
810 accepted_version: 3,
811 capabilities: Capabilities::PreloadCheck,
812 }))
813 .write(&frame_bytes(&AuthenticationAccepted))
814 .build();
815
816 let result = authenticate::<_, _, NoExtension>(
817 &mut reader,
818 &mut writer,
819 clients,
820 DEFAULT_MAX_PAYLOAD_LEN,
821 )
822 .await;
823
824 assert!(result.is_ok());
825 }
826
827 #[tokio::test]
828 async fn authentication_fails_with_wrong_password() {
829 let clients = test_clients().await;
830
831 let client_handshake = build_handshake(3, 3);
832 let authentication = build_authentication("client1", "wrong-secret", "127.0.0.1");
833
834 let mut reader = tokio_test::io::Builder::new()
835 .read(&client_handshake)
836 .read(&authentication)
837 .build();
838 let mut writer = tokio_test::io::Builder::new()
839 .write(&frame_bytes(&ServerHandshake {
840 accepted_version: 3,
841 capabilities: Capabilities::PreloadCheck,
842 }))
843 .build();
844
845 let result = authenticate::<_, _, NoExtension>(
846 &mut reader,
847 &mut writer,
848 clients,
849 DEFAULT_MAX_PAYLOAD_LEN,
850 )
851 .await;
852
853 assert!(matches!(result, Err(Error::InvalidCredentials)));
854 }
855
856 #[derive(
857 Clone,
858 Debug,
859 PartialEq,
860 bloop_protocol::Encode,
861 bloop_protocol::Decode,
862 bloop_protocol::Payload,
863 )]
864 #[bloop(opcode = 0x80)]
865 struct EchoRequest {
866 text: String,
867 }
868
869 #[derive(
870 Clone,
871 Debug,
872 PartialEq,
873 bloop_protocol::Encode,
874 bloop_protocol::Decode,
875 bloop_protocol::Payload,
876 )]
877 #[bloop(opcode = 0x81)]
878 struct EchoResponse {
879 text: String,
880 }
881
882 #[derive(Clone, Debug, PartialEq, bloop_protocol::MessageSet)]
883 enum TestRequest {
884 Echo(EchoRequest),
885 }
886
887 #[derive(Clone, Debug, PartialEq, bloop_protocol::MessageSet)]
888 enum TestResponse {
889 Echo(EchoResponse),
890 }
891
892 #[tokio::test]
898 async fn full_session_speaks_v3_bytes() {
899 use tokio::io::AsyncWriteExt;
900
901 let achievement_id = Uuid::from_bytes([7; 16]);
902 let audio_hash = DataHash::try_from(vec![9u8; 16]).unwrap();
903
904 let (engine_tx, mut engine_rx) =
906 mpsc::channel::<(EngineRequest, oneshot::Sender<ServerMessage>)>(8);
907 let engine_audio_hash = audio_hash.clone();
908
909 tokio::spawn(async move {
910 while let Some((request, response)) = engine_rx.recv().await {
911 let message: ServerMessage = match request {
912 EngineRequest::Bloop { .. } => BloopAccepted {
913 achievements: vec![AchievementRecord {
914 id: achievement_id,
915 audio_hash: Some(engine_audio_hash.clone()),
916 }],
917 }
918 .into(),
919 EngineRequest::RetrieveAudio { .. } => AudioData {
920 data: vec![1, 2, 3],
921 }
922 .into(),
923 EngineRequest::PreloadCheck { .. } => PreloadMatch.into(),
924 };
925
926 let _ = response.send(message);
927 }
928 });
929
930 let (custom_tx, mut custom_rx) =
933 mpsc::channel::<CustomRequestMessage<TestRequest, TestResponse>>(8);
934
935 tokio::spawn(async move {
936 while let Some(request) = custom_rx.recv().await {
937 assert_eq!(request.client_id, "client");
938 let TestRequest::Echo(echo) = request.request;
939
940 let outcome = if echo.text == "fail" {
941 CustomOutcome::Error(ErrorResponse::Custom(0x90))
942 } else {
943 CustomOutcome::Response(TestResponse::Echo(EchoResponse {
944 text: echo.text.chars().rev().collect(),
945 }))
946 };
947
948 let _ = request.response.send(outcome);
949 }
950 });
951
952 let (mut client, server) = tokio::io::duplex(64 * 1024);
953
954 let server_task = tokio::spawn(async move {
955 let (reader, writer) = io::split(server);
956 let mut reader = BufReader::new(reader);
957 let mut writer = BufWriter::new(writer);
958
959 let (client_id, _, _) = authenticate::<_, _, TestRequest>(
960 &mut reader,
961 &mut writer,
962 test_clients().await,
963 DEFAULT_MAX_PAYLOAD_LEN,
964 )
965 .await
966 .unwrap();
967
968 handle_connection(
969 &mut reader,
970 &mut writer,
971 &client_id,
972 engine_tx,
973 Some(custom_tx),
974 DEFAULT_MAX_PAYLOAD_LEN,
975 )
976 .await
977 .unwrap();
978 });
979
980 client.write_all(&[0x01, 2, 0, 0, 0, 3, 3]).await.unwrap();
982 let mut expected = vec![0x02, 9, 0, 0, 0, 3];
983 expected.extend_from_slice(&1u64.to_le_bytes());
984 assert_eq!(read_bytes(&mut client, expected.len()).await, expected);
985
986 client
988 .write_all(&build_authentication("client", "secret", "127.0.0.1"))
989 .await
990 .unwrap();
991 assert_eq!(read_bytes(&mut client, 5).await, [0x04, 0, 0, 0, 0]);
992
993 client.write_all(&[0x05, 0, 0, 0, 0]).await.unwrap();
995 assert_eq!(read_bytes(&mut client, 5).await, [0x06, 0, 0, 0, 0]);
996
997 client
999 .write_all(&[0x08, 5, 0, 0, 0, 4, 1, 2, 3, 4])
1000 .await
1001 .unwrap();
1002 let mut expected = vec![0x09, 34, 0, 0, 0, 1];
1003 expected.extend_from_slice(achievement_id.as_bytes());
1004 expected.push(16);
1005 expected.extend_from_slice(audio_hash.as_bytes());
1006 assert_eq!(read_bytes(&mut client, expected.len()).await, expected);
1007
1008 let mut request = vec![0x0a, 16, 0, 0, 0];
1010 request.extend_from_slice(achievement_id.as_bytes());
1011 client.write_all(&request).await.unwrap();
1012 assert_eq!(
1013 read_bytes(&mut client, 12).await,
1014 [0x0b, 7, 0, 0, 0, 3, 0, 0, 0, 1, 2, 3]
1015 );
1016
1017 client.write_all(&[0x0c, 1, 0, 0, 0, 0]).await.unwrap();
1019 assert_eq!(read_bytes(&mut client, 5).await, [0x0d, 0, 0, 0, 0]);
1020
1021 client
1023 .write_all(&[0x80, 3, 0, 0, 0, 2, b'h', b'i'])
1024 .await
1025 .unwrap();
1026 assert_eq!(
1027 read_bytes(&mut client, 8).await,
1028 [0x81, 3, 0, 0, 0, 2, b'i', b'h']
1029 );
1030
1031 client
1033 .write_all(&[0x80, 5, 0, 0, 0, 4, b'f', b'a', b'i', b'l'])
1034 .await
1035 .unwrap();
1036 assert_eq!(read_bytes(&mut client, 6).await, [0x00, 1, 0, 0, 0, 0x90]);
1037
1038 client.write_all(&[0x07, 0, 0, 0, 0]).await.unwrap();
1040 server_task.await.unwrap();
1041 }
1042
1043 #[tokio::test]
1044 async fn errors_map_to_the_owed_protocol_responses() {
1045 assert!(matches!(
1046 error_response(&Error::UnexpectedMessage(RawMessage::new(0xff, vec![]))),
1047 Some(ErrorResponse::UnexpectedMessage)
1048 ));
1049 assert!(matches!(
1050 error_response(&Error::MalformedMessage(MessageSetError::UnknownOpcode(
1051 0xff
1052 ))),
1053 Some(ErrorResponse::MalformedMessage)
1054 ));
1055 assert!(matches!(
1056 error_response(&Error::Frame(FrameError::PayloadTooLarge {
1057 length: 100,
1058 max: 10,
1059 })),
1060 Some(ErrorResponse::MalformedMessage)
1061 ));
1062 assert!(matches!(
1063 error_response(&Error::UnsupportedVersion(1, 2)),
1064 Some(ErrorResponse::UnsupportedVersionRange)
1065 ));
1066 assert!(matches!(
1067 error_response(&Error::InvalidCredentials),
1068 Some(ErrorResponse::InvalidCredentials)
1069 ));
1070
1071 let (tx, rx) = oneshot::channel::<()>();
1072 drop(tx);
1073 let recv_error = rx.await.unwrap_err();
1074
1075 assert!(error_response(&Error::Io(io::Error::other("gone"))).is_none());
1076 assert!(error_response(&Error::Oneshot(recv_error)).is_none());
1077 }
1078
1079 #[tokio::test]
1080 async fn error_responses_encode_to_v3_error_frames() {
1081 use tokio::io::AsyncWriteExt;
1082
1083 let (mut client, mut server) = tokio::io::duplex(1024);
1084
1085 write_payload(&mut client, &ErrorResponse::UnexpectedMessage)
1086 .await
1087 .unwrap();
1088 client.flush().await.unwrap();
1089
1090 assert_eq!(read_bytes(&mut server, 6).await, [0x00, 1, 0, 0, 0, 0x00]);
1091 }
1092
1093 #[tokio::test]
1094 async fn unknown_opcode_is_answered_with_unexpected_message() {
1095 let (engine_tx, _engine_rx) = mpsc::channel(1);
1096 let (mut client, server) = tokio::io::duplex(1024);
1097
1098 let server_task = tokio::spawn(async move {
1099 let (reader, writer) = io::split(server);
1100 let mut reader = BufReader::new(reader);
1101 let mut writer = BufWriter::new(writer);
1102
1103 let result = handle_connection::<_, _, NoExtension, NoExtension>(
1104 &mut reader,
1105 &mut writer,
1106 "client",
1107 engine_tx,
1108 None,
1109 DEFAULT_MAX_PAYLOAD_LEN,
1110 )
1111 .await;
1112
1113 assert!(matches!(result, Err(Error::UnexpectedMessage(_))));
1114 });
1115
1116 use tokio::io::AsyncWriteExt;
1117 client.write_all(&[0xff, 0, 0, 0, 0]).await.unwrap();
1118 server_task.await.unwrap();
1119 }
1120
1121 async fn read_bytes<S: AsyncRead + Unpin>(stream: &mut S, count: usize) -> Vec<u8> {
1122 use tokio::io::AsyncReadExt;
1123
1124 let mut bytes = vec![0; count];
1125 stream.read_exact(&mut bytes).await.unwrap();
1126 bytes
1127 }
1128
1129 async fn test_clients() -> Arc<RwLock<ClientRegistry>> {
1130 let clients = Arc::new(RwLock::new(HashMap::default()));
1131 clients.write().await.insert(
1132 "client".into(),
1133 PasswordHashString::new(
1134 "$argon2id$v=19$m=10,t=1,p=1$THh0RHE5YWNkQUZNa2lqUA$dmB4X7J49jjCGA",
1135 )
1136 .unwrap(),
1137 );
1138
1139 clients
1140 }
1141
1142 fn frame_bytes<M: Payload + Encode>(message: &M) -> Vec<u8> {
1143 let raw = encode_message(message).unwrap();
1144 let mut bytes = vec![raw.message_type];
1145 bytes.extend_from_slice(&(raw.payload.len() as u32).to_le_bytes());
1146 bytes.extend(raw.payload);
1147
1148 bytes
1149 }
1150
1151 fn dummy_engine_tx() -> mpsc::Sender<(EngineRequest, oneshot::Sender<ServerMessage>)> {
1152 let (tx, _rx) = mpsc::channel(1);
1153 tx
1154 }
1155
1156 fn dummy_event_tx() -> broadcast::Sender<Event> {
1157 let (tx, _rx) = broadcast::channel(1);
1158 tx
1159 }
1160
1161 fn build_handshake(min_version: u8, max_version: u8) -> Vec<u8> {
1162 let mut buf = Vec::new();
1163 let payload = [min_version, max_version];
1164
1165 buf.push(0x01);
1166 buf.extend(&(payload.len() as u32).to_le_bytes());
1167 buf.extend(&payload);
1168
1169 buf
1170 }
1171
1172 fn build_authentication(client_id: &str, password: &str, ip_addr: &str) -> Vec<u8> {
1173 use std::net::IpAddr;
1174
1175 let mut buf = Vec::new();
1176
1177 let client_id_bytes = client_id.as_bytes();
1178 let password_bytes = password.as_bytes();
1179
1180 let mut payload = Vec::new();
1181 payload.push(client_id_bytes.len() as u8);
1182 payload.extend(client_id_bytes);
1183
1184 payload.push(password_bytes.len() as u8);
1185 payload.extend(password_bytes);
1186
1187 let ip: IpAddr = ip_addr.parse().expect("Invalid IP address");
1188 match ip {
1189 IpAddr::V4(v4) => {
1190 payload.push(4); payload.extend(&v4.octets());
1192 }
1193 IpAddr::V6(v6) => {
1194 payload.push(6); payload.extend(&v6.octets());
1196 }
1197 }
1198
1199 buf.push(0x03);
1200 buf.extend(&(payload.len() as u32).to_le_bytes());
1201 buf.extend(payload);
1202
1203 buf
1204 }
1205}