1use std::collections::VecDeque;
4use std::env;
5use std::fmt;
6use std::net::IpAddr;
7use std::str::FromStr;
8use std::sync::Arc;
9
10use futures_util::StreamExt as _;
11use rustls::ClientConfig;
12use rustls::RootCertStore;
13use rustls::pki_types::ServerName;
14use tokio::io::{AsyncRead, AsyncWrite, ReadHalf, WriteHalf};
15use tokio::net::TcpStream;
16use tokio::sync::Mutex;
17use tokio_rustls::TlsConnector;
18use tokio_tungstenite::tungstenite::http::uri::Authority;
19use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
20use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async_with_config};
21
22use crate::wire::{
23 ClientFrame, ClientKind, ClientMessage, FrameReader, MAX_FRAME_BYTES, ServerFrame,
24 ServerMessage, framed_to_websocket, read_frame, validate_version, websocket_error,
25 websocket_to_framed, write_frame,
26};
27use crate::{Error, Result};
28
29const DEFAULT_ENDPOINT: &str = "tcp://127.0.0.1:8741";
30const WEBSOCKET_BRIDGE_BYTES: usize = 16 * 1024;
31pub const MAX_PENDING_FRAMES: usize = 1024;
33
34trait Transport: AsyncRead + AsyncWrite + Unpin + Send {}
35impl<T> Transport for T where T: AsyncRead + AsyncWrite + Unpin + Send {}
36
37type BoxedTransport = Box<dyn Transport>;
38type GatewayWebSocket = WebSocketStream<MaybeTlsStream<TcpStream>>;
39
40#[derive(Debug, Clone, PartialEq, Eq)]
42pub struct Endpoint {
43 security: Security,
44 host: String,
45 port: u16,
46}
47
48#[derive(Debug, Clone, Copy, PartialEq, Eq)]
49enum Security {
50 Plaintext,
51 Tls,
52 WebSocketTls,
53}
54
55#[derive(Debug, Clone, PartialEq, Eq)]
57pub struct PairedClient {
58 pub client_id: String,
59 pub token: String,
60}
61
62pub struct GatewayClient {
64 sender: GatewaySender,
65 events: GatewayEvents,
66}
67
68#[derive(Clone)]
70pub struct GatewaySender {
71 writer: Arc<Mutex<WriteHalf<BoxedTransport>>>,
72}
73
74pub struct GatewayEvents {
76 reader: FrameReader<ReadHalf<BoxedTransport>>,
77 pending: VecDeque<ServerFrame>,
78}
79
80impl Endpoint {
81 pub fn from_env() -> Result<Self> {
83 env::var("MOBIUS_GATEWAY_ENDPOINT")
84 .unwrap_or_else(|_| DEFAULT_ENDPOINT.into())
85 .parse()
86 }
87
88 #[must_use]
90 pub const fn is_plaintext(&self) -> bool {
91 matches!(self.security, Security::Plaintext)
92 }
93
94 #[must_use]
96 pub const fn is_websocket(&self) -> bool {
97 matches!(self.security, Security::WebSocketTls)
98 }
99
100 #[must_use]
102 pub(crate) fn host(&self) -> &str {
103 &self.host
104 }
105
106 async fn connect(&self) -> Result<BoxedTransport> {
107 if self.is_websocket() {
108 return self.connect_websocket().await;
109 }
110 let address = format_address(&self.host, self.port);
111 let stream = TcpStream::connect(&address).await?;
112 if self.security == Security::Plaintext {
113 let peer = stream.peer_addr()?;
114 if !peer.ip().is_loopback() {
115 return Err(Error::Config(
116 "plaintext gateway connections are restricted to loopback".into(),
117 ));
118 }
119 return Ok(Box::new(stream));
120 }
121
122 let mut roots = RootCertStore::empty();
123 roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
124 let config = ClientConfig::builder()
125 .with_root_certificates(roots)
126 .with_no_client_auth();
127 let name = ServerName::try_from(self.host.clone())
128 .map_err(|_| Error::Config("TLS endpoint has an invalid server name".into()))?;
129 let stream = TlsConnector::from(Arc::new(config))
130 .connect(name, stream)
131 .await
132 .map_err(|error| Error::Protocol(format!("TLS handshake failed: {error}")))?;
133 Ok(Box::new(stream))
134 }
135
136 async fn connect_websocket(&self) -> Result<BoxedTransport> {
137 let config = WebSocketConfig::default()
138 .max_message_size(Some(MAX_FRAME_BYTES))
139 .max_frame_size(Some(MAX_FRAME_BYTES));
140 let (websocket, _) = connect_async_with_config(self.to_string(), Some(config), false)
141 .await
142 .map_err(websocket_error)?;
143 let (transport, bridge) = tokio::io::duplex(WEBSOCKET_BRIDGE_BYTES);
144 tokio::spawn(async move {
145 let _result = bridge_websocket(websocket, bridge).await;
146 });
147 Ok(Box::new(transport))
148 }
149}
150
151impl FromStr for Endpoint {
152 type Err = Error;
153
154 fn from_str(value: &str) -> Result<Self> {
155 let (security, authority) = if let Some(authority) = value.strip_prefix("tcp://") {
156 (Security::Plaintext, authority)
157 } else if let Some(authority) = value.strip_prefix("tls://") {
158 (Security::Tls, authority)
159 } else if let Some(authority) = value.strip_prefix("wss://") {
160 (Security::WebSocketTls, authority)
161 } else {
162 return Err(Error::Config(
163 "gateway endpoint must use tcp://, tls://, or wss://".into(),
164 ));
165 };
166 if authority.contains(['/', '?', '#', '@']) {
167 return Err(Error::Config(
168 "gateway endpoint must contain only a host and port".into(),
169 ));
170 }
171 let authority = authority
172 .parse::<Authority>()
173 .map_err(|_| Error::Config("gateway endpoint has an invalid host or port".into()))?;
174 let host = authority
175 .host()
176 .strip_prefix('[')
177 .and_then(|host| host.strip_suffix(']'))
178 .unwrap_or_else(|| authority.host());
179 if host.is_empty() {
180 return Err(Error::Config("gateway endpoint requires a host".into()));
181 }
182 let port = match authority.port_u16() {
183 Some(port) => port,
184 None if authority.as_str().len() != authority.host().len() => {
185 return Err(Error::Config("gateway endpoint has an invalid port".into()));
186 }
187 None if security == Security::WebSocketTls => 443,
188 None => return Err(Error::Config("gateway endpoint requires a port".into())),
189 };
190 if port == 0 {
191 return Err(Error::Config(
192 "gateway endpoint port must be greater than zero".into(),
193 ));
194 }
195 if security == Security::Plaintext && !plaintext_host_is_loopback(host) {
196 return Err(Error::Config(
197 "tcp:// endpoints are restricted to loopback; use tls:// or wss:// remotely".into(),
198 ));
199 }
200 Ok(Self {
201 security,
202 host: host.into(),
203 port,
204 })
205 }
206}
207
208impl fmt::Display for Endpoint {
209 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
210 let scheme = match self.security {
211 Security::Plaintext => "tcp",
212 Security::Tls => "tls",
213 Security::WebSocketTls => "wss",
214 };
215 if self.security == Security::WebSocketTls && self.port == 443 {
216 if self.host.contains(':') {
217 return write!(formatter, "{scheme}://[{}]", self.host);
218 }
219 return write!(formatter, "{scheme}://{}", self.host);
220 }
221 write!(
222 formatter,
223 "{scheme}://{}",
224 format_address(&self.host, self.port)
225 )
226 }
227}
228
229async fn bridge_websocket(
230 websocket: GatewayWebSocket,
231 bridge: tokio::io::DuplexStream,
232) -> Result<()> {
233 let (outgoing, incoming) = websocket.split();
234 let (reader, writer) = tokio::io::split(bridge);
235 tokio::select! {
236 result = websocket_to_framed(incoming, writer) => result,
237 result = framed_to_websocket(reader, outgoing) => result,
238 }
239}
240
241impl GatewayClient {
242 pub async fn connect(
244 endpoint: &Endpoint,
245 token: impl Into<String>,
246 client_kind: ClientKind,
247 ) -> Result<Self> {
248 let transport = endpoint.connect().await?;
249 let (reader, writer) = tokio::io::split(transport);
250 let client = Self::from_parts(reader, writer);
251 client
252 .sender
253 .write(ClientMessage::Authenticate {
254 token: token.into(),
255 client_kind,
256 })
257 .await?;
258 client.expect_authenticated().await
259 }
260
261 pub async fn pair(
263 endpoint: &Endpoint,
264 code: impl Into<String>,
265 client_label: impl Into<String>,
266 client_kind: ClientKind,
267 ) -> Result<(Self, PairedClient)> {
268 let transport = endpoint.connect().await?;
269 let (reader, writer) = tokio::io::split(transport);
270 let mut client = Self::from_parts(reader, writer);
271 client
272 .sender
273 .write(ClientMessage::Pair {
274 code: code.into(),
275 client_label: client_label.into(),
276 client_kind,
277 })
278 .await?;
279 let frame = client
280 .events
281 .next()
282 .await?
283 .ok_or_else(|| Error::Protocol("gateway closed during pairing".into()))?;
284 let paired = match frame.message {
285 ServerMessage::Paired { client_id, token } => PairedClient { client_id, token },
286 ServerMessage::Error { code, message, .. } => {
287 return Err(connection_error(&code, message));
288 }
289 _ => {
290 return Err(Error::Protocol(
291 "gateway did not return a paired response".into(),
292 ));
293 }
294 };
295 client = client.expect_authenticated().await?;
296 Ok((client, paired))
297 }
298
299 #[must_use]
301 pub fn into_parts(self) -> (GatewaySender, GatewayEvents) {
302 (self.sender, self.events)
303 }
304
305 fn from_parts(reader: ReadHalf<BoxedTransport>, writer: WriteHalf<BoxedTransport>) -> Self {
306 Self {
307 sender: GatewaySender {
308 writer: Arc::new(Mutex::new(writer)),
309 },
310 events: GatewayEvents {
311 reader: FrameReader::new(reader),
312 pending: VecDeque::new(),
313 },
314 }
315 }
316
317 async fn expect_authenticated(mut self) -> Result<Self> {
318 let frame = self
319 .events
320 .next()
321 .await?
322 .ok_or_else(|| Error::Protocol("gateway closed during authentication".into()))?;
323 match frame.message {
324 ServerMessage::Authenticated => Ok(self),
325 ServerMessage::Error { code, message, .. } => Err(connection_error(&code, message)),
326 _ => Err(Error::Protocol(
327 "gateway did not acknowledge authentication".into(),
328 )),
329 }
330 }
331}
332
333fn connection_error(code: &str, message: String) -> Error {
334 if code == "unauthorized" {
335 Error::Unauthorized
336 } else {
337 Error::Protocol(message)
338 }
339}
340
341impl GatewaySender {
342 pub async fn send(&self, message: ClientMessage) -> Result<()> {
344 if matches!(
345 message,
346 ClientMessage::Pair { .. } | ClientMessage::Authenticate { .. }
347 ) {
348 return Err(Error::Protocol(
349 "authentication messages are valid only during connection setup".into(),
350 ));
351 }
352 self.write(message).await
353 }
354
355 async fn write(&self, message: ClientMessage) -> Result<()> {
356 let mut writer = self.writer.lock().await;
357 write_frame(&mut *writer, &ClientFrame::new(message)).await
358 }
359}
360
361impl GatewayEvents {
362 pub async fn next(&mut self) -> Result<Option<ServerFrame>> {
364 if let Some(frame) = self.pending.pop_front() {
365 return Ok(Some(frame));
366 }
367 let Some(frame) = read_frame::<ServerFrame>(&mut self.reader).await? else {
368 return Ok(None);
369 };
370 validate_version(frame.version)?;
371 Ok(Some(frame))
372 }
373
374 pub fn prepend(&mut self, frames: Vec<ServerFrame>) -> Result<()> {
376 if self.pending.len() + frames.len() > MAX_PENDING_FRAMES {
377 return Err(Error::Protocol(format!(
378 "gateway event backlog exceeds {MAX_PENDING_FRAMES} frames"
379 )));
380 }
381 for frame in &frames {
382 validate_version(frame.version)?;
383 }
384 for frame in frames.into_iter().rev() {
385 self.pending.push_front(frame);
386 }
387 Ok(())
388 }
389}
390
391pub fn token_from_env() -> Result<String> {
393 env::var("MOBIUS_GATEWAY_TOKEN")
394 .ok()
395 .filter(|token| !token.trim().is_empty())
396 .ok_or_else(|| Error::Config("set MOBIUS_GATEWAY_TOKEN before connecting".into()))
397}
398
399fn plaintext_host_is_loopback(host: &str) -> bool {
400 host.eq_ignore_ascii_case("localhost")
401 || host
402 .parse::<IpAddr>()
403 .is_ok_and(|address| address.is_loopback())
404}
405
406fn format_address(host: &str, port: u16) -> String {
407 if host.contains(':') {
408 format!("[{host}]:{port}")
409 } else {
410 format!("{host}:{port}")
411 }
412}
413
414#[cfg(test)]
415mod tests {
416 use super::*;
417
418 #[tokio::test]
419 async fn connect_authenticates_without_a_session_cursor() {
420 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
421 .await
422 .expect("bind gateway");
423 let endpoint = format!("tcp://{}", listener.local_addr().expect("gateway address"))
424 .parse::<Endpoint>()
425 .expect("gateway endpoint");
426 let gateway = tokio::spawn(async move {
427 let (stream, _) = listener.accept().await.expect("accept client");
428 let (reader, mut writer) = tokio::io::split(stream);
429 let mut reader = FrameReader::new(reader);
430 let frame = read_frame::<ClientFrame>(&mut reader)
431 .await
432 .expect("read authentication")
433 .expect("authentication frame");
434 write_frame(&mut writer, &ServerFrame::new(ServerMessage::Authenticated))
435 .await
436 .expect("acknowledge authentication");
437 frame
438 });
439
440 let _client = GatewayClient::connect(&endpoint, "secret", ClientKind::Cli)
441 .await
442 .expect("connect client");
443 let frame = gateway.await.expect("gateway task");
444
445 assert_eq!(
446 frame.message,
447 ClientMessage::Authenticate {
448 token: "secret".into(),
449 client_kind: ClientKind::Cli,
450 }
451 );
452 }
453
454 #[test]
455 fn endpoint_rejects_remote_plaintext() {
456 let error = "tcp://example.com:8741"
457 .parse::<Endpoint>()
458 .expect_err("remote plaintext must fail");
459
460 assert!(error.to_string().contains("use tls://"));
461 assert!("tcp://127.0.0.1:0".parse::<Endpoint>().is_err());
462 }
463
464 #[test]
465 fn endpoint_accepts_loopback_plaintext_and_remote_encrypted_transports() {
466 let loopback = "tcp://127.0.0.1:8741"
467 .parse::<Endpoint>()
468 .expect("loopback endpoint");
469 let remote = "tls://gateway.example:443"
470 .parse::<Endpoint>()
471 .expect("TLS endpoint");
472 let websocket = "wss://gateway.example"
473 .parse::<Endpoint>()
474 .expect("WSS endpoint");
475
476 assert_eq!(loopback.to_string(), "tcp://127.0.0.1:8741");
477 assert_eq!(remote.to_string(), "tls://gateway.example:443");
478 assert_eq!(websocket.to_string(), "wss://gateway.example");
479 assert!(loopback.is_plaintext());
480 assert!(!remote.is_plaintext());
481 assert!(websocket.is_websocket());
482 }
483
484 #[test]
485 fn authentication_errors_preserve_unauthorized_semantics() {
486 assert!(matches!(
487 connection_error("unauthorized", "authentication failed".into()),
488 Error::Unauthorized
489 ));
490 }
491
492 #[tokio::test]
493 async fn prepended_frames_are_returned_in_order() {
494 let (transport, _peer) = tokio::io::duplex(64);
495 let (reader, _writer) = tokio::io::split(Box::new(transport) as BoxedTransport);
496 let mut events = GatewayEvents {
497 reader: FrameReader::new(reader),
498 pending: VecDeque::new(),
499 };
500 events
501 .prepend(vec![
502 ServerFrame::new(ServerMessage::Accepted {
503 request_id: "first".into(),
504 }),
505 ServerFrame::new(ServerMessage::Accepted {
506 request_id: "second".into(),
507 }),
508 ])
509 .expect("defer frames");
510
511 for expected in ["first", "second"] {
512 let frame = events.next().await.expect("next frame").expect("frame");
513 assert!(matches!(
514 frame.message,
515 ServerMessage::Accepted { request_id } if request_id == expected
516 ));
517 }
518 let mut invalid = ServerFrame::new(ServerMessage::Accepted {
519 request_id: "invalid".into(),
520 });
521 invalid.version = 0;
522 assert!(events.prepend(vec![invalid]).is_err());
523 let frame = ServerFrame::new(ServerMessage::Accepted {
524 request_id: "overflow".into(),
525 });
526 assert!(events.prepend(vec![frame; MAX_PENDING_FRAMES + 1]).is_err());
527 }
528}