1use std::sync::Arc;
21
22use parking_lot::Mutex;
23use std::time::{Duration, SystemTime, UNIX_EPOCH};
24
25use rand::RngExt;
26use ring::aead::{AES_256_GCM, Aad, LessSafeKey, Nonce, UnboundKey};
27use ring::digest::{SHA256, digest};
28use ring::hkdf::{HKDF_SHA256, Salt};
29use tokio::io::{AsyncReadExt, AsyncWriteExt};
30
31use super::{
32 CodecMessageReader, CodecMessageWriter, DataLenType, MAX_MSG_LEN, MessageReader, MessageWriter,
33};
34use pb_mapper_auth::{
35 ADMIN_KEY_ID, AuthContext, AuthFailure, AuthRuntime, KeyId, LegacyConnectionGuard,
36};
37use pb_mapper_core::checksum::{
38 AesKeyType, Credential, get_process_credential, valid_checksum_for_key,
39};
40use pb_mapper_core::codec::{Aes256GcmDeCodec, Aes256GcmEnCodec, Decryptor};
41use pb_mapper_core::error::{Error, Result};
42
43pub const PROTOCOL_V2_MAGIC: [u8; 4] = *b"PBM2";
44pub const PROTOCOL_V2_VERSION: u8 = 2;
45const CONNECTION_SALT_LEN: usize = 16;
46const FIRST_PREFIX_REMAINDER_LEN: usize = 28;
47const FRAME_HEADER_LEN: usize = 12;
48const DIRECTION_CLIENT_TO_SERVER: u8 = 0;
49const DIRECTION_SERVER_TO_CLIENT: u8 = 1;
50const MAX_CONNECTION_CLOCK_SKEW_SECONDS: u64 = 5 * 60;
51const DEFAULT_REPLAY_WINDOW_SECONDS: u64 = MAX_CONNECTION_CLOCK_SKEW_SECONDS.saturating_mul(2);
55const DEFAULT_REPLAY_FILTER_BYTES: usize = 1024 * 1024;
56const MAX_INITIAL_PLAINTEXT_LEN: u32 = 64 * 1024;
57const MAX_INITIAL_CIPHERTEXT_LEN: u32 = MAX_INITIAL_PLAINTEXT_LEN + 16;
58
59#[derive(Clone, Copy, Debug, Eq, PartialEq)]
60pub enum HeaderProtocol {
61 Legacy,
62 V2,
63}
64
65pub struct ClientHeaderSession {
66 protocol: HeaderProtocol,
67 legacy_key: AesKeyType,
68 v2: Option<V2Material>,
69}
70
71impl ClientHeaderSession {
72 fn v2_material(&self) -> Result<&V2Material> {
77 self.v2
78 .as_ref()
79 .ok_or_else(|| protocol_error("v2 session is missing its key material"))
80 }
81
82 pub fn from_process() -> Result<Self> {
84 let credential = get_process_credential().map_err(protocol_error)?;
85 Self::new_v2(&credential)
86 }
87
88 pub fn new_v2(credential: &Credential) -> Result<Self> {
89 let mut salt = [0_u8; CONNECTION_SALT_LEN];
90 salt[..8].copy_from_slice(&unix_seconds().to_be_bytes());
91 let mut rng = rand::rng();
92 for byte in &mut salt[8..] {
93 *byte = rng.random();
94 }
95 let material =
96 derive_material(KeyId::from_u64(credential.key_id()), credential.key(), salt)?;
97 Ok(Self {
98 protocol: HeaderProtocol::V2,
99 legacy_key: *credential.key(),
100 v2: Some(material),
101 })
102 }
103
104 #[cfg(test)]
105 pub fn new_legacy(key: AesKeyType) -> Self {
106 Self {
107 protocol: HeaderProtocol::Legacy,
108 legacy_key: key,
109 v2: None,
110 }
111 }
112
113 pub fn protocol(&self) -> HeaderProtocol {
114 self.protocol
115 }
116
117 pub async fn write_initial<T: AsyncWriteExt + Unpin>(
118 &self,
119 writer: &mut T,
120 message: &[u8],
121 ) -> Result<()> {
122 match self.protocol {
123 HeaderProtocol::Legacy => {
124 legacy_message_writer(writer, &self.legacy_key, "legacy writer")?
125 .write_msg(message)
126 .await
127 }
128 HeaderProtocol::V2 => {
129 let material = self.v2_material()?;
130 writer
131 .write_all(&first_prefix(material))
132 .await
133 .map_err(|error| {
134 protocol_error(format!("failed to write v2 prefix: {error}"))
135 })?;
136 V2MessageWriter::new(writer, material.clone(), DIRECTION_CLIENT_TO_SERVER, 0)?
137 .write_msg(message)
138 .await
139 }
140 }
141 }
142
143 pub fn response_reader<'a, T: AsyncReadExt + Unpin>(
144 &self,
145 reader: &'a mut T,
146 ) -> Result<HeaderMessageReader<'a, T>> {
147 match self.protocol {
148 HeaderProtocol::Legacy => Ok(HeaderMessageReader::Legacy(legacy_message_reader(
149 reader,
150 &self.legacy_key,
151 "legacy reader",
152 )?)),
153 HeaderProtocol::V2 => Ok(HeaderMessageReader::V2(V2MessageReader::new(
154 reader,
155 self.v2_material()?.clone(),
156 DIRECTION_SERVER_TO_CLIENT,
157 0,
158 )?)),
159 }
160 }
161
162 pub async fn exchange<T: AsyncReadExt + AsyncWriteExt + Unpin>(
163 &self,
164 stream: &mut T,
165 payload: &[u8],
166 timeout: Duration,
167 ) -> Result<Vec<u8>> {
168 match tokio::time::timeout(timeout, self.write_initial(stream, payload)).await {
169 Ok(result) => result?,
170 Err(_) => {
171 return Err(protocol_error(format!(
172 "timed out writing first-flight request after {timeout:?}"
173 )));
174 }
175 }
176 let mut reader = self.response_reader(stream)?;
177 let message = match tokio::time::timeout(timeout, reader.read_msg()).await {
178 Ok(result) => result?,
179 Err(_) => {
180 return Err(protocol_error(format!(
181 "timed out reading first-flight response after {timeout:?}"
182 )));
183 }
184 };
185 Ok(message.to_vec())
186 }
187
188 pub fn continuation_writer<'a, T: AsyncWriteExt + Unpin>(
189 &self,
190 writer: &'a mut T,
191 ) -> Result<HeaderMessageWriter<'a, T>> {
192 match self.protocol {
193 HeaderProtocol::Legacy => Ok(HeaderMessageWriter::Legacy(legacy_message_writer(
194 writer,
195 &self.legacy_key,
196 "legacy writer",
197 )?)),
198 HeaderProtocol::V2 => Ok(HeaderMessageWriter::V2(V2MessageWriter::new(
199 writer,
200 self.v2_material()?.clone(),
201 DIRECTION_CLIENT_TO_SERVER,
202 1,
203 )?)),
204 }
205 }
206}
207
208pub struct ServerHeaderSession {
209 protocol: HeaderProtocol,
210 legacy_key: AesKeyType,
211 v2: Option<V2Material>,
212 context: Option<AuthContext>,
213 _legacy_guard: Option<LegacyConnectionGuard>,
214}
215
216impl fmt::Debug for ServerHeaderSession {
217 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
218 formatter
219 .debug_struct("ServerHeaderSession")
220 .field("protocol", &self.protocol)
221 .field("key_id", &self.key_id())
222 .field("authenticated", &self.context.is_some())
223 .finish()
224 }
225}
226
227impl ServerHeaderSession {
228 fn v2_material(&self) -> Result<&V2Material> {
231 self.v2
232 .as_ref()
233 .ok_or_else(|| protocol_error("v2 session is missing its key material"))
234 }
235
236 pub fn protocol(&self) -> HeaderProtocol {
237 self.protocol
238 }
239
240 pub fn framing_key(&self) -> AesKeyType {
241 self.legacy_key
242 }
243
244 pub fn key_id(&self) -> KeyId {
245 self.context
246 .as_ref()
247 .map(|context| context.key_id)
248 .unwrap_or_else(|| {
249 self.v2
250 .as_ref()
251 .map(|material| material.key_id)
252 .unwrap_or(ADMIN_KEY_ID)
253 })
254 }
255
256 pub fn context(&self) -> Result<&AuthContext> {
257 self.context
258 .as_ref()
259 .ok_or_else(|| protocol_error("server session was not authenticated"))
260 }
261
262 pub fn response_writer<'a, T: AsyncWriteExt + Unpin>(
263 &self,
264 writer: &'a mut T,
265 ) -> Result<HeaderMessageWriter<'a, T>> {
266 match self.protocol {
267 HeaderProtocol::Legacy => Ok(HeaderMessageWriter::Legacy(legacy_message_writer(
268 writer,
269 &self.legacy_key,
270 "legacy response writer",
271 )?)),
272 HeaderProtocol::V2 => Ok(HeaderMessageWriter::V2(V2MessageWriter::new(
273 writer,
274 self.v2_material()?.clone(),
275 DIRECTION_SERVER_TO_CLIENT,
276 0,
277 )?)),
278 }
279 }
280
281 pub fn continuation_reader<'a, T: AsyncReadExt + Unpin>(
282 &self,
283 reader: &'a mut T,
284 ) -> Result<HeaderMessageReader<'a, T>> {
285 match self.protocol {
286 HeaderProtocol::Legacy => Ok(HeaderMessageReader::Legacy(legacy_message_reader(
287 reader,
288 &self.legacy_key,
289 "legacy reader",
290 )?)),
291 HeaderProtocol::V2 => Ok(HeaderMessageReader::V2(V2MessageReader::new(
292 reader,
293 self.v2_material()?.clone(),
294 DIRECTION_CLIENT_TO_SERVER,
295 1,
296 )?)),
297 }
298 }
299}
300
301pub struct ServerInitialMessage {
302 pub payload: Vec<u8>,
303 pub session: ServerHeaderSession,
304 pub replay_fingerprint: Option<[u8; 32]>,
305 pub client_timestamp: Option<u64>,
306}
307
308pub struct ServerInitialError {
309 pub failure: AuthFailure,
310 pub response_session: Option<ServerHeaderSession>,
311 pub presented_key_id: Option<KeyId>,
312}
313
314impl ServerInitialError {
315 fn new(failure: AuthFailure) -> Self {
316 Self {
317 failure,
318 response_session: None,
319 presented_key_id: None,
320 }
321 }
322
323 fn fail(code: &'static str, message: impl Into<String>, retryable: bool) -> Self {
324 Self::new(AuthFailure::new(code, message, retryable))
325 }
326
327 fn fail_key(
328 code: &'static str,
329 message: impl Into<String>,
330 retryable: bool,
331 key_id: KeyId,
332 ) -> Self {
333 Self {
334 failure: AuthFailure::new(code, message, retryable),
335 response_session: None,
336 presented_key_id: Some(key_id),
337 }
338 }
339
340 fn from_failure_key(failure: AuthFailure, key_id: KeyId) -> Self {
341 Self {
342 failure,
343 response_session: None,
344 presented_key_id: Some(key_id),
345 }
346 }
347}
348
349impl fmt::Debug for ServerInitialError {
350 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
351 formatter
352 .debug_struct("ServerInitialError")
353 .field("failure", &self.failure)
354 .field("has_response_session", &self.response_session.is_some())
355 .field("presented_key_id", &self.presented_key_id)
356 .finish()
357 }
358}
359
360impl fmt::Display for ServerInitialError {
361 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
362 self.failure.fmt(formatter)
363 }
364}
365
366impl std::error::Error for ServerInitialError {}
367
368use std::fmt;
369
370#[derive(Clone)]
371pub struct ServerSecurity {
372 auth: AuthRuntime,
373 replay: Arc<Mutex<ReplayGuard>>,
374 failure_logs: Arc<Mutex<FailureLogLimiter>>,
375}
376
377#[allow(clippy::result_large_err)]
381impl ServerSecurity {
382 pub fn new(auth: AuthRuntime) -> Self {
383 let replay_path = auth.config().state_dir.join("connection.replay");
384 Self {
385 auth,
386 replay: Arc::new(Mutex::new(ReplayGuard::open(
387 Some(replay_path),
388 DEFAULT_REPLAY_FILTER_BYTES,
389 DEFAULT_REPLAY_WINDOW_SECONDS,
390 ))),
391 failure_logs: Arc::new(Mutex::new(FailureLogLimiter::default())),
392 }
393 }
394
395 pub fn auth(&self) -> &AuthRuntime {
396 &self.auth
397 }
398
399 pub fn record_failure_log(
400 &self,
401 peer_ip: std::net::IpAddr,
402 key_id: KeyId,
403 reason: &str,
404 ) -> FailureLogDecision {
405 self.failure_logs
406 .lock()
407 .record(peer_ip, key_id, reason, unix_seconds())
408 }
409
410 pub async fn read_initial<T: AsyncReadExt + Unpin>(
411 &self,
412 reader: &mut T,
413 ) -> std::result::Result<ServerInitialMessage, ServerInitialError> {
414 let mut first = [0_u8; 4];
415 reader.read_exact(&mut first).await.map_err(|error| {
416 ServerInitialError::new(AuthFailure::new(
417 "protocol_header_read_failed",
418 format!("failed to read initial protocol header: {error}"),
419 true,
420 ))
421 })?;
422 if first == PROTOCOL_V2_MAGIC {
423 self.read_v2_initial(reader).await
424 } else {
425 self.read_legacy_initial(reader, first).await
426 }
427 }
428
429 async fn read_legacy_initial<T: AsyncReadExt + Unpin>(
430 &self,
431 reader: &mut T,
432 checksum_bytes: [u8; 4],
433 ) -> std::result::Result<ServerInitialMessage, ServerInitialError> {
434 if !self.auth.legacy_protocol_allowed().unwrap_or(false) {
435 return Err(ServerInitialError::fail(
436 "legacy_protocol_disabled",
437 "legacy protocol is disabled by the administrator",
438 false,
439 ));
440 }
441 let key = self.auth.admin_key().map_err(ServerInitialError::new)?;
442 let checksum = u32::from_be_bytes(checksum_bytes);
443 let datalen = reader.read_u32().await.map_err(|error| {
444 ServerInitialError::fail(
445 "legacy_frame_invalid",
446 format!("failed to read legacy frame length: {error}"),
447 true,
448 )
449 })?;
450 if !valid_checksum_for_key(datalen, checksum, &key) || datalen > MAX_INITIAL_CIPHERTEXT_LEN
451 {
452 return Err(ServerInitialError::fail(
453 "legacy_frame_invalid",
454 "legacy frame checksum or length is invalid",
455 false,
456 ));
457 }
458 let mut encrypted = vec![0_u8; datalen as usize];
459 reader.read_exact(&mut encrypted).await.map_err(|error| {
460 ServerInitialError::fail(
461 "legacy_frame_invalid",
462 format!("failed to read legacy frame body: {error}"),
463 true,
464 )
465 })?;
466 let mut codec = Aes256GcmDeCodec::try_new(&key).map_err(|_| {
467 ServerInitialError::fail(
468 "legacy_decrypt_failed",
469 "failed to initialize legacy decryption",
470 false,
471 )
472 })?;
473 let plain = codec.decrypt(&mut encrypted).map_err(|_| {
474 ServerInitialError::fail(
475 "legacy_decrypt_failed",
476 "legacy credential or encrypted frame is invalid",
477 false,
478 )
479 })?;
480 let context = self
481 .auth
482 .authenticate_presented(ADMIN_KEY_ID, &key)
483 .map_err(ServerInitialError::new)?;
484 let legacy_guard = self
485 .auth
486 .record_legacy_connection()
487 .map_err(ServerInitialError::new)?;
488 Ok(ServerInitialMessage {
489 payload: plain.to_vec(),
490 session: ServerHeaderSession {
491 protocol: HeaderProtocol::Legacy,
492 legacy_key: key,
493 v2: None,
494 context: Some(context),
495 _legacy_guard: Some(legacy_guard),
496 },
497 replay_fingerprint: None,
498 client_timestamp: None,
499 })
500 }
501
502 async fn read_v2_initial<T: AsyncReadExt + Unpin>(
503 &self,
504 reader: &mut T,
505 ) -> std::result::Result<ServerInitialMessage, ServerInitialError> {
506 let mut remainder = [0_u8; FIRST_PREFIX_REMAINDER_LEN];
507 reader.read_exact(&mut remainder).await.map_err(|error| {
508 ServerInitialError::fail(
509 "protocol_v2_header_invalid",
510 format!("failed to read protocol-v2 header: {error}"),
511 true,
512 )
513 })?;
514 let version = remainder[0];
515 let flags = remainder[1];
516 let reserved = u16::from_be_bytes([remainder[2], remainder[3]]);
517 if version != PROTOCOL_V2_VERSION || flags != 0 || reserved != 0 {
518 return Err(ServerInitialError::fail(
519 if version != PROTOCOL_V2_VERSION {
520 "protocol_version_unsupported"
521 } else {
522 "protocol_v2_header_invalid"
523 },
524 format!(
525 "unsupported protocol header version={version} flags={flags} reserved={reserved}"
526 ),
527 false,
528 ));
529 }
530 let malformed =
534 || ServerInitialError::fail("protocol_error", "v2 prefix is malformed", false);
535 let key_id = KeyId::from_u64(u64::from_be_bytes(
536 remainder[4..12].try_into().map_err(|_| malformed())?,
537 ));
538 let salt: [u8; CONNECTION_SALT_LEN] =
539 remainder[12..28].try_into().map_err(|_| malformed())?;
540 let client_timestamp = u64::from_be_bytes(salt[..8].try_into().map_err(|_| malformed())?);
541 let now = unix_seconds();
542 if now.abs_diff(client_timestamp) > MAX_CONNECTION_CLOCK_SKEW_SECONDS {
543 return Err(ServerInitialError::fail_key(
544 "connection_timestamp_invalid",
545 "protocol-v2 connection timestamp is outside the accepted clock-skew window",
546 false,
547 key_id,
548 ));
549 }
550 let key = self
551 .auth
552 .derive_key(key_id)
553 .map_err(|failure| ServerInitialError::from_failure_key(failure, key_id))?;
554 let material = derive_material(key_id, &key, salt).map_err(|error| {
555 ServerInitialError::fail_key(
556 "protocol_v2_key_derivation_failed",
557 error.to_string(),
558 false,
559 key_id,
560 )
561 })?;
562 let mut session = v2_session(key, material.clone());
563 let (counter, ciphertext) = read_v2_frame(reader, 0, MAX_INITIAL_PLAINTEXT_LEN)
564 .await
565 .map_err(|error| {
566 ServerInitialError::fail_key(
567 "protocol_v2_decrypt_failed",
568 error.to_string(),
569 false,
570 key_id,
571 )
572 })?;
573 let mut current_ciphertext = ciphertext.clone();
574 let fingerprint = replay_fingerprint(key_id, &salt);
575 let work = match open_v2_payload(
576 &material,
577 DIRECTION_CLIENT_TO_SERVER,
578 counter,
579 &mut current_ciphertext,
580 ) {
581 Ok(payload) => FirstFlightWork::Live {
582 key,
583 payload,
584 error_session: session_without_context(&session),
585 },
586 Err(error) => {
587 match stale_root_first_flight(&self.auth, key_id, salt, counter, &ciphertext) {
588 Some(stale) => FirstFlightWork::Stale(stale),
589 None => {
590 return Err(first_flight_error(
591 "protocol_v2_decrypt_failed",
592 error.to_string(),
593 false,
594 key_id,
595 ));
596 }
597 }
598 }
599 };
600 let replay = self.replay.clone();
601 let auth = self.auth.clone();
602 let (payload, context) = tokio::task::spawn_blocking(move || {
603 evaluate_first_flight(&auth, &replay, key_id, fingerprint, work)
604 })
605 .await
606 .unwrap_or_else(|_| {
607 Err(first_flight_error(
608 "connection_replay_store_unavailable",
609 "failed to evaluate first-flight admission",
610 true,
611 key_id,
612 ))
613 })?;
614 session.context = Some(context);
615 Ok(ServerInitialMessage {
616 payload,
617 session,
618 replay_fingerprint: Some(fingerprint),
619 client_timestamp: Some(client_timestamp),
620 })
621 }
622}
623
624mod limiter;
625pub use limiter::FailureLogDecision;
626use limiter::FailureLogLimiter;
627
628pub enum HeaderMessageReader<'a, T: AsyncReadExt + Unpin> {
629 Legacy(CodecMessageReader<'a, T, Aes256GcmDeCodec>),
630 V2(V2MessageReader<'a, T>),
631}
632
633impl<T: AsyncReadExt + Unpin> MessageReader for HeaderMessageReader<'_, T> {
634 async fn read_msg(&mut self) -> Result<&'_ [u8]> {
635 match self {
636 Self::Legacy(reader) => reader.read_msg().await,
637 Self::V2(reader) => reader.read_msg().await,
638 }
639 }
640}
641
642pub enum HeaderMessageWriter<'a, T: AsyncWriteExt + Unpin> {
643 Legacy(CodecMessageWriter<'a, T, Aes256GcmEnCodec>),
644 V2(V2MessageWriter<'a, T>),
645}
646
647impl<T: AsyncWriteExt + Unpin> MessageWriter for HeaderMessageWriter<'_, T> {
648 async fn write_msg(&mut self, message: &[u8]) -> Result<()> {
649 match self {
650 Self::Legacy(writer) => writer.write_msg(message).await,
651 Self::V2(writer) => writer.write_msg(message).await,
652 }
653 }
654}
655
656mod frame;
657use frame::{V2Material, derive_material, first_prefix, open_v2_payload, read_v2_frame};
658pub use frame::{V2MessageReader, V2MessageWriter};
659mod replay;
660#[cfg(test)]
661use replay::RotatingBloom;
662use replay::{FirstFlightAdmit, ReplayGuard, replay_fingerprint};
663mod first_flight;
664use first_flight::*;
665fn legacy_message_reader<'a, T: AsyncReadExt + Unpin>(
666 reader: &'a mut T,
667 key: &AesKeyType,
668 action: &str,
669) -> Result<CodecMessageReader<'a, T, Aes256GcmDeCodec>> {
670 Ok(CodecMessageReader::for_session_key(
671 reader,
672 Aes256GcmDeCodec::try_new(key)
673 .map_err(|_| protocol_error(format!("failed to initialize {action}")))?,
674 *key,
675 ))
676}
677
678fn legacy_message_writer<'a, T: AsyncWriteExt + Unpin>(
679 writer: &'a mut T,
680 key: &AesKeyType,
681 action: &str,
682) -> Result<CodecMessageWriter<'a, T, Aes256GcmEnCodec>> {
683 Ok(CodecMessageWriter::for_session_key(
684 writer,
685 Aes256GcmEnCodec::try_new(key)
686 .map_err(|_| protocol_error(format!("failed to initialize {action}")))?,
687 *key,
688 ))
689}
690
691fn protocol_error(detail: impl Into<String>) -> Error {
692 Error::MsgProtocol {
693 detail: detail.into(),
694 }
695}
696
697fn unix_seconds() -> u64 {
698 SystemTime::now()
699 .duration_since(UNIX_EPOCH)
700 .unwrap_or_default()
701 .as_secs()
702}
703
704#[cfg(test)]
705mod tests;