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