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::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#[derive(Clone, Debug)]
34pub struct ClientConfig {
35 pub server: String,
37 pub credential: String,
39 pub keep_alive: bool,
40 pub namespace: Option<u64>,
43}
44
45#[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#[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#[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 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 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 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 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 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 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 let io_timeout = control_io_timeout();
286 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
306async fn resolve(addr: &str) -> Result<ResolvedAddrs> {
312 resolve_addrs_async(addr)
313 .await
314 .context(AddressSnafu { addr })
315}
316
317fn 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
332fn watch_callback(tx: watch::Sender<TunnelStatus>) -> StatusCallback {
337 Box::new(move |status: &str| {
338 let _ = tx.send(TunnelStatus::from_callback(status));
339 })
340}
341
342struct WorkerContext {
345 local_addr: ResolvedAddrs,
346 remote_addr: ResolvedAddrs,
347 key: Arc<str>,
348 status_callback: StatusCallback,
349 shutdown: CancellationToken,
350}
351
352struct 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}