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| {
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    /// Authenticates an existing client and leaves the gateway Ready frame for `events`.
245    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    /// Consumes a pending pairing code and returns a connected independent client.
264    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    /// Separates the clonable command writer from the single event reader.
302    #[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    /// Sends one authenticated operation.
345    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    /// Receives the next version-checked server frame.
365    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    /// Restores temporarily consumed frames ahead of unread transport data.
377    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
393/// Resolves the bearer token expected by the reusable CLI client.
394pub 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}