1use 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#[derive(Clone)]
32pub struct ClientConfig {
33 pub server: String,
35 pub credential: String,
37 pub keep_alive: bool,
38 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#[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#[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#[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 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 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 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 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 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 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 let io_timeout = control_io_timeout();
295 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
315async fn resolve(addr: &str) -> Result<ResolvedAddrs> {
321 resolve_addrs_async(addr)
322 .await
323 .context(AddressSnafu { addr })
324}
325
326fn 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
341fn watch_callback(tx: watch::Sender<TunnelStatus>) -> StatusCallback {
346 Box::new(move |status: &str| {
347 let _ = tx.send(TunnelStatus::from_callback(status));
348 })
349}
350
351struct WorkerContext {
354 local_addr: ResolvedAddrs,
355 remote_addr: ResolvedAddrs,
356 key: Arc<str>,
357 status_callback: StatusCallback,
358 shutdown: CancellationToken,
359}
360
361struct 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}