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