Skip to main content

bloop_server_framework/
network.rs

1//! TLS-based network server for handling authenticated client connections.
2//!
3//! This module provides a [`NetworkListener`] that manages client
4//! authentication, request processing, and message dispatching over a secure
5//! TCP connection.
6//!
7//! The listener is generic over a pair of extension message sets. Custom
8//! client messages decode into `Req` and are forwarded, fully typed, over the
9//! custom request channel; the handler answers with a [`CustomOutcome`]
10//! carrying either a `Res` message or a protocol error. Servers without
11//! protocol extensions use the default
12//! [`NoExtension`](bloop_protocol::set::NoExtension) sets and never touch any
13//! of this.
14
15use 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
50/// Maps client IDs to their stored password hashes (Argon2).
51///
52/// Used during client authentication.
53pub type ClientRegistry = HashMap<String, PasswordHashString>;
54
55/// Default maximum payload length accepted from clients.
56///
57/// Standard client-to-server payloads are tiny (authentication is the
58/// largest); the default leaves generous room for custom extension messages
59/// while bounding what a hostile length field can allocate.
60pub 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/// The application's answer to a custom request.
96#[derive(Debug)]
97pub enum CustomOutcome<Res> {
98    /// A custom response message.
99    Response(Res),
100
101    /// A protocol error response, e.g. a custom error code.
102    Error(ErrorResponse),
103}
104
105/// A custom client request forwarded to the application layer.
106///
107/// The protocol is strictly request-response, so the handler must send an
108/// outcome for every request; dropping the sender tears down the client
109/// connection.
110#[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
117/// A TLS-secured TCP server that accepts client connections, authenticates
118/// them, and dispatches their requests to the appropriate handlers.
119///
120/// This listener handles authentication, version negotiation, and supports
121/// custom client messages through the `Req`/`Res` extension sets.
122///
123/// # Examples
124///
125/// ```no_run
126/// use std::sync::Arc;
127/// use tokio::sync::{mpsc, broadcast, RwLock};
128/// use bloop_server_framework::network::NetworkListenerBuilder;
129///
130/// #[tokio::main]
131/// async fn main() {
132///   let clients = Arc::new(RwLock::new(Default::default()));
133///   let (engine_tx, _) = mpsc::channel(10);
134///   let (event_tx, _) = broadcast::channel(10);
135///
136///   let listener = NetworkListenerBuilder::new()
137///       .address("127.0.0.1:12345")
138///       .cert_path("server.crt")
139///       .key_path("server.key")
140///       .clients(clients)
141///       .engine_tx(engine_tx)
142///       .event_tx(event_tx)
143///       .build()
144///       .unwrap();
145///
146///   listener.listen().await.unwrap();
147/// }
148/// ```
149pub 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    /// Starts listening for incoming TCP connections.
178    ///
179    /// This method blocks indefinitely, accepting and processing new connections.
180    /// Each connection is handled asynchronously in its own task.
181    ///
182    /// Returns an error only if the server fails to bind to the specified address.
183    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
300/// Maps a connection error onto the error response owed to the client.
301///
302/// Returns `None` for transport errors, where the connection is gone and no
303/// response can be delivered.
304fn 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/// Builder for [`NetworkListener`].
368///
369/// This allows configuring the address, TLS certificates, client registry,
370/// message channels, and custom request handlers. Registering a custom
371/// request channel via [`custom_req_tx`](Self::custom_req_tx) selects the
372/// listener's extension sets; without one the listener speaks the standard
373/// protocol only.
374///
375/// # Examples
376///
377/// ```
378/// use std::sync::Arc;
379/// use tokio::sync::{broadcast, RwLock};
380/// use tokio::sync::mpsc;
381/// use bloop_server_framework::network::NetworkListenerBuilder;
382///
383/// let (engine_tx, engine_rx) = mpsc::channel(512);
384/// let (event_tx, event_rx) = broadcast::channel(512);
385///
386/// # let _ = rustls::crypto::aws_lc_rs::default_provider().install_default();
387/// let builder = NetworkListenerBuilder::new()
388///     .address("127.0.0.1:12345")
389///     .clients(Arc::new(RwLock::new(Default::default())))
390///     .engine_tx(engine_tx)
391///     .event_tx(event_tx)
392///     .cert_path("examples/cert.pem")
393///     .key_path("examples/key.pem")
394///     .build()
395///     .unwrap();
396/// ```
397#[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    /// Creates a builder for a listener without protocol extensions.
411    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    /// Registers the custom request channel, selecting the listener's
433    /// extension sets from the channel's message type.
434    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    /// Overrides the maximum payload length accepted from clients.
484    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    /// Builds the [`NetworkListener`] from the provided configuration.
490    ///
491    /// Returns an error if required fields are missing, or TLS setup fails.
492    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
543/// Handles an authenticated client connection.
544///
545/// Reads client messages from the stream, dispatches them to the appropriate
546/// handlers, and sends back server responses.
547async 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
630/// Authenticates a client by performing handshake and credential verification.
631///
632/// Returns the client ID, its IP address, and the negotiated protocol version
633/// on success.
634async 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    /// Scripted raw-byte session covering the whole v3 message flow.
893    ///
894    /// This is the load-bearing wire-compatibility test: the client side of
895    /// the duplex writes literal protocol bytes and asserts literal response
896    /// bytes, so codec regressions cannot cancel out.
897    #[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        // Fake engine answering fixed responses.
905        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        // Custom handler echoing the text reversed, or answering a custom
931        // error code for the magic word.
932        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        // Handshake.
981        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        // Authentication.
987        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        // Ping.
994        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        // Bloop.
998        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        // Retrieve audio.
1009        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        // Preload check without a stored hash.
1018        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        // Custom echo request.
1022        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        // Custom request answered with a custom error code.
1032        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        // Quit.
1039        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); // IPv4
1191                payload.extend(&v4.octets());
1192            }
1193            IpAddr::V6(v6) => {
1194                payload.push(6); // IPv6
1195                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}