Skip to main content

pb_mapper_client/sdk/
client.rs

1//! The session object: [`Client`], its request types, and the tunnel spawner.
2//!
3//! Everything a caller starts goes through here. `register` and `connect` are
4//! the same shape — resolve both endpoints, open a status channel, spawn the
5//! transport's worker — so that lifecycle lives in one place (`TunnelWorker`)
6//! and each call site is left with only the worker it invokes.
7
8use std::sync::{Arc, RwLock};
9
10use pb_mapper_core::checksum::{Credential, parse_credential};
11use pb_mapper_core::config::{ResolvedAddrs, control_io_timeout, resolve_addrs_async};
12use pb_mapper_protocol::command::{PbConnStatusReq, PbConnStatusResp};
13use tokio::sync::watch;
14use tokio_util::sync::CancellationToken;
15use uni_stream::stream::{
16    TcpListenerProvider, TcpStreamProvider, UdpListenerProvider, UdpStreamProvider,
17};
18
19use snafu::ResultExt;
20
21use super::Error;
22use super::admin::Admin;
23use super::error::{AddressSnafu, ConnectSnafu, Result, StatusSnafu};
24use super::handle::{Connection, LiveTunnel, Registration};
25use super::types::{RemoteId, ServiceConnection, Transport, TunnelStatus};
26use crate::client::run_client_side_cli_with_shutdown;
27use crate::client::status::get_status_with_credential;
28use crate::server::{ServerTunnelOptions, StatusCallback, run_server_side_cli_with_shutdown};
29
30/// Configuration for a [`Client`] session.
31#[derive(Clone)]
32pub struct ClientConfig {
33    /// Relay address (`host:port`).
34    pub server: String,
35    /// Administrator key (32 printable bytes) or a `pbmt1_` temporary credential.
36    pub credential: String,
37    pub keep_alive: bool,
38    /// Administrator-only target namespace. Temporary credentials always use
39    /// their own key id when this is `None`.
40    pub namespace: Option<u64>,
41}
42
43impl std::fmt::Debug for ClientConfig {
44    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
45        f.debug_struct("ClientConfig")
46            .field("server", &self.server)
47            .field("credential", &"[redacted]")
48            .field("keep_alive", &self.keep_alive)
49            .field("namespace", &self.namespace)
50            .finish()
51    }
52}
53
54/// Register a local TCP/UDP service with the relay.
55#[derive(Clone, Debug)]
56pub struct RegisterRequest {
57    pub key: String,
58    pub local_addr: String,
59    pub transport: Transport,
60    pub codec: bool,
61    pub force_namespace: bool,
62}
63
64/// Subscribe to a registered service and listen locally.
65#[derive(Clone, Debug)]
66pub struct ConnectRequest {
67    pub key: String,
68    pub local_addr: String,
69    pub transport: Transport,
70}
71
72pub(crate) struct ClientInner {
73    pub(crate) server: String,
74    pub(crate) credential: RwLock<Credential>,
75    pub(crate) keep_alive: bool,
76    pub(crate) namespace: Option<u64>,
77}
78
79/// Session against one deployed relay.
80///
81/// The credential is per-client, not process-global. Two clients in the same
82/// process may use different keys.
83#[derive(Clone)]
84pub struct Client {
85    pub(crate) inner: Arc<ClientInner>,
86}
87
88impl Client {
89    pub fn new(config: ClientConfig) -> Result<Self> {
90        if config.server.trim().is_empty() {
91            return Err(Error::invalid_config("server address is required"));
92        }
93        if config.credential.trim().is_empty() {
94            return Err(Error::invalid_config("credential is required"));
95        }
96        let credential =
97            parse_credential(config.credential.trim()).map_err(Error::invalid_config)?;
98        Ok(Self::from_credential(
99            config.server,
100            credential,
101            config.keep_alive,
102            config.namespace,
103        ))
104    }
105
106    /// Build a client from an already-parsed credential.
107    pub fn from_credential(
108        server: impl Into<String>,
109        credential: Credential,
110        keep_alive: bool,
111        namespace: Option<u64>,
112    ) -> Self {
113        Self {
114            inner: Arc::new(ClientInner {
115                server: server.into(),
116                credential: RwLock::new(credential),
117                keep_alive,
118                namespace,
119            }),
120        }
121    }
122
123    pub fn server(&self) -> &str {
124        &self.inner.server
125    }
126
127    pub fn namespace(&self) -> Option<u64> {
128        self.inner.namespace
129    }
130
131    pub(crate) fn credential(&self) -> Credential {
132        *self
133            .inner
134            .credential
135            .read()
136            .unwrap_or_else(|poisoned| poisoned.into_inner())
137    }
138
139    /// Administrator RPCs. Fails locally when the session is not an admin key.
140    pub fn admin(&self) -> Result<Admin> {
141        if !self.credential().is_admin() {
142            return Err(Error::NotAdministrator);
143        }
144        Ok(Admin {
145            inner: Arc::clone(&self.inner),
146        })
147    }
148
149    /// Publish a local service on the relay under `request.key`.
150    ///
151    /// Returns as soon as the worker is spawned; the handle's
152    /// [`Registration::wait_ready`] is what waits for the relay to accept it.
153    pub async fn register(&self, request: RegisterRequest) -> Result<Registration> {
154        if request.key.trim().is_empty() {
155            return Err(Error::invalid_config("service key is required"));
156        }
157        let options = ServerTunnelOptions {
158            need_codec: request.codec,
159            is_datagram: request.transport.is_datagram(),
160            keep_alive: self.inner.keep_alive,
161            namespace: self.inner.namespace,
162            force_namespace: request.force_namespace,
163        };
164        let worker = self
165            .prepare_worker(&request.key, &request.local_addr)
166            .await?;
167        let credential = self.credential();
168        let handle = match request.transport {
169            Transport::Tcp => worker.spawn(move |context| {
170                run_server_side_cli_with_shutdown::<TcpStreamProvider>(
171                    context.local_addr,
172                    context.remote_addr,
173                    context.key,
174                    options,
175                    Some(context.status_callback),
176                    credential,
177                    context.shutdown,
178                )
179            }),
180            Transport::Udp => worker.spawn(move |context| {
181                run_server_side_cli_with_shutdown::<UdpStreamProvider>(
182                    context.local_addr,
183                    context.remote_addr,
184                    context.key,
185                    options,
186                    Some(context.status_callback),
187                    credential,
188                    context.shutdown,
189                )
190            }),
191        };
192        Ok(Registration::new(handle, request.key))
193    }
194
195    /// Subscribe to a registered service and forward it from a local listener.
196    ///
197    /// Returns as soon as the worker is spawned; the handle's
198    /// [`Connection::wait_ready`] is what waits for the local listener to bind
199    /// and the relay to confirm the service.
200    pub async fn connect(&self, request: ConnectRequest) -> Result<Connection> {
201        if request.key.trim().is_empty() {
202            return Err(Error::invalid_config("service key is required"));
203        }
204        let worker = self
205            .prepare_worker(&request.key, &request.local_addr)
206            .await?;
207        let credential = self.credential();
208        let keep_alive = self.inner.keep_alive;
209        let namespace = self.inner.namespace;
210        let handle = match request.transport {
211            Transport::Tcp => worker.spawn(move |context| {
212                run_client_side_cli_with_shutdown::<TcpListenerProvider>(
213                    context.local_addr,
214                    context.remote_addr,
215                    context.key,
216                    keep_alive,
217                    namespace,
218                    Some(context.status_callback),
219                    Some(credential),
220                    context.shutdown,
221                )
222            }),
223            Transport::Udp => worker.spawn(move |context| {
224                run_client_side_cli_with_shutdown::<UdpListenerProvider>(
225                    context.local_addr,
226                    context.remote_addr,
227                    context.key,
228                    keep_alive,
229                    namespace,
230                    Some(context.status_callback),
231                    Some(credential),
232                    context.shutdown,
233                )
234            }),
235        };
236        Ok(Connection::new(handle, request.key))
237    }
238
239    /// Service keys visible to this credential's namespace.
240    pub async fn list_keys(&self) -> Result<Vec<String>> {
241        match self.status_request(PbConnStatusReq::Keys).await? {
242            PbConnStatusResp::Keys(keys) => Ok(keys),
243            other => Err(Error::protocol(format!(
244                "expected keys status, got {other:?}"
245            ))),
246        }
247    }
248
249    pub async fn service_status(&self, key: impl Into<String>) -> Result<Vec<ServiceConnection>> {
250        let key = key.into();
251        match self
252            .status_request(PbConnStatusReq::Service { key })
253            .await?
254        {
255            PbConnStatusResp::Service { connections, .. } => Ok(connections
256                .into_iter()
257                .map(ServiceConnection::from)
258                .collect()),
259            other => Err(Error::protocol(format!(
260                "expected service status, got {other:?}"
261            ))),
262        }
263    }
264
265    pub async fn remote_id(&self) -> Result<RemoteId> {
266        RemoteId::from_status(self.status_request(PbConnStatusReq::RemoteId).await?)
267    }
268
269    /// Resolve both ends of a tunnel and set up its status channel.
270    ///
271    /// Shared by `register` and `connect`: the two differ only in which worker
272    /// they hand the resolved context to.
273    async fn prepare_worker(&self, key: &str, local_addr: &str) -> Result<TunnelWorker> {
274        let local_addr = resolve(local_addr).await?;
275        let remote_addr = resolve(&self.inner.server).await?;
276        let (status_tx, status_rx) = watch::channel(TunnelStatus::Starting);
277        Ok(TunnelWorker {
278            local_addr,
279            remote_addr,
280            key: Arc::from(key),
281            shutdown: CancellationToken::new(),
282            status_tx,
283            status_rx,
284        })
285    }
286
287    async fn status_request(&self, request: PbConnStatusReq) -> Result<PbConnStatusResp> {
288        let addrs = resolve(&self.inner.server).await?;
289        let credential = self.credential();
290        // The connect is inside the timeout, not just the exchange that follows it.
291        // A relay that drops SYNs silently leaves `TcpStream::connect` waiting on
292        // the OS timeout — minutes — so the SDK's own bound has to cover it, the
293        // way the administrator path already does.
294        let io_timeout = control_io_timeout();
295        // Every candidate, under one shared bound: `each_addr` moves on to the next
296        // address when one refuses, and the timeout covers the whole sequence so a
297        // list of blackholed addresses cannot multiply the wait by its length.
298        let connect = crate::addr::connect_tcp(&addrs);
299        let mut stream = match tokio::time::timeout(io_timeout, connect).await {
300            Ok(result) => result.context(ConnectSnafu {
301                addr: addrs.to_string(),
302            })?,
303            Err(_) => {
304                return Err(Error::TimedOut {
305                    timeout: io_timeout,
306                });
307            }
308        };
309        get_status_with_credential(&mut stream, request, self.inner.namespace, &credential)
310            .await
311            .context(StatusSnafu)
312    }
313}
314
315/// Resolve one endpoint to every address it names.
316///
317/// The list, not just its first entry: a worker handed the whole list dials each
318/// candidate in turn, so a relay hostname with several records survives one of
319/// them being unreachable.
320async fn resolve(addr: &str) -> Result<ResolvedAddrs> {
321    resolve_addrs_async(addr)
322        .await
323        .context(AddressSnafu { addr })
324}
325
326/// Mark a finished worker as `Stopped`, unless it already reported why it will
327/// never come up. The status is a watch channel, so it keeps only the newest
328/// value: overwriting a `Failed(reason)` that the worker set on its way out
329/// would replace the only description of a permanent rejection the caller ever
330/// gets, leaving `wait_ready` to report a bare stop instead.
331fn settle_stopped(tx: &watch::Sender<TunnelStatus>) {
332    tx.send_if_modified(|status| match status {
333        TunnelStatus::Failed(_) => false,
334        _ => {
335            *status = TunnelStatus::Stopped;
336            true
337        }
338    });
339}
340
341/// Bridge the worker's string status callback onto the handle's watch channel.
342///
343/// `StatusCallback` and `ClientStatusCallback` are the same boxed closure type,
344/// so both tunnel directions share this.
345fn watch_callback(tx: watch::Sender<TunnelStatus>) -> StatusCallback {
346    Box::new(move |status: &str| {
347        let _ = tx.send(TunnelStatus::from_callback(status));
348    })
349}
350
351/// What a spawned tunnel worker is given: its resolved endpoints, its key, the
352/// callback that publishes its status, and the token that stops it.
353struct WorkerContext {
354    local_addr: ResolvedAddrs,
355    remote_addr: ResolvedAddrs,
356    key: Arc<str>,
357    status_callback: StatusCallback,
358    shutdown: CancellationToken,
359}
360
361/// A tunnel resolved and wired up, waiting only for the transport-specific
362/// worker that will drive it.
363///
364/// Register and connect, over TCP and UDP, are four workers with four different
365/// signatures but one lifecycle: spawn, publish status, settle as `Stopped`.
366/// This owns that lifecycle so each call site is left with just its own call.
367struct TunnelWorker {
368    local_addr: ResolvedAddrs,
369    remote_addr: ResolvedAddrs,
370    key: Arc<str>,
371    shutdown: CancellationToken,
372    status_tx: watch::Sender<TunnelStatus>,
373    status_rx: watch::Receiver<TunnelStatus>,
374}
375
376impl TunnelWorker {
377    fn spawn<F, Fut>(self, start: F) -> LiveTunnel
378    where
379        F: FnOnce(WorkerContext) -> Fut + Send + 'static,
380        Fut: std::future::Future<Output = ()> + Send,
381    {
382        let Self {
383            local_addr,
384            remote_addr,
385            key,
386            shutdown,
387            status_tx,
388            status_rx,
389        } = self;
390        let worker_shutdown = shutdown.clone();
391        let join = tokio::spawn(async move {
392            start(WorkerContext {
393                local_addr,
394                remote_addr,
395                key,
396                status_callback: watch_callback(status_tx.clone()),
397                shutdown: worker_shutdown,
398            })
399            .await;
400            settle_stopped(&status_tx);
401        });
402        LiveTunnel::new(shutdown, join, status_rx)
403    }
404}
405
406#[cfg(test)]
407mod tests {
408    use super::*;
409
410    #[test]
411    fn config_debug_redacts_credentials() {
412        let config = ClientConfig {
413            server: "localhost:7666".into(),
414            credential: "0123456789abcdefghijklmnopqrstuv".into(),
415            keep_alive: false,
416            namespace: None,
417        };
418        let debug = format!("{config:?}");
419        assert!(!debug.contains(&config.credential));
420        assert!(debug.contains("[redacted]"));
421    }
422
423    #[test]
424    fn admin_requires_administrator_credential() {
425        let admin = Client::from_credential(
426            "127.0.0.1:7666",
427            Credential::Admin(*b"0123456789abcdefghijklmnopqrstuv"),
428            false,
429            None,
430        );
431        assert!(admin.admin().is_ok());
432
433        let temporary = Client::from_credential(
434            "127.0.0.1:7666",
435            Credential::Temporary {
436                key_id: 1,
437                key: [0_u8; 32],
438            },
439            false,
440            None,
441        );
442        assert!(matches!(temporary.admin(), Err(Error::NotAdministrator)));
443    }
444}