1use std::sync::Arc;
10
11use serde::{Deserialize, Serialize};
12use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt as _};
13use tokio::net::TcpStream;
14use tokio::net::UnixStream;
15use tokio_rustls::TlsConnector;
16use tokio_rustls::rustls::pki_types::pem::PemObject;
17use tokio_rustls::rustls::pki_types::{CertificateDer, ServerName};
18use tokio_rustls::rustls::{ClientConfig, RootCertStore};
19
20use crate::caliband::wire::Endpoint;
21
22pub trait Conn: AsyncRead + AsyncWrite + Unpin + Send {}
24impl<T: AsyncRead + AsyncWrite + Unpin + Send> Conn for T {}
25
26pub type BoxConn = Box<dyn Conn>;
28
29#[derive(Clone)]
31pub struct TlsClient {
32 pub connector: TlsConnector,
34 pub server_name: String,
36}
37
38fn ensure_crypto_provider() {
40 use std::sync::Once;
41 static INIT: Once = Once::new();
42 INIT.call_once(|| {
43 let _ = tokio_rustls::rustls::crypto::ring::default_provider().install_default();
44 });
45}
46
47pub fn tls_client_from_pem(ca_pem: &[u8], server_name: &str) -> std::io::Result<TlsClient> {
49 ensure_crypto_provider();
50 let mut roots = RootCertStore::empty();
51 for cert in CertificateDer::pem_slice_iter(ca_pem) {
52 roots
53 .add(cert.map_err(|e| std::io::Error::other(e.to_string()))?)
54 .map_err(std::io::Error::other)?;
55 }
56 if roots.is_empty() {
61 return Err(std::io::Error::new(
62 std::io::ErrorKind::InvalidData,
63 "no certificates found in CA PEM",
64 ));
65 }
66 let config = ClientConfig::builder()
67 .with_root_certificates(roots)
68 .with_no_client_auth();
69 Ok(TlsClient {
70 connector: TlsConnector::from(Arc::new(config)),
71 server_name: server_name.to_string(),
72 })
73}
74
75#[derive(Serialize, Deserialize)]
80struct TokenPreamble {
81 bearer: String,
82}
83
84async fn client_send_token(conn: &mut BoxConn, token: &str) -> std::io::Result<()> {
85 let mut line = serde_json::to_vec(&TokenPreamble {
86 bearer: token.to_string(),
87 })
88 .map_err(std::io::Error::other)?;
89 line.push(b'\n');
90 conn.write_all(&line).await?;
91 conn.flush().await
92}
93
94pub struct ConnectSpec {
96 pub endpoint: Endpoint,
98 pub tls: Option<TlsClient>,
100 pub token: Option<String>,
102}
103
104pub async fn connect(spec: &ConnectSpec) -> std::io::Result<BoxConn> {
107 match &spec.endpoint {
108 Endpoint::Unix { path } => Ok(Box::new(UnixStream::connect(path).await?)),
109 Endpoint::Tcp { addr } => {
110 let stream = TcpStream::connect(addr).await?;
111 let mut conn: BoxConn = match &spec.tls {
112 None => Box::new(stream),
113 Some(t) => {
114 let name = ServerName::try_from(t.server_name.clone())
115 .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidInput, e))?;
116 Box::new(t.connector.connect(name, stream).await?)
117 }
118 };
119 if let Some(token) = &spec.token {
120 client_send_token(&mut conn, token).await?;
121 }
122 Ok(conn)
123 }
124 }
125}
126
127#[cfg(any(test, feature = "testkit"))]
129mod server {
130 use super::{BoxConn, Endpoint, TokenPreamble, ensure_crypto_provider};
131 use std::sync::Arc;
132 use tokio::io::AsyncReadExt as _;
133 use tokio::net::{TcpListener, UnixListener};
134 use tokio_rustls::TlsAcceptor;
135 use tokio_rustls::rustls::ServerConfig;
136 use tokio_rustls::rustls::pki_types::pem::PemObject;
137 use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer};
138
139 #[derive(Clone)]
141 pub struct TlsServer {
142 pub acceptor: TlsAcceptor,
144 }
145
146 pub fn tls_server_from_pem(cert_pem: &[u8], key_pem: &[u8]) -> std::io::Result<TlsServer> {
148 ensure_crypto_provider();
149 let certs: Vec<CertificateDer<'static>> = CertificateDer::pem_slice_iter(cert_pem)
150 .collect::<Result<_, _>>()
151 .map_err(|e| std::io::Error::other(e.to_string()))?;
152 let key: PrivateKeyDer<'static> = PrivateKeyDer::from_pem_slice(key_pem)
153 .map_err(|e| std::io::Error::other(e.to_string()))?;
154 let config = ServerConfig::builder()
155 .with_no_client_auth()
156 .with_single_cert(certs, key)
157 .map_err(std::io::Error::other)?;
158 Ok(TlsServer {
159 acceptor: TlsAcceptor::from(Arc::new(config)),
160 })
161 }
162
163 async fn read_preamble_line(conn: &mut BoxConn) -> std::io::Result<String> {
164 let mut buf = Vec::with_capacity(128);
165 let mut byte = [0u8; 1];
166 loop {
167 let n = conn.read(&mut byte).await?;
168 if n == 0 {
169 return Err(std::io::Error::new(
170 std::io::ErrorKind::UnexpectedEof,
171 "no token preamble",
172 ));
173 }
174 if byte[0] == b'\n' {
175 break;
176 }
177 buf.push(byte[0]);
178 if buf.len() > 4096 {
179 return Err(std::io::Error::new(
180 std::io::ErrorKind::InvalidData,
181 "token preamble too long",
182 ));
183 }
184 }
185 String::from_utf8(buf).map_err(std::io::Error::other)
186 }
187
188 fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
196 if a.len() != b.len() {
197 return false;
198 }
199 let mut diff: u8 = 0;
200 for (x, y) in a.iter().zip(b.iter()) {
201 diff |= x ^ y;
202 }
203 diff == 0
204 }
205
206 async fn server_check_token(conn: &mut BoxConn, expected: &str) -> std::io::Result<()> {
207 let line = read_preamble_line(conn).await?;
208 let preamble: TokenPreamble = serde_json::from_str(&line).map_err(std::io::Error::other)?;
209 if constant_time_eq(preamble.bearer.as_bytes(), expected.as_bytes()) {
210 Ok(())
211 } else {
212 Err(std::io::Error::new(
213 std::io::ErrorKind::PermissionDenied,
214 "bad bearer token",
215 ))
216 }
217 }
218
219 pub struct BindSpec {
221 pub endpoint: Endpoint,
223 pub tls: Option<TlsServer>,
225 pub token: Option<String>,
227 }
228
229 pub enum Listener {
231 Unix(UnixListener),
233 Tcp {
235 listener: TcpListener,
237 tls: Option<TlsServer>,
239 token: Option<String>,
241 },
242 }
243
244 impl Listener {
245 pub async fn bind(spec: &BindSpec) -> std::io::Result<Listener> {
247 match &spec.endpoint {
248 Endpoint::Unix { path } => {
249 if let Some(parent) = path.parent() {
250 tokio::fs::create_dir_all(parent).await?;
251 }
252 let _ = tokio::fs::remove_file(path).await;
253 Ok(Listener::Unix(UnixListener::bind(path)?))
254 }
255 Endpoint::Tcp { addr } => Ok(Listener::Tcp {
256 listener: TcpListener::bind(addr).await?,
257 tls: spec.tls.clone(),
258 token: spec.token.clone(),
259 }),
260 }
261 }
262
263 pub fn local_addr(&self) -> Option<String> {
265 match self {
266 Listener::Unix(_) => None,
267 Listener::Tcp { listener, .. } => listener.local_addr().ok().map(|a| a.to_string()),
268 }
269 }
270
271 pub async fn accept(&self) -> std::io::Result<BoxConn> {
274 match self {
275 Listener::Unix(l) => {
276 let (stream, _addr) = l.accept().await?;
277 Ok(Box::new(stream))
278 }
279 Listener::Tcp {
280 listener,
281 tls,
282 token,
283 } => {
284 let (stream, _addr) = listener.accept().await?;
285 let mut conn: BoxConn = match tls {
286 None => Box::new(stream),
287 Some(t) => Box::new(t.acceptor.accept(stream).await?),
288 };
289 if let Some(expected) = token {
290 server_check_token(&mut conn, expected).await?;
291 }
292 Ok(conn)
293 }
294 }
295 }
296 }
297}
298
299#[cfg(any(test, feature = "testkit"))]
300pub use server::{BindSpec, Listener, TlsServer, tls_server_from_pem};
301
302#[cfg(test)]
303mod tests {
304 use super::*;
305 use tokio::io::AsyncReadExt as _;
306
307 async fn echo_once(listener: Listener) {
308 let mut c = listener.accept().await.expect("accept");
309 let mut buf = [0u8; 5];
310 c.read_exact(&mut buf).await.expect("read");
311 c.write_all(&buf).await.expect("write");
312 c.flush().await.expect("flush");
313 }
314
315 #[tokio::test]
316 async fn tcp_tls_token_round_trip() {
317 let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
318 let cert_pem = cert.cert.pem().into_bytes();
319 let key_pem = cert.key_pair.serialize_pem().into_bytes();
320
321 let listener = Listener::bind(&BindSpec {
322 endpoint: Endpoint::Tcp {
323 addr: "127.0.0.1:0".into(),
324 },
325 tls: Some(tls_server_from_pem(&cert_pem, &key_pem).unwrap()),
326 token: Some("s3cr3t".into()),
327 })
328 .await
329 .unwrap();
330 let addr = listener.local_addr().unwrap();
331 let server = tokio::spawn(echo_once(listener));
332
333 let mut c = connect(&ConnectSpec {
334 endpoint: Endpoint::Tcp { addr },
335 tls: Some(tls_client_from_pem(&cert_pem, "localhost").unwrap()),
336 token: Some("s3cr3t".into()),
337 })
338 .await
339 .unwrap();
340 c.write_all(b"hello").await.unwrap();
341 let mut got = [0u8; 5];
342 c.read_exact(&mut got).await.unwrap();
343 assert_eq!(&got, b"hello");
344 server.await.unwrap();
345 }
346
347 #[tokio::test]
348 async fn bad_token_is_rejected() {
349 let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
350 let cert_pem = cert.cert.pem().into_bytes();
351 let key_pem = cert.key_pair.serialize_pem().into_bytes();
352 let listener = Listener::bind(&BindSpec {
353 endpoint: Endpoint::Tcp {
354 addr: "127.0.0.1:0".into(),
355 },
356 tls: Some(tls_server_from_pem(&cert_pem, &key_pem).unwrap()),
357 token: Some("right".into()),
358 })
359 .await
360 .unwrap();
361 let addr = listener.local_addr().unwrap();
362 tokio::spawn(async move {
363 let _ = listener.accept().await;
364 });
365 let r = connect(&ConnectSpec {
366 endpoint: Endpoint::Tcp { addr },
367 tls: Some(tls_client_from_pem(&cert_pem, "localhost").unwrap()),
368 token: Some("wrong".into()),
369 })
370 .await;
371 if let Ok(mut c) = r {
374 let mut b = [0u8; 1];
375 assert!(c.read(&mut b).await.map(|n| n == 0).unwrap_or(true));
376 }
377 }
378
379 #[tokio::test]
380 async fn unix_round_trip() {
381 let dir = tempfile::tempdir().unwrap();
382 let path = dir.path().join("t.sock");
383 let listener = Listener::bind(&BindSpec {
384 endpoint: Endpoint::Unix { path: path.clone() },
385 tls: None,
386 token: None,
387 })
388 .await
389 .unwrap();
390 let server = tokio::spawn(echo_once(listener));
391 let mut c = connect(&ConnectSpec {
392 endpoint: Endpoint::Unix { path },
393 tls: None,
394 token: None,
395 })
396 .await
397 .unwrap();
398 c.write_all(b"world").await.unwrap();
399 let mut got = [0u8; 5];
400 c.read_exact(&mut got).await.unwrap();
401 assert_eq!(&got, b"world");
402 server.await.unwrap();
403 }
404}