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| {
133 Error::Protocol(format!("TLS handshake failed: {:?}", error.kind()))
134 })?;
135 Ok(Box::new(stream))
136 }
137
138 async fn connect_websocket(&self) -> Result<BoxedTransport> {
139 let config = WebSocketConfig::default()
140 .max_message_size(Some(MAX_FRAME_BYTES))
141 .max_frame_size(Some(MAX_FRAME_BYTES));
142 let (websocket, _) = connect_async_with_config(self.to_string(), Some(config), false)
143 .await
144 .map_err(websocket_error)?;
145 let (transport, bridge) = tokio::io::duplex(WEBSOCKET_BRIDGE_BYTES);
146 tokio::spawn(async move {
147 let _result = bridge_websocket(websocket, bridge).await;
148 });
149 Ok(Box::new(transport))
150 }
151}
152
153impl FromStr for Endpoint {
154 type Err = Error;
155
156 fn from_str(value: &str) -> Result<Self> {
157 let (security, authority) = if let Some(authority) = value.strip_prefix("tcp://") {
158 (Security::Plaintext, authority)
159 } else if let Some(authority) = value.strip_prefix("tls://") {
160 (Security::Tls, authority)
161 } else if let Some(authority) = value.strip_prefix("wss://") {
162 (Security::WebSocketTls, authority)
163 } else {
164 return Err(Error::Config(
165 "gateway endpoint must use tcp://, tls://, or wss://".into(),
166 ));
167 };
168 if authority.contains(['/', '?', '#', '@']) {
169 return Err(Error::Config(
170 "gateway endpoint must contain only a host and port".into(),
171 ));
172 }
173 let authority = authority
174 .parse::<Authority>()
175 .map_err(|_| Error::Config("gateway endpoint has an invalid host or port".into()))?;
176 let host = authority
177 .host()
178 .strip_prefix('[')
179 .and_then(|host| host.strip_suffix(']'))
180 .unwrap_or_else(|| authority.host());
181 if host.is_empty() {
182 return Err(Error::Config("gateway endpoint requires a host".into()));
183 }
184 let port = match authority.port_u16() {
185 Some(port) => port,
186 None if authority.as_str().len() != authority.host().len() => {
187 return Err(Error::Config("gateway endpoint has an invalid port".into()));
188 }
189 None if security == Security::WebSocketTls => 443,
190 None => return Err(Error::Config("gateway endpoint requires a port".into())),
191 };
192 if port == 0 {
193 return Err(Error::Config(
194 "gateway endpoint port must be greater than zero".into(),
195 ));
196 }
197 if security == Security::Plaintext && !plaintext_host_is_loopback(host) {
198 return Err(Error::Config(
199 "tcp:// endpoints are restricted to loopback; use tls:// or wss:// remotely".into(),
200 ));
201 }
202 Ok(Self {
203 security,
204 host: host.into(),
205 port,
206 })
207 }
208}
209
210impl fmt::Display for Endpoint {
211 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
212 let scheme = match self.security {
213 Security::Plaintext => "tcp",
214 Security::Tls => "tls",
215 Security::WebSocketTls => "wss",
216 };
217 if self.security == Security::WebSocketTls && self.port == 443 {
218 if self.host.contains(':') {
219 return write!(formatter, "{scheme}://[{}]", self.host);
220 }
221 return write!(formatter, "{scheme}://{}", self.host);
222 }
223 write!(
224 formatter,
225 "{scheme}://{}",
226 format_address(&self.host, self.port)
227 )
228 }
229}
230
231async fn bridge_websocket(
232 websocket: GatewayWebSocket,
233 bridge: tokio::io::DuplexStream,
234) -> Result<()> {
235 let (outgoing, incoming) = websocket.split();
236 let (reader, writer) = tokio::io::split(bridge);
237 tokio::select! {
238 result = websocket_to_framed(incoming, writer) => result,
239 result = framed_to_websocket(reader, outgoing) => result,
240 }
241}
242
243impl GatewayClient {
244 pub async fn connect(
246 endpoint: &Endpoint,
247 token: impl Into<String>,
248 client_kind: ClientKind,
249 ) -> Result<Self> {
250 let transport = endpoint.connect().await?;
251 let (reader, writer) = tokio::io::split(transport);
252 let client = Self::from_parts(reader, writer);
253 client
254 .sender
255 .write(ClientMessage::Authenticate {
256 token: token.into(),
257 client_kind,
258 })
259 .await?;
260 client.expect_authenticated().await
261 }
262
263 pub async fn pair(
265 endpoint: &Endpoint,
266 code: impl Into<String>,
267 client_label: impl Into<String>,
268 client_kind: ClientKind,
269 ) -> Result<(Self, PairedClient)> {
270 let transport = endpoint.connect().await?;
271 let (reader, writer) = tokio::io::split(transport);
272 let mut client = Self::from_parts(reader, writer);
273 client
274 .sender
275 .write(ClientMessage::Pair {
276 code: code.into(),
277 client_label: client_label.into(),
278 client_kind,
279 })
280 .await?;
281 let frame = client
282 .events
283 .next()
284 .await?
285 .ok_or_else(|| Error::Protocol("gateway closed during pairing".into()))?;
286 let paired = match frame.message {
287 ServerMessage::Paired { client_id, token } => PairedClient { client_id, token },
288 ServerMessage::Error { code, message, .. } => {
289 return Err(connection_error(&code, message));
290 }
291 _ => {
292 return Err(Error::Protocol(
293 "gateway did not return a paired response".into(),
294 ));
295 }
296 };
297 client = client.expect_authenticated().await?;
298 Ok((client, paired))
299 }
300
301 #[must_use]
303 pub fn into_parts(self) -> (GatewaySender, GatewayEvents) {
304 (self.sender, self.events)
305 }
306
307 fn from_parts(reader: ReadHalf<BoxedTransport>, writer: WriteHalf<BoxedTransport>) -> Self {
308 Self {
309 sender: GatewaySender {
310 writer: Arc::new(Mutex::new(writer)),
311 },
312 events: GatewayEvents {
313 reader: FrameReader::new(reader),
314 pending: VecDeque::new(),
315 },
316 }
317 }
318
319 async fn expect_authenticated(mut self) -> Result<Self> {
320 let frame = self
321 .events
322 .next()
323 .await?
324 .ok_or_else(|| Error::Protocol("gateway closed during authentication".into()))?;
325 match frame.message {
326 ServerMessage::Authenticated => Ok(self),
327 ServerMessage::Error { code, message, .. } => Err(connection_error(&code, message)),
328 _ => Err(Error::Protocol(
329 "gateway did not acknowledge authentication".into(),
330 )),
331 }
332 }
333}
334
335fn connection_error(code: &str, message: String) -> Error {
336 if code == "unauthorized" {
337 Error::Unauthorized
338 } else {
339 Error::Protocol(message)
340 }
341}
342
343impl GatewaySender {
344 pub async fn send(&self, message: ClientMessage) -> Result<()> {
346 if matches!(
347 message,
348 ClientMessage::Pair { .. } | ClientMessage::Authenticate { .. }
349 ) {
350 return Err(Error::Protocol(
351 "authentication messages are valid only during connection setup".into(),
352 ));
353 }
354 self.write(message).await
355 }
356
357 async fn write(&self, message: ClientMessage) -> Result<()> {
358 let mut writer = self.writer.lock().await;
359 write_frame(&mut *writer, &ClientFrame::new(message)).await
360 }
361}
362
363impl GatewayEvents {
364 pub async fn next(&mut self) -> Result<Option<ServerFrame>> {
366 if let Some(frame) = self.pending.pop_front() {
367 return Ok(Some(frame));
368 }
369 let Some(frame) = read_frame::<ServerFrame>(&mut self.reader).await? else {
370 return Ok(None);
371 };
372 validate_version(frame.version)?;
373 Ok(Some(frame))
374 }
375
376 pub fn prepend(&mut self, frames: Vec<ServerFrame>) -> Result<()> {
378 if self.pending.len() + frames.len() > MAX_PENDING_FRAMES {
379 return Err(Error::Protocol(format!(
380 "gateway event backlog exceeds {MAX_PENDING_FRAMES} frames"
381 )));
382 }
383 for frame in &frames {
384 validate_version(frame.version)?;
385 }
386 for frame in frames.into_iter().rev() {
387 self.pending.push_front(frame);
388 }
389 Ok(())
390 }
391}
392
393pub fn token_from_env() -> Result<String> {
395 env::var("MOBIUS_GATEWAY_TOKEN")
396 .ok()
397 .filter(|token| !token.trim().is_empty())
398 .ok_or_else(|| Error::Config("set MOBIUS_GATEWAY_TOKEN before connecting".into()))
399}
400
401fn plaintext_host_is_loopback(host: &str) -> bool {
402 host.eq_ignore_ascii_case("localhost")
403 || host
404 .parse::<IpAddr>()
405 .is_ok_and(|address| address.is_loopback())
406}
407
408fn format_address(host: &str, port: u16) -> String {
409 if host.contains(':') {
410 format!("[{host}]:{port}")
411 } else {
412 format!("{host}:{port}")
413 }
414}
415
416#[cfg(test)]
417mod tests {
418 use super::*;
419
420 #[tokio::test]
421 async fn connect_authenticates_without_a_session_cursor() {
422 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
423 .await
424 .expect("bind gateway");
425 let endpoint = format!("tcp://{}", listener.local_addr().expect("gateway address"))
426 .parse::<Endpoint>()
427 .expect("gateway endpoint");
428 let gateway = tokio::spawn(async move {
429 let (stream, _) = listener.accept().await.expect("accept client");
430 let (reader, mut writer) = tokio::io::split(stream);
431 let mut reader = FrameReader::new(reader);
432 let frame = read_frame::<ClientFrame>(&mut reader)
433 .await
434 .expect("read authentication")
435 .expect("authentication frame");
436 write_frame(&mut writer, &ServerFrame::new(ServerMessage::Authenticated))
437 .await
438 .expect("acknowledge authentication");
439 frame
440 });
441
442 let _client = GatewayClient::connect(&endpoint, "secret", ClientKind::Cli)
443 .await
444 .expect("connect client");
445 let frame = gateway.await.expect("gateway task");
446
447 assert_eq!(
448 frame.message,
449 ClientMessage::Authenticate {
450 token: "secret".into(),
451 client_kind: ClientKind::Cli,
452 }
453 );
454 }
455
456 #[test]
457 fn endpoint_rejects_remote_plaintext() {
458 let error = "tcp://example.com:8741"
459 .parse::<Endpoint>()
460 .expect_err("remote plaintext must fail");
461
462 assert!(error.to_string().contains("use tls://"));
463 assert!("tcp://127.0.0.1:0".parse::<Endpoint>().is_err());
464 }
465
466 #[test]
467 fn endpoint_accepts_loopback_plaintext_and_remote_encrypted_transports() {
468 let loopback = "tcp://127.0.0.1:8741"
469 .parse::<Endpoint>()
470 .expect("loopback endpoint");
471 let remote = "tls://gateway.example:443"
472 .parse::<Endpoint>()
473 .expect("TLS endpoint");
474 let websocket = "wss://gateway.example"
475 .parse::<Endpoint>()
476 .expect("WSS endpoint");
477
478 assert_eq!(loopback.to_string(), "tcp://127.0.0.1:8741");
479 assert_eq!(remote.to_string(), "tls://gateway.example:443");
480 assert_eq!(websocket.to_string(), "wss://gateway.example");
481 assert!(loopback.is_plaintext());
482 assert!(!remote.is_plaintext());
483 assert!(websocket.is_websocket());
484 }
485
486 #[test]
487 fn authentication_errors_preserve_unauthorized_semantics() {
488 assert!(matches!(
489 connection_error("unauthorized", "authentication failed".into()),
490 Error::Unauthorized
491 ));
492 }
493
494 #[tokio::test]
495 async fn prepended_frames_are_returned_in_order() {
496 let (transport, _peer) = tokio::io::duplex(64);
497 let (reader, _writer) = tokio::io::split(Box::new(transport) as BoxedTransport);
498 let mut events = GatewayEvents {
499 reader: FrameReader::new(reader),
500 pending: VecDeque::new(),
501 };
502 events
503 .prepend(vec![
504 ServerFrame::new(ServerMessage::Accepted {
505 request_id: "first".into(),
506 }),
507 ServerFrame::new(ServerMessage::Accepted {
508 request_id: "second".into(),
509 }),
510 ])
511 .expect("defer frames");
512
513 for expected in ["first", "second"] {
514 let frame = events.next().await.expect("next frame").expect("frame");
515 assert!(matches!(
516 frame.message,
517 ServerMessage::Accepted { request_id } if request_id == expected
518 ));
519 }
520 let mut invalid = ServerFrame::new(ServerMessage::Accepted {
521 request_id: "invalid".into(),
522 });
523 invalid.version = 0;
524 assert!(events.prepend(vec![invalid]).is_err());
525 let frame = ServerFrame::new(ServerMessage::Accepted {
526 request_id: "overflow".into(),
527 });
528 assert!(events.prepend(vec![frame; MAX_PENDING_FRAMES + 1]).is_err());
529 }
530}