Skip to main content

mobius_gateway/
client.rs

1//! Reusable async client for CLI and native frontends.
2
3use 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;
31/// Maximum number of frames a focused client flow may temporarily defer.
32pub 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/// Validated plaintext-loopback, authenticated-root TLS, or WSS endpoint.
41#[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/// Token returned while pairing a new client.
56#[derive(Debug, Clone, PartialEq, Eq)]
57pub struct PairedClient {
58    pub client_id: String,
59    pub token: String,
60}
61
62/// Connected client before its command and event halves are separated.
63pub struct GatewayClient {
64    sender: GatewaySender,
65    events: GatewayEvents,
66}
67
68/// Cloneable framed command writer.
69#[derive(Clone)]
70pub struct GatewaySender {
71    writer: Arc<Mutex<WriteHalf<BoxedTransport>>>,
72}
73
74/// Single-owner framed event reader.
75pub struct GatewayEvents {
76    reader: FrameReader<ReadHalf<BoxedTransport>>,
77    pending: VecDeque<ServerFrame>,
78}
79
80impl Endpoint {
81    /// Resolves `MOBIUS_GATEWAY_ENDPOINT`, defaulting to local plaintext.
82    pub fn from_env() -> Result<Self> {
83        env::var("MOBIUS_GATEWAY_ENDPOINT")
84            .unwrap_or_else(|_| DEFAULT_ENDPOINT.into())
85            .parse()
86    }
87
88    /// Returns whether this endpoint uses loopback-only plaintext transport.
89    #[must_use]
90    pub const fn is_plaintext(&self) -> bool {
91        matches!(self.security, Security::Plaintext)
92    }
93
94    /// Returns whether this endpoint uses secure WebSocket transport.
95    #[must_use]
96    pub const fn is_websocket(&self) -> bool {
97        matches!(self.security, Security::WebSocketTls)
98    }
99
100    /// Returns the validated endpoint host for server-side routing policy.
101    #[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    /// Authenticates an existing client and leaves the gateway Ready frame for `events`.
243    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    /// Consumes a pending pairing code and returns a connected independent client.
262    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    /// Separates the clonable command writer from the single event reader.
300    #[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    /// Sends one authenticated operation.
343    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    /// Receives the next version-checked server frame.
363    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    /// Restores temporarily consumed frames ahead of unread transport data.
375    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
391/// Resolves the bearer token expected by the reusable CLI client.
392pub 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}