1use 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#[derive(Clone)]
34pub struct ClientConfig {
35 pub server: String,
37 pub credential: String,
39 pub keep_alive: bool,
40 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#[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#[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#[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 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 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 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 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 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 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 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 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
339async fn resolve(addr: &str) -> Result<ResolvedAddrs> {
345 resolve_addrs_async(addr)
346 .await
347 .context(AddressSnafu { addr })
348}
349
350fn 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
365fn watch_callback(tx: watch::Sender<TunnelStatus>) -> StatusCallback {
370 Box::new(move |status: &str| {
371 let _ = tx.send(TunnelStatus::from_callback(status));
372 })
373}
374
375struct 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
386struct 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}