1use std::collections::VecDeque;
4use std::env;
5use std::fmt;
6use std::net::IpAddr;
7use std::str::FromStr;
8use std::sync::Arc;
9
10use rustls::ClientConfig;
11use rustls::RootCertStore;
12use rustls::pki_types::ServerName;
13use tokio::io::{AsyncRead, AsyncWrite, ReadHalf, WriteHalf};
14use tokio::net::TcpStream;
15use tokio::sync::Mutex;
16use tokio_rustls::TlsConnector;
17use tokio_tungstenite::connect_async_with_config;
18use tokio_tungstenite::tungstenite::client::IntoClientRequest as _;
19use tokio_tungstenite::tungstenite::http::uri::Authority;
20use tokio_tungstenite::tungstenite::http::{
21 HeaderValue, Request,
22 header::{AUTHORIZATION, SEC_WEBSOCKET_PROTOCOL},
23};
24use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
25
26#[cfg(unix)]
27use crate::wire::read_frame_with_limit;
28use crate::wire::{
29 CatalogHint, ClientFrame, ClientKind, ClientMessage, FrameReader, MAX_PRE_AUTH_FRAME_BYTES,
30 ServerFrame, ServerMessage, read_frame, validate_version, websocket_error, write_frame,
31};
32use crate::{Error, Result};
33
34const DEFAULT_ENDPOINT: &str = "tcp://127.0.0.1:8741";
35use crate::wire::WEBSOCKET_BRIDGE_BYTES;
36pub const MAX_PENDING_FRAMES: usize = 1024;
38
39trait Transport: AsyncRead + AsyncWrite + Unpin + Send {}
40impl<T> Transport for T where T: AsyncRead + AsyncWrite + Unpin + Send {}
41
42type BoxedTransport = Box<dyn Transport>;
43
44#[derive(Debug, Clone, PartialEq, Eq)]
46pub struct Endpoint {
47 security: Security,
48 host: String,
49 port: u16,
50 websocket_authorization: Option<HeaderValue>,
51}
52
53#[derive(Debug, Clone, Copy, PartialEq, Eq)]
54enum Security {
55 Plaintext,
56 Tls,
57 WebSocketTls,
58}
59
60#[derive(Debug, Clone, PartialEq, Eq)]
62pub struct PairedClient {
63 pub client_id: String,
65 pub token: String,
67}
68
69#[derive(Debug, Default)]
71pub struct ConnectOptions {
72 pub catalog: CatalogHint,
74 pub pipelined: Vec<ClientMessage>,
76}
77
78pub struct GatewayClient {
80 sender: GatewaySender,
81 events: GatewayEvents,
82}
83
84#[derive(Clone)]
86pub struct GatewaySender {
87 writer: Arc<Mutex<Option<WriteHalf<BoxedTransport>>>>,
88}
89
90pub struct GatewayEvents {
92 reader: FrameReader<ReadHalf<BoxedTransport>>,
93 pending: VecDeque<ServerFrame>,
94 scoped_deferred: usize,
95}
96
97pub struct GatewayEventScope<'a> {
99 events: &'a mut GatewayEvents,
100 deferred: Vec<ServerFrame>,
101}
102
103impl Endpoint {
104 pub fn from_env() -> Result<Self> {
109 env::var("MOBIUS_GATEWAY_ENDPOINT")
110 .unwrap_or_else(|_| DEFAULT_ENDPOINT.into())
111 .parse()
112 }
113
114 #[must_use]
116 pub const fn is_plaintext(&self) -> bool {
117 matches!(self.security, Security::Plaintext)
118 }
119
120 #[must_use]
122 pub const fn is_websocket(&self) -> bool {
123 matches!(self.security, Security::WebSocketTls)
124 }
125
126 pub fn with_websocket_bearer(mut self, token: &str) -> Result<Self> {
133 if !self.is_websocket()
134 || token.is_empty()
135 || token.bytes().any(|byte| !byte.is_ascii_graphic())
136 {
137 return Err(Error::Config(
138 "invalid secure WebSocket bearer credential".into(),
139 ));
140 }
141 let mut value = HeaderValue::from_str(&format!("Bearer {token}"))
142 .map_err(|_| Error::Config("invalid secure WebSocket bearer credential".into()))?;
143 value.set_sensitive(true);
144 self.websocket_authorization = Some(value);
145 Ok(self)
146 }
147
148 #[must_use]
150 pub(crate) fn host(&self) -> &str {
151 &self.host
152 }
153
154 async fn connect(&self, credential: &str) -> Result<BoxedTransport> {
155 if self.is_websocket() {
156 crate::channel::credential_key(credential)?;
157 return self.connect_websocket(credential).await;
158 }
159 let address = format_address(&self.host, self.port);
160 let stream = TcpStream::connect(&address).await?;
161 self.secure_tcp(stream).await
162 }
163
164 async fn secure_tcp(&self, stream: TcpStream) -> Result<BoxedTransport> {
165 if self.security == Security::Plaintext {
166 let peer = stream.peer_addr()?;
167 if !peer.ip().is_loopback() {
168 return Err(Error::Config(
169 "plaintext gateway connections are restricted to loopback".into(),
170 ));
171 }
172 return Ok(Box::new(stream));
173 }
174
175 let mut roots = RootCertStore::empty();
176 roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
177 let config = ClientConfig::builder()
178 .with_root_certificates(roots)
179 .with_no_client_auth();
180 let name = ServerName::try_from(self.host.clone())
181 .map_err(|_| Error::Config("TLS endpoint has an invalid server name".into()))?;
182 let stream = TlsConnector::from(Arc::new(config))
183 .connect(name, stream)
184 .await
185 .map_err(|error| {
186 Error::Protocol(format!("TLS handshake failed: {:?}", error.kind()))
187 })?;
188 Ok(Box::new(stream))
189 }
190
191 #[cfg(unix)]
192 pub(crate) async fn local_gateway_version(
193 &self,
194 mut listen: std::net::SocketAddr,
195 token: &str,
196 ) -> Result<String> {
197 if self.is_websocket() {
198 return Err(Error::Config(
199 "local gateway control requires TCP or TLS".into(),
200 ));
201 }
202 if listen.ip().is_unspecified() {
203 listen.set_ip(match listen.ip() {
204 IpAddr::V4(_) => std::net::Ipv4Addr::LOCALHOST.into(),
205 IpAddr::V6(_) => std::net::Ipv6Addr::LOCALHOST.into(),
206 });
207 }
208 let mut version = crate::wire::PROTOCOL_VERSION;
209 for attempt in 0..2 {
210 let transport = self.secure_tcp(TcpStream::connect(listen).await?).await?;
211 let (reader, mut writer) = tokio::io::split(transport);
212 let mut reader = FrameReader::new(reader);
213 write_frame(
214 &mut writer,
215 &ClientFrame {
216 version,
217 message: ClientMessage::Authenticate {
218 token: token.into(),
219 client_kind: ClientKind::GatewayDashboard,
220 catalog: CatalogHint::default(),
221 },
222 },
223 )
224 .await?;
225 let response = read_frame_with_limit::<ServerFrame>(&mut reader, 4 * 1024)
226 .await?
227 .ok_or_else(|| {
228 Error::Protocol("gateway closed during version authentication".into())
229 })?;
230 match response.message {
231 ServerMessage::Error { code, .. }
232 if attempt == 0 && code == "protocol_version" && response.version > 0 =>
233 {
234 version = response.version;
235 continue;
236 }
237 ServerMessage::Authenticated if response.version == version => {}
238 ServerMessage::Error { code, message, .. } => {
239 return Err(connection_error(&code, message));
240 }
241 _ => {
242 return Err(Error::Protocol(
243 "gateway did not authenticate the version check".into(),
244 ));
245 }
246 }
247 let frame = read_frame::<serde_json::Value>(&mut reader)
248 .await?
249 .ok_or_else(|| {
250 Error::Protocol("gateway disconnected before reporting its version".into())
251 })?;
252 if frame["type"] != "ready" || frame["version"].as_u64() != Some(u64::from(version)) {
253 return Err(Error::Protocol(
254 "gateway did not report a valid ready frame".into(),
255 ));
256 }
257 return frame["payload"]["gateway_version"]
258 .as_str()
259 .map(str::to_owned)
260 .ok_or_else(|| Error::Protocol("gateway did not report its version".into()));
261 }
262 Err(Error::Protocol("gateway did not report its version".into()))
263 }
264
265 async fn connect_websocket(&self, credential: &str) -> Result<BoxedTransport> {
266 let config = WebSocketConfig::default()
267 .max_message_size(Some(crate::channel::MAX_RECORD))
268 .max_frame_size(Some(crate::channel::MAX_RECORD));
269 let (mut websocket, response) =
270 connect_async_with_config(self.websocket_request()?, Some(config), false)
271 .await
272 .map_err(|error| match error {
273 tokio_tungstenite::tungstenite::Error::Http(response)
274 if self.websocket_authorization.is_some()
275 && matches!(response.status().as_u16(), 401 | 403) =>
276 {
277 Error::Config(
278 "gateway access was denied; sign in to the gateway service again"
279 .into(),
280 )
281 }
282 error => websocket_error(error),
283 })?;
284 if response
285 .headers()
286 .get(SEC_WEBSOCKET_PROTOCOL)
287 .and_then(|value| value.to_str().ok())
288 != Some(crate::channel::SUBPROTOCOL)
289 {
290 return Err(Error::Config(
291 "gateway does not support the encrypted WebSocket protocol; update the gateway"
292 .into(),
293 ));
294 }
295 let state = tokio::time::timeout(
296 std::time::Duration::from_secs(5),
297 crate::channel::client_handshake(&mut websocket, credential),
298 )
299 .await
300 .map_err(|_| Error::Unauthorized)??;
301 let (transport, bridge) = tokio::io::duplex(WEBSOCKET_BRIDGE_BYTES);
302 tokio::spawn(async move {
303 let _result = crate::channel::bridge(websocket, state, bridge).await;
304 });
305 Ok(Box::new(transport))
306 }
307
308 fn websocket_request(&self) -> Result<Request<()>> {
309 let mut request = self
310 .to_string()
311 .into_client_request()
312 .map_err(websocket_error)?;
313 request.headers_mut().insert(
314 SEC_WEBSOCKET_PROTOCOL,
315 HeaderValue::from_static(crate::channel::SUBPROTOCOL),
316 );
317 if let Some(authorization) = &self.websocket_authorization {
318 request
319 .headers_mut()
320 .insert(AUTHORIZATION, authorization.clone());
321 }
322 Ok(request)
323 }
324}
325
326impl FromStr for Endpoint {
327 type Err = Error;
328
329 fn from_str(value: &str) -> Result<Self> {
330 let (security, authority) = if let Some(authority) = value.strip_prefix("tcp://") {
331 (Security::Plaintext, authority)
332 } else if let Some(authority) = value.strip_prefix("tls://") {
333 (Security::Tls, authority)
334 } else if let Some(authority) = value.strip_prefix("wss://") {
335 (Security::WebSocketTls, authority)
336 } else {
337 return Err(Error::Config(
338 "gateway endpoint must use tcp://, tls://, or wss://".into(),
339 ));
340 };
341 if authority.contains(['/', '?', '#', '@']) {
342 return Err(Error::Config(
343 "gateway endpoint must contain only a host and port".into(),
344 ));
345 }
346 let authority = authority
347 .parse::<Authority>()
348 .map_err(|_| Error::Config("gateway endpoint has an invalid host or port".into()))?;
349 let host = authority
350 .host()
351 .strip_prefix('[')
352 .and_then(|host| host.strip_suffix(']'))
353 .unwrap_or_else(|| authority.host());
354 if host.is_empty() {
355 return Err(Error::Config("gateway endpoint requires a host".into()));
356 }
357 let port = match authority.port_u16() {
358 Some(port) => port,
359 None if authority.as_str().len() != authority.host().len() => {
360 return Err(Error::Config("gateway endpoint has an invalid port".into()));
361 }
362 None if security == Security::WebSocketTls => 443,
363 None => return Err(Error::Config("gateway endpoint requires a port".into())),
364 };
365 if port == 0 {
366 return Err(Error::Config(
367 "gateway endpoint port must be greater than zero".into(),
368 ));
369 }
370 if security == Security::Plaintext && !plaintext_host_is_loopback(host) {
371 return Err(Error::Config(
372 "tcp:// endpoints are restricted to loopback; use tls:// or wss:// remotely".into(),
373 ));
374 }
375 Ok(Self {
376 security,
377 host: host.into(),
378 port,
379 websocket_authorization: None,
380 })
381 }
382}
383
384impl fmt::Display for Endpoint {
385 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
386 let scheme = match self.security {
387 Security::Plaintext => "tcp",
388 Security::Tls => "tls",
389 Security::WebSocketTls => "wss",
390 };
391 if self.security == Security::WebSocketTls && self.port == 443 {
392 if self.host.contains(':') {
393 return write!(formatter, "{scheme}://[{}]", self.host);
394 }
395 return write!(formatter, "{scheme}://{}", self.host);
396 }
397 write!(
398 formatter,
399 "{scheme}://{}",
400 format_address(&self.host, self.port)
401 )
402 }
403}
404
405impl GatewayClient {
406 pub async fn connect(
411 endpoint: &Endpoint,
412 token: &str,
413 client_kind: ClientKind,
414 ) -> Result<Self> {
415 Self::connect_with(endpoint, token, client_kind, &mut ConnectOptions::default()).await
416 }
417
418 pub async fn connect_with(
424 endpoint: &Endpoint,
425 token: &str,
426 client_kind: ClientKind,
427 options: &mut ConnectOptions,
428 ) -> Result<Self> {
429 let pipelined_bytes = options.pipelined.iter().try_fold(0, |total, message| {
431 serde_json::to_vec(message).map(|encoded| total + encoded.len())
432 })?;
433 if pipelined_bytes > MAX_PRE_AUTH_FRAME_BYTES / 2 {
434 return Err(Error::Config(
435 "requests sent with authentication must stay under 2 KiB".into(),
436 ));
437 }
438 let transport = endpoint.connect(token).await?;
439 let (reader, writer) = tokio::io::split(transport);
440 let client = Self::from_parts(reader, writer);
441 let authentication = ClientFrame::new(ClientMessage::Authenticate {
442 token: token.to_owned(),
444 client_kind,
445 catalog: std::mem::take(&mut options.catalog),
446 });
447 let pipelined: Vec<_> = std::mem::take(&mut options.pipelined)
448 .into_iter()
449 .map(ClientFrame::new)
450 .collect();
451 let written = client
452 .sender
453 .write_frames(&authentication, &pipelined)
454 .await;
455 if let ClientMessage::Authenticate { catalog, .. } = authentication.message {
456 options.catalog = catalog;
457 }
458 options.pipelined = pipelined.into_iter().map(|frame| frame.message).collect();
459 written?;
460 client.expect_authenticated().await
461 }
462
463 pub async fn pair(
468 endpoint: &Endpoint,
469 code: impl Into<String>,
470 client_label: impl Into<String>,
471 client_kind: ClientKind,
472 ) -> Result<(Self, PairedClient)> {
473 let code = code.into();
474 let transport = endpoint.connect(&code).await?;
475 let (reader, writer) = tokio::io::split(transport);
476 let mut client = Self::from_parts(reader, writer);
477 client
478 .sender
479 .write(ClientMessage::Pair {
480 code,
481 client_label: client_label.into(),
482 client_kind,
483 })
484 .await?;
485 let frame = client
486 .events
487 .next()
488 .await?
489 .ok_or_else(|| Error::Protocol("gateway closed during pairing".into()))?;
490 let paired = match frame.message {
491 ServerMessage::Paired { client_id, token } => PairedClient { client_id, token },
492 ServerMessage::Error { code, message, .. } => {
493 return Err(connection_error(&code, message));
494 }
495 _ => {
496 return Err(Error::Protocol(
497 "gateway did not return a paired response".into(),
498 ));
499 }
500 };
501 client = client.expect_authenticated().await?;
502 Ok((client, paired))
503 }
504
505 #[must_use]
507 pub fn into_parts(self) -> (GatewaySender, GatewayEvents) {
508 (self.sender, self.events)
509 }
510
511 fn from_parts(reader: ReadHalf<BoxedTransport>, writer: WriteHalf<BoxedTransport>) -> Self {
512 Self {
513 sender: GatewaySender {
514 writer: Arc::new(Mutex::new(Some(writer))),
515 },
516 events: GatewayEvents {
517 reader: FrameReader::new(reader),
518 pending: VecDeque::new(),
519 scoped_deferred: 0,
520 },
521 }
522 }
523
524 async fn expect_authenticated(mut self) -> Result<Self> {
525 let frame = self
526 .events
527 .next()
528 .await?
529 .ok_or_else(|| Error::Protocol("gateway closed during authentication".into()))?;
530 match frame.message {
531 ServerMessage::Authenticated => Ok(self),
532 ServerMessage::Error { code, message, .. } => Err(connection_error(&code, message)),
533 _ => Err(Error::Protocol(
534 "gateway did not acknowledge authentication".into(),
535 )),
536 }
537 }
538}
539
540fn connection_error(code: &str, message: String) -> Error {
541 if code == "unauthorized" {
542 Error::Unauthorized
543 } else {
544 Error::Protocol(message)
545 }
546}
547
548impl GatewaySender {
549 pub async fn send(&self, message: ClientMessage) -> Result<()> {
554 if matches!(
555 message,
556 ClientMessage::Pair { .. }
557 | ClientMessage::RepairPairing { .. }
558 | ClientMessage::Authenticate { .. }
559 ) {
560 return Err(Error::Protocol(
561 "authentication messages are valid only during connection setup".into(),
562 ));
563 }
564 self.write(message).await
565 }
566
567 async fn write(&self, message: ClientMessage) -> Result<()> {
568 self.write_frames(&ClientFrame::new(message), &[]).await
569 }
570
571 async fn write_frames(&self, first: &ClientFrame, then: &[ClientFrame]) -> Result<()> {
573 let mut slot = self.writer.lock().await;
574 let mut writer = slot.take().ok_or_else(|| {
575 Error::Protocol("gateway writer is closed after a failed or cancelled write".into())
576 })?;
577 write_frame(&mut writer, first).await?;
578 for frame in then {
579 write_frame(&mut writer, frame).await?;
580 }
581 *slot = Some(writer);
582 Ok(())
583 }
584}
585
586impl GatewayEvents {
587 pub async fn next(&mut self) -> Result<Option<ServerFrame>> {
592 if let Some(frame) = self.pending.pop_front() {
593 return Ok(Some(frame));
594 }
595 let Some(frame) = read_frame::<ServerFrame>(&mut self.reader).await? else {
596 return Ok(None);
597 };
598 validate_version(frame.version)?;
599 Ok(Some(frame))
600 }
601
602 pub fn scoped(&mut self) -> GatewayEventScope<'_> {
604 GatewayEventScope {
605 events: self,
606 deferred: Vec::new(),
607 }
608 }
609}
610
611impl GatewayEventScope<'_> {
612 pub async fn next(&mut self) -> Result<Option<ServerFrame>> {
617 self.events.next().await
618 }
619
620 pub fn reborrow(&mut self) -> &mut GatewayEvents {
622 self.events
623 }
624
625 pub fn defer(&mut self, frame: ServerFrame) -> Result<()> {
630 validate_version(frame.version)?;
631 if self
632 .events
633 .pending
634 .len()
635 .saturating_add(self.events.scoped_deferred)
636 >= MAX_PENDING_FRAMES
637 {
638 return Err(Error::Protocol(format!(
639 "gateway event backlog exceeds {MAX_PENDING_FRAMES} frames"
640 )));
641 }
642 self.deferred.push(frame);
643 self.events.scoped_deferred += 1;
644 Ok(())
645 }
646}
647
648impl Drop for GatewayEventScope<'_> {
649 fn drop(&mut self) {
650 self.events.scoped_deferred -= self.deferred.len();
651 for frame in self.deferred.drain(..).rev() {
652 self.events.pending.push_front(frame);
653 }
654 }
655}
656
657pub fn token_from_env() -> Result<String> {
662 env::var("MOBIUS_GATEWAY_TOKEN")
663 .ok()
664 .filter(|token| !token.trim().is_empty())
665 .ok_or_else(|| Error::Config("set MOBIUS_GATEWAY_TOKEN before connecting".into()))
666}
667
668fn plaintext_host_is_loopback(host: &str) -> bool {
669 host.eq_ignore_ascii_case("localhost")
670 || host
671 .parse::<IpAddr>()
672 .is_ok_and(|address| address.is_loopback())
673}
674
675fn format_address(host: &str, port: u16) -> String {
676 if host.contains(':') {
677 format!("[{host}]:{port}")
678 } else {
679 format!("{host}:{port}")
680 }
681}
682
683#[cfg(test)]
684mod tests {
685 use super::*;
686
687 #[test]
688 fn websocket_admission_is_sensitive_and_separate_from_the_endpoint() {
689 let endpoint: Endpoint = "wss://gateway.example".parse().expect("endpoint");
690 assert!(
691 !endpoint
692 .websocket_request()
693 .expect("request")
694 .headers()
695 .contains_key(AUTHORIZATION)
696 );
697 let endpoint = endpoint
698 .with_websocket_bearer("cloud-secret")
699 .expect("bearer");
700 let request = endpoint.websocket_request().expect("request");
701 assert_eq!(request.headers()[AUTHORIZATION], "Bearer cloud-secret");
702 assert!(request.headers()[AUTHORIZATION].is_sensitive());
703 assert_eq!(endpoint.to_string(), "wss://gateway.example");
704 assert!(!format!("{endpoint:?} {request:?}").contains("cloud-secret"));
705 for invalid in ["", "two words", "secret\r\nInjected: value", "nonascii-é"] {
706 assert!(endpoint.clone().with_websocket_bearer(invalid).is_err());
707 }
708 assert!(
709 "tcp://127.0.0.1:8741"
710 .parse::<Endpoint>()
711 .expect("loopback")
712 .with_websocket_bearer("secret")
713 .is_err()
714 );
715 }
716
717 #[tokio::test]
718 async fn websocket_redirects_are_rejected_without_exposing_response_secrets() {
719 use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
720
721 let endpoint = "wss://gateway.example"
722 .parse::<Endpoint>()
723 .expect("endpoint")
724 .with_websocket_bearer("cloud-secret")
725 .expect("bearer");
726 let (client, mut server) = tokio::io::duplex(2048);
727 let response = tokio::spawn(async move {
728 let mut headers = Vec::new();
729 while !headers.ends_with(b"\r\n\r\n") {
730 headers.push(server.read_u8().await.expect("request header"));
731 }
732 assert!(
733 String::from_utf8(headers)
734 .expect("headers")
735 .contains("Bearer cloud-secret")
736 );
737 server.write_all(b"HTTP/1.1 302 Found\r\nLocation: https://evil.example\r\nContent-Length: 12\r\n\r\ncloud-secret")
738 .await.expect("redirect response");
739 });
740 let error =
741 tokio_tungstenite::client_async(endpoint.websocket_request().expect("request"), client)
742 .await
743 .expect_err("a redirect must not follow the bearer");
744 let error = websocket_error(error).to_string();
745 assert!(error.contains("HTTP 302"));
746 assert!(!error.contains("cloud-secret"));
747 response.await.expect("server");
748 }
749
750 #[tokio::test(start_paused = true)]
751 async fn timed_out_writer_cannot_send_another_frame() {
752 let (transport, _peer) = tokio::io::duplex(4);
753 let (reader, writer) = tokio::io::split(Box::new(transport) as BoxedTransport);
754 let (sender, _events) = GatewayClient::from_parts(reader, writer).into_parts();
755 let request = ClientMessage::ListSessions {
756 request_id: "request".into(),
757 };
758 assert!(
759 matches!(sender.send(request.clone()).await, Err(Error::Io(error)) if error.kind() == std::io::ErrorKind::TimedOut)
760 );
761 assert!(
762 matches!(sender.send(request).await, Err(Error::Protocol(message)) if message.contains("writer is closed"))
763 );
764 }
765
766 #[tokio::test]
767 async fn connect_authenticates_without_a_session_cursor() {
768 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
769 .await
770 .expect("bind gateway");
771 let endpoint = format!("tcp://{}", listener.local_addr().expect("gateway address"))
772 .parse::<Endpoint>()
773 .expect("gateway endpoint");
774 let gateway = tokio::spawn(async move {
775 let (stream, _) = listener.accept().await.expect("accept client");
776 let (reader, mut writer) = tokio::io::split(stream);
777 let mut reader = FrameReader::new(reader);
778 let frame = read_frame::<ClientFrame>(&mut reader)
779 .await
780 .expect("read authentication")
781 .expect("authentication frame");
782 write_frame(&mut writer, &ServerFrame::new(ServerMessage::Authenticated))
783 .await
784 .expect("acknowledge authentication");
785 frame
786 });
787
788 let _client = GatewayClient::connect(&endpoint, "secret", ClientKind::Cli)
789 .await
790 .expect("connect client");
791 let frame = gateway.await.expect("gateway task");
792
793 assert_eq!(
794 frame.message,
795 ClientMessage::Authenticate {
796 token: "secret".into(),
797 client_kind: ClientKind::Cli,
798 catalog: CatalogHint::default(),
799 }
800 );
801 }
802
803 #[test]
804 fn endpoint_rejects_remote_plaintext() {
805 let error = "tcp://example.com:8741"
806 .parse::<Endpoint>()
807 .expect_err("remote plaintext must fail");
808
809 assert!(error.to_string().contains("use tls://"));
810 assert!("tcp://127.0.0.1:0".parse::<Endpoint>().is_err());
811 }
812
813 #[test]
814 fn endpoint_accepts_loopback_plaintext_and_remote_encrypted_transports() {
815 let loopback = "tcp://127.0.0.1:8741"
816 .parse::<Endpoint>()
817 .expect("loopback endpoint");
818 let remote = "tls://gateway.example:443"
819 .parse::<Endpoint>()
820 .expect("TLS endpoint");
821 let websocket = "wss://gateway.example"
822 .parse::<Endpoint>()
823 .expect("WSS endpoint");
824
825 assert_eq!(loopback.to_string(), "tcp://127.0.0.1:8741");
826 assert_eq!(remote.to_string(), "tls://gateway.example:443");
827 assert_eq!(websocket.to_string(), "wss://gateway.example");
828 assert!(loopback.is_plaintext());
829 assert!(!remote.is_plaintext());
830 assert!(websocket.is_websocket());
831 }
832
833 #[test]
834 fn authentication_errors_preserve_unauthorized_semantics() {
835 assert!(matches!(
836 connection_error("unauthorized", "authentication failed".into()),
837 Error::Unauthorized
838 ));
839 }
840
841 #[tokio::test]
842 async fn scoped_wait_restores_deferred_frames_on_early_return() {
843 let (transport, _peer) = tokio::io::duplex(64);
844 let (reader, _writer) = tokio::io::split(Box::new(transport) as BoxedTransport);
845 let mut events = GatewayEvents {
846 reader: FrameReader::new(reader),
847 pending: VecDeque::from([
848 ServerFrame::new(ServerMessage::Accepted {
849 request_id: "unrelated".into(),
850 }),
851 ServerFrame::new(ServerMessage::Accepted {
852 request_id: "expected".into(),
853 }),
854 ]),
855 scoped_deferred: 0,
856 };
857
858 {
859 let mut scope = events.scoped();
860 let unrelated = scope.next().await.expect("next").expect("unrelated");
861 scope.defer(unrelated).expect("defer");
862 let expected = scope.next().await.expect("next").expect("expected");
863 assert!(matches!(
864 expected.message,
865 ServerMessage::Accepted { request_id } if request_id == "expected"
866 ));
867 }
868
869 let restored = events.next().await.expect("next").expect("restored");
870 assert!(matches!(
871 restored.message,
872 ServerMessage::Accepted { request_id } if request_id == "unrelated"
873 ));
874 }
875
876 #[tokio::test]
877 async fn nested_scopes_share_the_bounded_deferred_backlog() {
878 let (transport, mut peer) = tokio::io::duplex(1024);
879 write_frame(
880 &mut peer,
881 &ServerFrame::new(ServerMessage::Accepted {
882 request_id: "overflow".into(),
883 }),
884 )
885 .await
886 .expect("write overflow frame");
887 let (reader, _writer) = tokio::io::split(Box::new(transport) as BoxedTransport);
888 let pending = (0..MAX_PENDING_FRAMES)
889 .map(|index| {
890 ServerFrame::new(ServerMessage::Accepted {
891 request_id: index.to_string(),
892 })
893 })
894 .collect();
895 let mut events = GatewayEvents {
896 reader: FrameReader::new(reader),
897 pending,
898 scoped_deferred: 0,
899 };
900
901 {
902 let mut outer = events.scoped();
903 let frame = outer.next().await.expect("next").expect("outer frame");
904 outer.defer(frame).expect("outer defer");
905 let mut inner = outer.reborrow().scoped();
906 for _ in 1..MAX_PENDING_FRAMES {
907 let frame = inner.next().await.expect("next").expect("inner frame");
908 inner.defer(frame).expect("inner defer");
909 }
910 let overflow = inner.next().await.expect("next").expect("overflow frame");
911 assert!(inner.defer(overflow).is_err());
912 }
913
914 assert_eq!(events.pending.len(), MAX_PENDING_FRAMES);
915 }
916
917 #[tokio::test]
918 async fn scoped_defer_rejects_an_invalid_protocol_version() {
919 let (transport, _peer) = tokio::io::duplex(64);
920 let (reader, _writer) = tokio::io::split(Box::new(transport) as BoxedTransport);
921 let mut events = GatewayEvents {
922 reader: FrameReader::new(reader),
923 pending: VecDeque::new(),
924 scoped_deferred: 0,
925 };
926 let mut frame = ServerFrame::new(ServerMessage::Accepted {
927 request_id: "invalid".into(),
928 });
929 frame.version = crate::wire::PROTOCOL_VERSION.saturating_sub(1);
930
931 assert!(events.scoped().defer(frame).is_err());
932 assert!(events.pending.is_empty());
933 }
934}