1use std::{
2 fmt,
3 future::Future,
4 io::{Read, Write},
5 net::{SocketAddr, TcpStream, ToSocketAddrs},
6 sync::{
7 atomic::{AtomicBool, AtomicU64, Ordering},
8 Arc, Condvar, Mutex,
9 },
10 time::{Duration, Instant},
11};
12
13use rhiza_core::{LogHash, StoredCommand};
14use rhiza_quepaxa::{
15 DecisionProof, Error, Membership, ReadFenceObservation, ReadFenceRequest, RecordRequest,
16 RecordSummary, RecorderRpc, RejectReason,
17};
18use serde::{de::DeserializeOwned, Deserialize, Serialize};
19use tokio::io::{AsyncReadExt, AsyncWriteExt};
20use tokio_rustls::TlsAcceptor;
21
22use crate::{
23 authenticated_proposer_admitted, map_quorum_record_transport_error,
24 peer_credentials_authenticated, valid_recorder_command, valid_recorder_record, PeerConfig,
25 DEFAULT_PEER_CONCURRENCY, MAX_HTTP_BODY_BYTES, QUORUM_RECORD_REQUEST_TIMEOUT,
26 READ_FENCE_REQUEST_TIMEOUT,
27};
28
29const WIRE_VERSION: u16 = 3;
30const CONNECT_TIMEOUT: Duration = Duration::from_secs(2);
31const CALL_TIMEOUT: Duration = Duration::from_secs(10);
32const CONNECTIONS_PER_LANE: usize = 2;
33const MAX_SERVER_CONNECTIONS: usize = DEFAULT_PEER_CONCURRENCY * 4;
34const RECORDER_TLS_ALPN: &[u8] = b"rhiza-recorder/3";
35
36#[cfg(feature = "recorder-postcard-rpc")]
37mod postcard_rpc;
38#[cfg(feature = "recorder-postcard-rpc")]
39pub use postcard_rpc::{
40 serve_recorder_postcard_rpc, serve_recorder_postcard_rpc_tls,
41 RecorderPostcardRpcTlsClientConfig, RecorderPostcardRpcTlsServerConfig,
42 TcpPostcardRpcRecorderClient,
43};
44
45#[derive(Clone)]
46pub struct RecorderTlsServerConfig {
47 inner: Arc<rustls::ServerConfig>,
48}
49
50impl fmt::Debug for RecorderTlsServerConfig {
51 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
52 formatter
53 .debug_struct("RecorderTlsServerConfig")
54 .finish_non_exhaustive()
55 }
56}
57
58impl RecorderTlsServerConfig {
59 pub fn from_pem(certificate_chain_pem: &[u8], private_key_pem: &[u8]) -> Result<Self, String> {
60 let certificates = rustls_pemfile::certs(&mut std::io::Cursor::new(certificate_chain_pem))
61 .collect::<Result<Vec<_>, _>>()
62 .map_err(|_| "invalid recorder TLS certificate PEM".to_string())?;
63 if certificates.is_empty() {
64 return Err("recorder TLS certificate chain is empty".into());
65 }
66 let mut key_reader = std::io::Cursor::new(private_key_pem);
67 let private_key = rustls_pemfile::private_key(&mut key_reader)
68 .map_err(|_| "invalid recorder TLS private key PEM".to_string())?
69 .ok_or_else(|| "recorder TLS private key is missing".to_string())?;
70 if rustls_pemfile::private_key(&mut key_reader)
71 .map_err(|_| "invalid recorder TLS private key PEM".to_string())?
72 .is_some()
73 {
74 return Err("recorder TLS private key PEM contains multiple keys".into());
75 }
76 let mut config = rustls::ServerConfig::builder_with_provider(Arc::new(
77 rustls::crypto::ring::default_provider(),
78 ))
79 .with_protocol_versions(&[&rustls::version::TLS13])
80 .map_err(|_| "recorder TLS crypto provider does not support TLS 1.3".to_string())?
81 .with_no_client_auth()
82 .with_single_cert(certificates, private_key)
83 .map_err(|_| {
84 "recorder TLS certificate and private key are invalid or mismatched".to_string()
85 })?;
86 config.alpn_protocols = vec![RECORDER_TLS_ALPN.to_vec()];
87 config.max_early_data_size = 0;
88 Ok(Self {
89 inner: Arc::new(config),
90 })
91 }
92}
93
94#[derive(Clone)]
95pub struct RecorderTlsClientConfig {
96 inner: Arc<rustls::ClientConfig>,
97 server_name: rustls::pki_types::ServerName<'static>,
98}
99
100impl fmt::Debug for RecorderTlsClientConfig {
101 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
102 formatter
103 .debug_struct("RecorderTlsClientConfig")
104 .field("server_name", &self.server_name)
105 .finish_non_exhaustive()
106 }
107}
108
109impl RecorderTlsClientConfig {
110 pub fn from_ca_pem(ca_bundle_pem: &[u8], server_name: &str) -> Result<Self, String> {
111 let certificates = rustls_pemfile::certs(&mut std::io::Cursor::new(ca_bundle_pem))
112 .collect::<Result<Vec<_>, _>>()
113 .map_err(|_| "invalid recorder TLS CA bundle PEM".to_string())?;
114 if certificates.is_empty() {
115 return Err("recorder TLS CA bundle is empty".into());
116 }
117 let mut roots = rustls::RootCertStore::empty();
118 for certificate in certificates {
119 roots.add(certificate).map_err(|_| {
120 "recorder TLS CA bundle contains an invalid certificate".to_string()
121 })?;
122 }
123 let server_name = rustls::pki_types::ServerName::try_from(server_name.to_owned())
124 .map_err(|_| "invalid recorder TLS server name".to_string())?;
125 let mut config = rustls::ClientConfig::builder_with_provider(Arc::new(
126 rustls::crypto::ring::default_provider(),
127 ))
128 .with_protocol_versions(&[&rustls::version::TLS13])
129 .map_err(|_| "recorder TLS crypto provider does not support TLS 1.3".to_string())?
130 .with_root_certificates(roots)
131 .with_no_client_auth();
132 config.alpn_protocols = vec![RECORDER_TLS_ALPN.to_vec()];
133 config.enable_early_data = false;
134 Ok(Self {
135 inner: Arc::new(config),
136 server_name,
137 })
138 }
139}
140
141#[derive(Debug, Deserialize, Serialize)]
142struct Hello {
143 version: u16,
144 node_id: String,
145 recovery_generation: u64,
146 token: String,
147}
148
149#[derive(Debug, Deserialize, Serialize)]
150enum HelloReply {
151 Accepted { version: u16, recorder_id: String },
152 Rejected,
153}
154
155#[derive(Debug, Deserialize, Serialize)]
156struct RequestFrame {
157 version: u16,
158 request_id: u64,
159 remaining_deadline_ms: u32,
160 body: RecorderRequestBody,
161}
162
163#[derive(Debug, Deserialize, Serialize)]
164enum RecorderRequestBody {
165 Identity,
166 StoreCommand {
167 cluster_id: String,
168 epoch: u64,
169 config_id: u64,
170 config_digest: LogHash,
171 command_hash: LogHash,
172 command: StoredCommand,
173 },
174 FetchCommand {
175 cluster_id: String,
176 epoch: u64,
177 config_id: u64,
178 config_digest: LogHash,
179 command_hash: LogHash,
180 },
181 Record(RecordRequest),
182 InstallDecisionProof {
183 proof: DecisionProof,
184 members: Vec<String>,
185 },
186 InspectDecisionProof {
187 slot: u64,
188 },
189 InspectRecordSummary {
190 slot: u64,
191 },
192 ObserveReadFence(ReadFenceRequest),
193}
194
195#[derive(Debug, Deserialize, Serialize)]
196struct ResponseFrame {
197 version: u16,
198 request_id: u64,
199 body: RecorderResponseBody,
200}
201
202#[derive(Debug, Deserialize, Serialize)]
203enum RecorderResponseBody {
204 Identity(RpcResult<String>),
205 StoreCommand(RpcResult<()>),
206 FetchCommand(RpcResult<Option<StoredCommand>>),
207 Record(RpcResult<RecordSummary>),
208 InstallDecisionProof(RpcResult<()>),
209 InspectDecisionProof(RpcResult<Option<DecisionProof>>),
210 InspectRecordSummary(RpcResult<Option<RecordSummary>>),
211 ObserveReadFence(RpcResult<ReadFenceObservation>),
212}
213
214#[derive(Debug, Deserialize, Serialize)]
215enum RpcResult<T> {
216 Ok(T),
217 Rejected(RejectReason),
218 Error(String),
219 Overloaded,
220}
221
222impl<T> RpcResult<T> {
223 fn from_result(result: rhiza_quepaxa::Result<T>) -> Self {
224 match result {
225 Ok(value) => Self::Ok(value),
226 Err(Error::Rejected(reason)) => Self::Rejected(reason),
227 Err(error) => Self::Error(error.to_string()),
228 }
229 }
230
231 fn into_result(self) -> rhiza_quepaxa::Result<T> {
232 match self {
233 Self::Ok(value) => Ok(value),
234 Self::Rejected(reason) => Err(Error::Rejected(reason)),
235 Self::Error(message) => Err(Error::Io(message)),
236 Self::Overloaded => Err(Error::Io("recorder RPC overloaded".into())),
237 }
238 }
239}
240
241pub async fn serve_recorder_tcp<R, F>(
242 listener: tokio::net::TcpListener,
243 recorder: R,
244 peers: Vec<PeerConfig>,
245 recovery_generation: u64,
246 shutdown: F,
247) -> Result<(), String>
248where
249 R: RecorderRpc + Clone + Send + Sync + 'static,
250 F: Future<Output = ()> + Send,
251{
252 serve_recorder_tcp_inner(
253 listener,
254 recorder,
255 peers,
256 recovery_generation,
257 None,
258 shutdown,
259 )
260 .await
261}
262
263pub async fn serve_recorder_tcp_tls<R, F>(
264 listener: tokio::net::TcpListener,
265 recorder: R,
266 peers: Vec<PeerConfig>,
267 recovery_generation: u64,
268 tls: RecorderTlsServerConfig,
269 shutdown: F,
270) -> Result<(), String>
271where
272 R: RecorderRpc + Clone + Send + Sync + 'static,
273 F: Future<Output = ()> + Send,
274{
275 serve_recorder_tcp_inner(
276 listener,
277 recorder,
278 peers,
279 recovery_generation,
280 Some(tls.inner),
281 shutdown,
282 )
283 .await
284}
285
286async fn serve_recorder_tcp_inner<R, F>(
287 listener: tokio::net::TcpListener,
288 recorder: R,
289 peers: Vec<PeerConfig>,
290 recovery_generation: u64,
291 tls: Option<Arc<rustls::ServerConfig>>,
292 shutdown: F,
293) -> Result<(), String>
294where
295 R: RecorderRpc + Clone + Send + Sync + 'static,
296 F: Future<Output = ()> + Send,
297{
298 let peers: Arc<[PeerConfig]> = peers.into();
299 let slots = Arc::new(tokio::sync::Semaphore::new(DEFAULT_PEER_CONCURRENCY));
300 let connections = Arc::new(tokio::sync::Semaphore::new(MAX_SERVER_CONNECTIONS));
301 let reported_connection_error = Arc::new(AtomicBool::new(false));
302 let mut tasks = tokio::task::JoinSet::new();
303 tokio::pin!(shutdown);
304 loop {
305 tokio::select! {
306 () = &mut shutdown => break,
307 Some(_) = tasks.join_next(), if !tasks.is_empty() => {}
308 accepted = listener.accept() => {
309 let (stream, _) = accepted.map_err(|error| format!("recorder TCP accept failed: {error}"))?;
310 let Ok(connection) = connections.clone().try_acquire_owned() else {
311 continue;
312 };
313 let _ = stream.set_nodelay(true);
314 let recorder = recorder.clone();
315 let peers = peers.clone();
316 let slots = slots.clone();
317 let tls = tls.clone();
318 let reported_connection_error = Arc::clone(&reported_connection_error);
319 tasks.spawn(async move {
320 let _connection = connection;
321 let result = if let Some(config) = tls {
322 let acceptor = TlsAcceptor::from(config);
323 match tokio::time::timeout(CONNECT_TIMEOUT, acceptor.accept(stream)).await {
324 Ok(Ok(tls_stream)) => {
325 if tls_stream.get_ref().1.alpn_protocol() != Some(RECORDER_TLS_ALPN) {
326 Err("recorder TLS ALPN negotiation failed".to_string())
327 } else {
328 serve_connection(tls_stream, recorder, peers, recovery_generation, slots).await
329 }
330 }
331 Ok(Err(_)) => Err("recorder TLS handshake failed".to_string()),
332 Err(_) => Err("recorder TLS handshake timed out".to_string()),
333 }
334 } else {
335 serve_connection(stream, recorder, peers, recovery_generation, slots).await
336 };
337 if let Err(error) = result {
338 if error != "connection closed"
339 && !reported_connection_error.swap(true, Ordering::Relaxed)
340 {
341 eprintln!("recorder TCP connection rejected: {error}");
342 }
343 }
344 });
345 }
346 }
347 }
348 tasks.abort_all();
349 while tasks.join_next().await.is_some() {}
350 let _drained = slots
351 .acquire_many_owned(u32::try_from(DEFAULT_PEER_CONCURRENCY).unwrap_or(u32::MAX))
352 .await
353 .map_err(|_| "recorder operation semaphore closed during shutdown".to_string())?;
354 Ok(())
355}
356
357async fn serve_connection<R, S>(
358 mut stream: S,
359 recorder: R,
360 peers: Arc<[PeerConfig]>,
361 recovery_generation: u64,
362 slots: Arc<tokio::sync::Semaphore>,
363) -> Result<(), String>
364where
365 R: RecorderRpc + Clone + Send + Sync + 'static,
366 S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
367{
368 let hello_bytes = tokio::time::timeout(CALL_TIMEOUT, read_frame_async(&mut stream))
369 .await
370 .map_err(|_| "recorder HELLO timed out".to_string())??;
371 let hello: Hello = decode_exact(&hello_bytes)?;
372 if !hello_authenticated(&hello, &peers, recovery_generation) {
373 let _ = write_value_async_with_timeout(
374 &mut stream,
375 &HelloReply::Rejected,
376 "recorder HELLO rejection",
377 )
378 .await;
379 return Err("recorder HELLO rejected".into());
380 }
381 let identity_recorder = recorder.clone();
382 let recorder_id = tokio::task::spawn_blocking(move || identity_recorder.recorder_id())
383 .await
384 .map_err(|error| format!("recorder identity task failed: {error}"))?
385 .map_err(|error| error.to_string())?;
386 write_value_async_with_timeout(
387 &mut stream,
388 &HelloReply::Accepted {
389 version: WIRE_VERSION,
390 recorder_id,
391 },
392 "recorder HELLO response",
393 )
394 .await?;
395
396 loop {
397 let request = match read_frame_async(&mut stream).await {
398 Ok(bytes) => decode_exact::<RequestFrame>(&bytes)?,
399 Err(error) if error == "connection closed" => return Ok(()),
400 Err(error) => return Err(error),
401 };
402 if request.version != WIRE_VERSION || request.remaining_deadline_ms == 0 {
403 return Err("invalid recorder request envelope".into());
404 }
405 let request_id = request.request_id;
406 let operation = response_operation(&request.body);
407 let dispatch_deadline = Instant::now()
408 + Duration::from_millis(u64::from(request.remaining_deadline_ms)).min(CALL_TIMEOUT);
409 let permit = match slots.clone().try_acquire_owned() {
410 Ok(permit) => permit,
411 Err(_) => {
412 write_value_async_with_timeout(
413 &mut stream,
414 &ResponseFrame {
415 version: WIRE_VERSION,
416 request_id,
417 body: overloaded_response(operation),
418 },
419 "recorder overload response",
420 )
421 .await?;
422 continue;
423 }
424 };
425 let body = dispatch_with_deadline(
426 recorder.clone(),
427 request.body,
428 operation,
429 permit,
430 dispatch_deadline,
431 hello.node_id.clone(),
432 Arc::clone(&peers),
433 )
434 .await;
435 write_value_async_with_timeout(
436 &mut stream,
437 &ResponseFrame {
438 version: WIRE_VERSION,
439 request_id,
440 body,
441 },
442 "recorder response",
443 )
444 .await?;
445 }
446}
447
448async fn dispatch_with_deadline<R>(
449 recorder: R,
450 body: RecorderRequestBody,
451 operation: Operation,
452 permit: tokio::sync::OwnedSemaphorePermit,
453 deadline: Instant,
454 authenticated_peer_id: String,
455 peers: Arc<[PeerConfig]>,
456) -> RecorderResponseBody
457where
458 R: RecorderRpc + Send + Sync + 'static,
459{
460 if deadline <= Instant::now() {
461 return error_response(operation, "recorder RPC deadline exceeded".into());
462 }
463 let dispatched = tokio::task::spawn_blocking(move || {
464 let _permit = permit;
465 dispatch(recorder, body, &authenticated_peer_id, &peers)
466 });
467 match tokio::time::timeout_at(deadline.into(), dispatched).await {
468 Ok(Ok(response)) => response,
469 Ok(Err(error)) => error_response(operation, error.to_string()),
470 Err(_) => error_response(operation, "recorder RPC deadline exceeded".into()),
471 }
472}
473
474fn hello_authenticated(hello: &Hello, peers: &[PeerConfig], recovery_generation: u64) -> bool {
475 hello.version == WIRE_VERSION
476 && hello.recovery_generation == recovery_generation
477 && peer_credentials_authenticated(&hello.node_id, &hello.token, peers)
478}
479
480#[derive(Clone, Copy, Eq, PartialEq)]
481enum Operation {
482 Identity,
483 StoreCommand,
484 FetchCommand,
485 Record,
486 InstallDecisionProof,
487 InspectDecisionProof,
488 InspectRecordSummary,
489 ObserveReadFence,
490}
491
492fn response_operation(request: &RecorderRequestBody) -> Operation {
493 match request {
494 RecorderRequestBody::Identity => Operation::Identity,
495 RecorderRequestBody::StoreCommand { .. } => Operation::StoreCommand,
496 RecorderRequestBody::FetchCommand { .. } => Operation::FetchCommand,
497 RecorderRequestBody::Record(_) => Operation::Record,
498 RecorderRequestBody::InstallDecisionProof { .. } => Operation::InstallDecisionProof,
499 RecorderRequestBody::InspectDecisionProof { .. } => Operation::InspectDecisionProof,
500 RecorderRequestBody::InspectRecordSummary { .. } => Operation::InspectRecordSummary,
501 RecorderRequestBody::ObserveReadFence(_) => Operation::ObserveReadFence,
502 }
503}
504
505fn dispatch<R: RecorderRpc>(
506 recorder: R,
507 request: RecorderRequestBody,
508 authenticated_peer_id: &str,
509 peers: &[PeerConfig],
510) -> RecorderResponseBody {
511 match request {
512 RecorderRequestBody::Identity => {
513 RecorderResponseBody::Identity(RpcResult::from_result(recorder.recorder_id()))
514 }
515 RecorderRequestBody::StoreCommand {
516 cluster_id,
517 epoch,
518 config_id,
519 config_digest,
520 command_hash,
521 command,
522 } => {
523 let result = if !valid_recorder_command(&command) {
524 Err(Error::Rejected(RejectReason::InvalidRequest))
525 } else {
526 recorder.store_command_for(
527 cluster_id,
528 epoch,
529 config_id,
530 config_digest,
531 command_hash,
532 command,
533 )
534 };
535 RecorderResponseBody::StoreCommand(RpcResult::from_result(result))
536 }
537 RecorderRequestBody::FetchCommand {
538 cluster_id,
539 epoch,
540 config_id,
541 config_digest,
542 command_hash,
543 } => RecorderResponseBody::FetchCommand(RpcResult::from_result(
544 recorder.fetch_command_for(cluster_id, epoch, config_id, config_digest, command_hash),
545 )),
546 RecorderRequestBody::Record(request) => {
547 let result = if !valid_recorder_record(&request)
548 || !authenticated_proposer_admitted(
549 authenticated_peer_id,
550 &request.proposal.proposer_id,
551 peers,
552 ) {
553 Err(Error::Rejected(RejectReason::InvalidRequest))
554 } else {
555 recorder.record(request)
556 };
557 RecorderResponseBody::Record(RpcResult::from_result(result))
558 }
559 RecorderRequestBody::InstallDecisionProof { proof, members } => {
560 let result = if !authenticated_proposer_admitted(
561 authenticated_peer_id,
562 &proof.proposal().proposer_id,
563 peers,
564 ) {
565 Err(Error::Rejected(RejectReason::InvalidRequest))
566 } else {
567 Membership::from_voters(members)
568 .and_then(|membership| recorder.install_decision_proof(proof, &membership))
569 };
570 RecorderResponseBody::InstallDecisionProof(RpcResult::from_result(result))
571 }
572 RecorderRequestBody::InspectDecisionProof { slot } => {
573 RecorderResponseBody::InspectDecisionProof(RpcResult::from_result(
574 recorder.inspect_decision_proof(slot),
575 ))
576 }
577 RecorderRequestBody::InspectRecordSummary { slot } => {
578 RecorderResponseBody::InspectRecordSummary(RpcResult::from_result(
579 recorder.inspect_record_summary(slot),
580 ))
581 }
582 RecorderRequestBody::ObserveReadFence(request) => RecorderResponseBody::ObserveReadFence(
583 RpcResult::from_result(recorder.observe_read_fence(request)),
584 ),
585 }
586}
587
588fn overloaded_response(operation: Operation) -> RecorderResponseBody {
589 match operation {
590 Operation::Identity => RecorderResponseBody::Identity(RpcResult::Overloaded),
591 Operation::StoreCommand => RecorderResponseBody::StoreCommand(RpcResult::Overloaded),
592 Operation::FetchCommand => RecorderResponseBody::FetchCommand(RpcResult::Overloaded),
593 Operation::Record => RecorderResponseBody::Record(RpcResult::Overloaded),
594 Operation::InstallDecisionProof => {
595 RecorderResponseBody::InstallDecisionProof(RpcResult::Overloaded)
596 }
597 Operation::InspectDecisionProof => {
598 RecorderResponseBody::InspectDecisionProof(RpcResult::Overloaded)
599 }
600 Operation::InspectRecordSummary => {
601 RecorderResponseBody::InspectRecordSummary(RpcResult::Overloaded)
602 }
603 Operation::ObserveReadFence => {
604 RecorderResponseBody::ObserveReadFence(RpcResult::Overloaded)
605 }
606 }
607}
608
609fn error_response(operation: Operation, message: String) -> RecorderResponseBody {
610 match operation {
611 Operation::Identity => RecorderResponseBody::Identity(RpcResult::Error(message)),
612 Operation::StoreCommand => RecorderResponseBody::StoreCommand(RpcResult::Error(message)),
613 Operation::FetchCommand => RecorderResponseBody::FetchCommand(RpcResult::Error(message)),
614 Operation::Record => RecorderResponseBody::Record(RpcResult::Error(message)),
615 Operation::InstallDecisionProof => {
616 RecorderResponseBody::InstallDecisionProof(RpcResult::Error(message))
617 }
618 Operation::InspectDecisionProof => {
619 RecorderResponseBody::InspectDecisionProof(RpcResult::Error(message))
620 }
621 Operation::InspectRecordSummary => {
622 RecorderResponseBody::InspectRecordSummary(RpcResult::Error(message))
623 }
624 Operation::ObserveReadFence => {
625 RecorderResponseBody::ObserveReadFence(RpcResult::Error(message))
626 }
627 }
628}
629
630async fn read_frame_async<R: tokio::io::AsyncRead + Unpin>(
631 reader: &mut R,
632) -> Result<Vec<u8>, String> {
633 let mut length = [0_u8; 4];
634 match reader.read_exact(&mut length).await {
635 Ok(_) => {}
636 Err(error) if error.kind() == std::io::ErrorKind::UnexpectedEof => {
637 return Err("connection closed".into())
638 }
639 Err(error) => return Err(error.to_string()),
640 }
641 let length = usize::try_from(u32::from_be_bytes(length)).unwrap_or(usize::MAX);
642 if length == 0 || length > MAX_HTTP_BODY_BYTES {
643 return Err("invalid recorder frame length".into());
644 }
645 let mut frame = vec![0; length];
646 reader
647 .read_exact(&mut frame)
648 .await
649 .map_err(|error| error.to_string())?;
650 Ok(frame)
651}
652
653async fn write_value_async<W: tokio::io::AsyncWrite + Unpin, T: Serialize>(
654 writer: &mut W,
655 value: &T,
656) -> Result<(), String> {
657 let encoded = postcard::to_allocvec(value).map_err(|error| error.to_string())?;
658 write_frame_async(writer, &encoded).await
659}
660
661async fn write_value_async_with_timeout<W: tokio::io::AsyncWrite + Unpin, T: Serialize>(
662 writer: &mut W,
663 value: &T,
664 operation: &str,
665) -> Result<(), String> {
666 tokio::time::timeout(CALL_TIMEOUT, write_value_async(writer, value))
667 .await
668 .map_err(|_| format!("{operation} timed out"))?
669}
670
671async fn write_frame_async<W: tokio::io::AsyncWrite + Unpin>(
672 writer: &mut W,
673 frame: &[u8],
674) -> Result<(), String> {
675 let length = frame_length(frame)?;
676 writer
677 .write_all(&length)
678 .await
679 .map_err(|error| error.to_string())?;
680 writer
681 .write_all(frame)
682 .await
683 .map_err(|error| error.to_string())
684}
685
686fn decode_exact<T: DeserializeOwned>(bytes: &[u8]) -> Result<T, String> {
687 let (value, trailing) = postcard::take_from_bytes(bytes).map_err(|error| error.to_string())?;
688 if !trailing.is_empty() {
689 return Err("trailing recorder frame bytes".into());
690 }
691 Ok(value)
692}
693
694fn frame_length(frame: &[u8]) -> Result<[u8; 4], String> {
695 if frame.is_empty() || frame.len() > MAX_HTTP_BODY_BYTES {
696 return Err("invalid recorder frame length".into());
697 }
698 let length = u32::try_from(frame.len()).map_err(|_| "recorder frame is too large")?;
699 Ok(length.to_be_bytes())
700}
701
702struct ConnectionPool {
703 state: Mutex<PoolState>,
704 available: Condvar,
705}
706
707#[derive(Default)]
708struct PoolState {
709 idle: Vec<RecorderClientStream>,
710 open: usize,
711}
712
713trait DeadlineClock {
714 fn now(&self) -> Instant;
715}
716
717#[derive(Clone, Copy)]
718struct SystemClock;
719
720impl DeadlineClock for SystemClock {
721 fn now(&self) -> Instant {
722 Instant::now()
723 }
724}
725
726trait SocketTimeouts {
727 fn set_read_timeout(&self, timeout: Option<Duration>) -> std::io::Result<()>;
728 fn set_write_timeout(&self, timeout: Option<Duration>) -> std::io::Result<()>;
729}
730
731impl SocketTimeouts for TcpStream {
732 fn set_read_timeout(&self, timeout: Option<Duration>) -> std::io::Result<()> {
733 TcpStream::set_read_timeout(self, timeout)
734 }
735
736 fn set_write_timeout(&self, timeout: Option<Duration>) -> std::io::Result<()> {
737 TcpStream::set_write_timeout(self, timeout)
738 }
739}
740
741struct DeadlineStream<S, C = SystemClock> {
742 inner: S,
743 deadline: Instant,
744 clock: C,
745}
746
747impl<S> DeadlineStream<S> {
748 fn new(inner: S, deadline: Instant) -> Self {
749 Self::new_with_clock(inner, deadline, SystemClock)
750 }
751}
752
753impl<S, C> DeadlineStream<S, C> {
754 fn new_with_clock(inner: S, deadline: Instant, clock: C) -> Self {
755 Self {
756 inner,
757 deadline,
758 clock,
759 }
760 }
761
762 fn set_deadline(&mut self, deadline: Instant) {
763 self.deadline = deadline;
764 }
765}
766
767impl<S, C: DeadlineClock> DeadlineStream<S, C> {
768 fn remaining(&self) -> std::io::Result<Duration> {
769 let remaining = self.deadline.saturating_duration_since(self.clock.now());
770 if remaining.is_zero() {
771 Err(std::io::Error::new(
772 std::io::ErrorKind::TimedOut,
773 "recorder RPC deadline exceeded",
774 ))
775 } else {
776 Ok(remaining)
777 }
778 }
779}
780
781impl<S: Read + SocketTimeouts, C: DeadlineClock> Read for DeadlineStream<S, C> {
782 fn read(&mut self, buffer: &mut [u8]) -> std::io::Result<usize> {
783 self.inner.set_read_timeout(Some(self.remaining()?))?;
784 self.inner.read(buffer)
785 }
786}
787
788impl<S: Write + SocketTimeouts, C: DeadlineClock> Write for DeadlineStream<S, C> {
789 fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
790 self.inner.set_write_timeout(Some(self.remaining()?))?;
791 self.inner.write(buffer)
792 }
793
794 fn flush(&mut self) -> std::io::Result<()> {
795 self.inner.set_write_timeout(Some(self.remaining()?))?;
796 self.inner.flush()
797 }
798}
799
800enum RecorderClientStream {
801 Plain(DeadlineStream<TcpStream>),
802 Tls(Box<rustls::StreamOwned<rustls::ClientConnection, DeadlineStream<TcpStream>>>),
803}
804
805impl RecorderClientStream {
806 fn set_deadline(&mut self, deadline: Instant) {
807 match self {
808 Self::Plain(stream) => stream.set_deadline(deadline),
809 Self::Tls(stream) => stream.sock.set_deadline(deadline),
810 }
811 }
812
813 fn ensure_deadline(&self) -> std::io::Result<()> {
814 match self {
815 Self::Plain(stream) => stream.remaining().map(|_| ()),
816 Self::Tls(stream) => stream.sock.remaining().map(|_| ()),
817 }
818 }
819}
820
821impl Read for RecorderClientStream {
822 fn read(&mut self, buffer: &mut [u8]) -> std::io::Result<usize> {
823 self.ensure_deadline()?;
824 match self {
825 Self::Plain(stream) => stream.read(buffer),
826 Self::Tls(stream) => stream.read(buffer),
827 }
828 }
829}
830
831impl Write for RecorderClientStream {
832 fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
833 self.ensure_deadline()?;
834 match self {
835 Self::Plain(stream) => stream.write(buffer),
836 Self::Tls(stream) => stream.write(buffer),
837 }
838 }
839
840 fn flush(&mut self) -> std::io::Result<()> {
841 self.ensure_deadline()?;
842 match self {
843 Self::Plain(stream) => stream.flush(),
844 Self::Tls(stream) => stream.flush(),
845 }
846 }
847}
848
849#[derive(Clone)]
850enum ClientTransport {
851 Plain,
852 Tls(RecorderTlsClientConfig),
853}
854
855impl ConnectionPool {
856 fn new() -> Self {
857 Self {
858 state: Mutex::new(PoolState::default()),
859 available: Condvar::new(),
860 }
861 }
862}
863
864pub struct TcpPostcardRecorderClient {
865 address: String,
866 expected_recorder_id: String,
867 local_node_id: String,
868 peer_token: String,
869 recovery_generation: u64,
870 transport: ClientTransport,
871 call_timeout: Duration,
872 consensus: ConnectionPool,
873 control: ConnectionPool,
874 next_request_id: AtomicU64,
875}
876
877impl fmt::Debug for TcpPostcardRecorderClient {
878 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
879 formatter
880 .debug_struct("TcpPostcardRecorderClient")
881 .field("address", &self.address)
882 .field("expected_recorder_id", &self.expected_recorder_id)
883 .field("local_node_id", &self.local_node_id)
884 .field("peer_token", &"[redacted]")
885 .field("recovery_generation", &self.recovery_generation)
886 .field("call_timeout", &self.call_timeout)
887 .field(
888 "transport",
889 &match self.transport {
890 ClientTransport::Plain => "plain",
891 ClientTransport::Tls(_) => "tls",
892 },
893 )
894 .finish()
895 }
896}
897
898impl TcpPostcardRecorderClient {
899 pub fn new(
900 address: impl ToString,
901 expected_recorder_id: impl Into<String>,
902 local_node_id: impl Into<String>,
903 peer_token: impl Into<String>,
904 recovery_generation: u64,
905 ) -> Result<Self, String> {
906 Self::new_with_transport(
907 address,
908 expected_recorder_id,
909 local_node_id,
910 peer_token,
911 recovery_generation,
912 ClientTransport::Plain,
913 )
914 }
915
916 pub fn new_tls(
917 address: impl ToString,
918 expected_recorder_id: impl Into<String>,
919 local_node_id: impl Into<String>,
920 peer_token: impl Into<String>,
921 recovery_generation: u64,
922 tls: RecorderTlsClientConfig,
923 ) -> Result<Self, String> {
924 Self::new_with_transport(
925 address,
926 expected_recorder_id,
927 local_node_id,
928 peer_token,
929 recovery_generation,
930 ClientTransport::Tls(tls),
931 )
932 }
933
934 fn new_with_transport(
935 address: impl ToString,
936 expected_recorder_id: impl Into<String>,
937 local_node_id: impl Into<String>,
938 peer_token: impl Into<String>,
939 recovery_generation: u64,
940 transport: ClientTransport,
941 ) -> Result<Self, String> {
942 Self::new_with_transport_and_timeout(
943 address,
944 expected_recorder_id,
945 local_node_id,
946 peer_token,
947 recovery_generation,
948 transport,
949 CALL_TIMEOUT,
950 )
951 }
952
953 fn new_with_transport_and_timeout(
954 address: impl ToString,
955 expected_recorder_id: impl Into<String>,
956 local_node_id: impl Into<String>,
957 peer_token: impl Into<String>,
958 recovery_generation: u64,
959 transport: ClientTransport,
960 call_timeout: Duration,
961 ) -> Result<Self, String> {
962 let address = address.to_string();
963 validate_recorder_tcp_endpoint(&address)?;
964 let expected_recorder_id = expected_recorder_id.into();
965 let local_node_id = local_node_id.into();
966 let peer_token = peer_token.into();
967 if expected_recorder_id.trim().is_empty()
968 || local_node_id.trim().is_empty()
969 || peer_token.trim().is_empty()
970 || recovery_generation == 0
971 || call_timeout.is_zero()
972 {
973 return Err("invalid recorder TCP client identity".into());
974 }
975 Ok(Self {
976 address,
977 expected_recorder_id,
978 local_node_id,
979 peer_token,
980 recovery_generation,
981 transport,
982 call_timeout,
983 consensus: ConnectionPool::new(),
984 control: ConnectionPool::new(),
985 next_request_id: AtomicU64::new(1),
986 })
987 }
988
989 fn exchange(
990 &self,
991 request: RecorderRequestBody,
992 consensus: bool,
993 ) -> rhiza_quepaxa::Result<RecorderResponseBody> {
994 self.exchange_with_timeout(request, consensus, self.call_timeout)
995 }
996
997 fn exchange_with_timeout(
998 &self,
999 request: RecorderRequestBody,
1000 consensus: bool,
1001 timeout: Duration,
1002 ) -> rhiza_quepaxa::Result<RecorderResponseBody> {
1003 let deadline = Instant::now() + timeout.min(self.call_timeout);
1004 let pool = if consensus {
1005 &self.consensus
1006 } else {
1007 &self.control
1008 };
1009 let mut stream = self.checkout(pool, deadline)?;
1010 let request_id = self.next_request_id.fetch_add(1, Ordering::Relaxed);
1011 let operation = response_operation(&request);
1012 stream.set_deadline(deadline);
1013 let remaining_deadline_ms = match advertised_remaining_deadline_ms(deadline) {
1014 Ok(remaining) => remaining,
1015 Err(error) => {
1016 self.discard(pool);
1017 return Err(error);
1018 }
1019 };
1020 let frame = RequestFrame {
1021 version: WIRE_VERSION,
1022 request_id,
1023 remaining_deadline_ms,
1024 body: request,
1025 };
1026 let result = write_value_sync(&mut stream, &frame)
1027 .and_then(|()| read_frame_sync(&mut stream))
1028 .and_then(|bytes| decode_exact::<ResponseFrame>(&bytes));
1029 match result {
1030 Ok(response)
1031 if response.version == WIRE_VERSION
1032 && response.request_id == request_id
1033 && response_matches(operation, &response.body) =>
1034 {
1035 self.checkin(pool, stream);
1036 Ok(response.body)
1037 }
1038 Ok(_) => {
1039 self.discard(pool);
1040 Err(Error::Decode("recorder response envelope mismatch".into()))
1041 }
1042 Err(error) => {
1043 self.discard(pool);
1044 Err(Error::Io(error))
1045 }
1046 }
1047 }
1048
1049 fn checkout(
1050 &self,
1051 pool: &ConnectionPool,
1052 deadline: Instant,
1053 ) -> rhiza_quepaxa::Result<RecorderClientStream> {
1054 loop {
1055 let mut state = pool
1056 .state
1057 .lock()
1058 .map_err(|_| Error::Io("recorder connection pool lock poisoned".into()))?;
1059 if let Some(stream) = state.idle.pop() {
1060 return Ok(stream);
1061 }
1062 if state.open < CONNECTIONS_PER_LANE {
1063 state.open += 1;
1064 drop(state);
1065 return match self.connect(deadline) {
1066 Ok(stream) => Ok(stream),
1067 Err(error) => {
1068 self.discard(pool);
1069 Err(Error::Io(error))
1070 }
1071 };
1072 }
1073 let remaining = deadline.saturating_duration_since(Instant::now());
1074 if remaining.is_zero() {
1075 return Err(Error::Io("recorder connection checkout timed out".into()));
1076 }
1077 let (next, wait) = pool
1078 .available
1079 .wait_timeout(state, remaining)
1080 .map_err(|_| Error::Io("recorder connection pool lock poisoned".into()))?;
1081 drop(next);
1082 if wait.timed_out() {
1083 return Err(Error::Io("recorder connection checkout timed out".into()));
1084 }
1085 }
1086 }
1087
1088 fn connect(&self, deadline: Instant) -> Result<RecorderClientStream, String> {
1089 let remaining = deadline.saturating_duration_since(Instant::now());
1090 let connect_timeout = CONNECT_TIMEOUT.min(remaining);
1091 if connect_timeout.is_zero() {
1092 return Err("recorder connect deadline exceeded".into());
1093 }
1094 let mut last_error = None;
1095 let mut socket = None;
1096 let resolved_addresses = self
1097 .address
1098 .to_socket_addrs()
1099 .map_err(|error| format!("cannot resolve recorder TCP address: {error}"))?
1100 .collect::<Vec<SocketAddr>>();
1101 if resolved_addresses.is_empty() {
1102 return Err("recorder TCP address resolved to no endpoints".into());
1103 }
1104 for address in &resolved_addresses {
1105 let remaining = deadline.saturating_duration_since(Instant::now());
1106 if remaining.is_zero() {
1107 break;
1108 }
1109 match TcpStream::connect_timeout(address, connect_timeout.min(remaining)) {
1110 Ok(connected) => {
1111 socket = Some(connected);
1112 break;
1113 }
1114 Err(error) => last_error = Some(error),
1115 }
1116 }
1117 let socket = socket.ok_or_else(|| {
1118 format!(
1119 "recorder TCP connect failed: {}",
1120 last_error
1121 .map(|error| error.to_string())
1122 .unwrap_or_else(|| "deadline exceeded".into())
1123 )
1124 })?;
1125 socket
1126 .set_nodelay(true)
1127 .map_err(|error| format!("cannot set recorder TCP_NODELAY: {error}"))?;
1128 let socket = DeadlineStream::new(socket, deadline);
1129 let mut stream = match &self.transport {
1130 ClientTransport::Plain => RecorderClientStream::Plain(socket),
1131 ClientTransport::Tls(tls) => {
1132 let connection =
1133 rustls::ClientConnection::new(Arc::clone(&tls.inner), tls.server_name.clone())
1134 .map_err(|_| "cannot initialize recorder TLS connection".to_string())?;
1135 let mut stream = rustls::StreamOwned::new(connection, socket);
1136 while stream.conn.is_handshaking() {
1137 let remaining = deadline.saturating_duration_since(Instant::now());
1138 if remaining.is_zero() {
1139 return Err("recorder TLS handshake timed out".into());
1140 }
1141 stream
1142 .conn
1143 .complete_io(&mut stream.sock)
1144 .map_err(|_| "recorder TLS handshake failed".to_string())?;
1145 }
1146 if stream.conn.alpn_protocol() != Some(RECORDER_TLS_ALPN) {
1147 return Err("recorder TLS ALPN negotiation failed".into());
1148 }
1149 RecorderClientStream::Tls(Box::new(stream))
1150 }
1151 };
1152 write_value_sync(
1153 &mut stream,
1154 &Hello {
1155 version: WIRE_VERSION,
1156 node_id: self.local_node_id.clone(),
1157 recovery_generation: self.recovery_generation,
1158 token: self.peer_token.clone(),
1159 },
1160 )?;
1161 let reply: HelloReply = decode_exact(&read_frame_sync(&mut stream)?)?;
1162 match reply {
1163 HelloReply::Accepted {
1164 version,
1165 recorder_id,
1166 } if version == WIRE_VERSION && recorder_id == self.expected_recorder_id => Ok(stream),
1167 HelloReply::Accepted { .. } => Err("recorder identity mismatch".into()),
1168 HelloReply::Rejected => Err("recorder HELLO rejected".into()),
1169 }
1170 }
1171
1172 fn checkin(&self, pool: &ConnectionPool, stream: RecorderClientStream) {
1173 if let Ok(mut state) = pool.state.lock() {
1174 state.idle.push(stream);
1175 pool.available.notify_one();
1176 }
1177 }
1178
1179 fn discard(&self, pool: &ConnectionPool) {
1180 if let Ok(mut state) = pool.state.lock() {
1181 state.open = state.open.saturating_sub(1);
1182 pool.available.notify_one();
1183 }
1184 }
1185}
1186
1187pub fn validate_recorder_tcp_endpoint(address: &str) -> Result<(), String> {
1188 let parsed = reqwest::Url::parse(&format!("tcp://{address}"))
1189 .map_err(|_| "invalid recorder TCP address".to_string())?;
1190 if parsed.host_str().is_none()
1191 || parsed.port().is_none()
1192 || !matches!(parsed.path(), "" | "/")
1193 || parsed.query().is_some()
1194 || parsed.fragment().is_some()
1195 {
1196 return Err("invalid recorder TCP address".into());
1197 }
1198 Ok(())
1199}
1200
1201fn response_matches(operation: Operation, response: &RecorderResponseBody) -> bool {
1202 matches!(
1203 (operation, response),
1204 (Operation::Identity, RecorderResponseBody::Identity(_))
1205 | (
1206 Operation::StoreCommand,
1207 RecorderResponseBody::StoreCommand(_)
1208 )
1209 | (
1210 Operation::FetchCommand,
1211 RecorderResponseBody::FetchCommand(_)
1212 )
1213 | (Operation::Record, RecorderResponseBody::Record(_))
1214 | (
1215 Operation::InstallDecisionProof,
1216 RecorderResponseBody::InstallDecisionProof(_)
1217 )
1218 | (
1219 Operation::InspectDecisionProof,
1220 RecorderResponseBody::InspectDecisionProof(_)
1221 )
1222 | (
1223 Operation::InspectRecordSummary,
1224 RecorderResponseBody::InspectRecordSummary(_)
1225 )
1226 | (
1227 Operation::ObserveReadFence,
1228 RecorderResponseBody::ObserveReadFence(_)
1229 )
1230 )
1231}
1232
1233impl RecorderRpc for TcpPostcardRecorderClient {
1234 fn recorder_id(&self) -> rhiza_quepaxa::Result<String> {
1235 match self.exchange(RecorderRequestBody::Identity, false)? {
1236 RecorderResponseBody::Identity(result) => result.into_result(),
1237 _ => Err(Error::Decode("recorder response operation mismatch".into())),
1238 }
1239 }
1240
1241 fn store_command_for(
1242 &self,
1243 cluster_id: String,
1244 epoch: u64,
1245 config_id: u64,
1246 config_digest: LogHash,
1247 command_hash: LogHash,
1248 command: StoredCommand,
1249 ) -> rhiza_quepaxa::Result<()> {
1250 let request = RecorderRequestBody::StoreCommand {
1251 cluster_id,
1252 epoch,
1253 config_id,
1254 config_digest,
1255 command_hash,
1256 command,
1257 };
1258 match self.exchange(request, false)? {
1259 RecorderResponseBody::StoreCommand(result) => result.into_result(),
1260 _ => Err(Error::Decode("recorder response operation mismatch".into())),
1261 }
1262 }
1263
1264 fn fetch_command_for(
1265 &self,
1266 cluster_id: String,
1267 epoch: u64,
1268 config_id: u64,
1269 config_digest: LogHash,
1270 command_hash: LogHash,
1271 ) -> rhiza_quepaxa::Result<Option<StoredCommand>> {
1272 let request = RecorderRequestBody::FetchCommand {
1273 cluster_id,
1274 epoch,
1275 config_id,
1276 config_digest,
1277 command_hash,
1278 };
1279 match self.exchange(request, false)? {
1280 RecorderResponseBody::FetchCommand(result) => result.into_result(),
1281 _ => Err(Error::Decode("recorder response operation mismatch".into())),
1282 }
1283 }
1284
1285 fn record(&self, request: RecordRequest) -> rhiza_quepaxa::Result<RecordSummary> {
1286 let response = self
1287 .exchange_with_timeout(
1288 RecorderRequestBody::Record(request),
1289 true,
1290 QUORUM_RECORD_REQUEST_TIMEOUT,
1291 )
1292 .map_err(map_quorum_record_transport_error)?;
1293 match response {
1294 RecorderResponseBody::Record(result) => result.into_result(),
1295 _ => Err(Error::Decode("recorder response operation mismatch".into())),
1296 }
1297 .map_err(map_quorum_record_transport_error)
1298 }
1299
1300 fn install_decision_proof(
1301 &self,
1302 proof: DecisionProof,
1303 membership: &Membership,
1304 ) -> rhiza_quepaxa::Result<()> {
1305 let request = RecorderRequestBody::InstallDecisionProof {
1306 proof,
1307 members: membership.members().to_vec(),
1308 };
1309 match self.exchange(request, true)? {
1310 RecorderResponseBody::InstallDecisionProof(result) => result.into_result(),
1311 _ => Err(Error::Decode("recorder response operation mismatch".into())),
1312 }
1313 }
1314
1315 fn inspect_decision_proof(&self, slot: u64) -> rhiza_quepaxa::Result<Option<DecisionProof>> {
1316 let request = RecorderRequestBody::InspectDecisionProof { slot };
1317 match self.exchange(request, false)? {
1318 RecorderResponseBody::InspectDecisionProof(result) => result.into_result(),
1319 _ => Err(Error::Decode("recorder response operation mismatch".into())),
1320 }
1321 }
1322
1323 fn inspect_record_summary(&self, slot: u64) -> rhiza_quepaxa::Result<Option<RecordSummary>> {
1324 let request = RecorderRequestBody::InspectRecordSummary { slot };
1325 match self.exchange(request, false)? {
1326 RecorderResponseBody::InspectRecordSummary(result) => result.into_result(),
1327 _ => Err(Error::Decode("recorder response operation mismatch".into())),
1328 }
1329 }
1330
1331 fn supports_context_read_fence(&self) -> bool {
1332 true
1333 }
1334
1335 fn observe_read_fence(
1336 &self,
1337 request: ReadFenceRequest,
1338 ) -> rhiza_quepaxa::Result<ReadFenceObservation> {
1339 match self.exchange_with_timeout(
1340 RecorderRequestBody::ObserveReadFence(request),
1341 false,
1342 READ_FENCE_REQUEST_TIMEOUT,
1343 )? {
1344 RecorderResponseBody::ObserveReadFence(result) => result.into_result(),
1345 _ => Err(Error::Decode("recorder response operation mismatch".into())),
1346 }
1347 }
1348}
1349
1350fn advertised_remaining_deadline_ms(deadline: Instant) -> rhiza_quepaxa::Result<u32> {
1351 let remaining = deadline.saturating_duration_since(Instant::now());
1352 if remaining.is_zero() {
1353 return Err(Error::Io("recorder RPC deadline exceeded".into()));
1354 }
1355 Ok(u32::try_from(remaining.as_millis())
1356 .unwrap_or(u32::MAX)
1357 .max(1))
1358}
1359
1360fn read_frame_sync(reader: &mut impl Read) -> Result<Vec<u8>, String> {
1361 let mut length = [0_u8; 4];
1362 reader
1363 .read_exact(&mut length)
1364 .map_err(|error| error.to_string())?;
1365 let length = usize::try_from(u32::from_be_bytes(length)).unwrap_or(usize::MAX);
1366 if length == 0 || length > MAX_HTTP_BODY_BYTES {
1367 return Err("invalid recorder frame length".into());
1368 }
1369 let mut frame = vec![0; length];
1370 reader
1371 .read_exact(&mut frame)
1372 .map_err(|error| error.to_string())?;
1373 Ok(frame)
1374}
1375
1376fn write_value_sync(writer: &mut impl Write, value: &impl Serialize) -> Result<(), String> {
1377 let encoded = postcard::to_allocvec(value).map_err(|error| error.to_string())?;
1378 let length = frame_length(&encoded)?;
1379 writer
1380 .write_all(&length)
1381 .map_err(|error| error.to_string())?;
1382 writer
1383 .write_all(&encoded)
1384 .map_err(|error| error.to_string())?;
1385 writer.flush().map_err(|error| error.to_string())
1386}
1387
1388#[cfg(test)]
1389mod tests {
1390 use super::*;
1391 use std::{
1392 cell::{Cell, RefCell},
1393 collections::VecDeque,
1394 net::TcpListener,
1395 rc::Rc,
1396 sync::{
1397 atomic::{AtomicUsize, Ordering},
1398 mpsc,
1399 },
1400 thread,
1401 };
1402
1403 #[derive(Clone)]
1404 struct FakeClock {
1405 origin: Instant,
1406 elapsed: Rc<Cell<Duration>>,
1407 }
1408
1409 impl DeadlineClock for FakeClock {
1410 fn now(&self) -> Instant {
1411 self.origin + self.elapsed.get()
1412 }
1413 }
1414
1415 struct SlowPartialIo {
1416 clock: FakeClock,
1417 step: Duration,
1418 input: VecDeque<u8>,
1419 read_timeout: Cell<Option<Duration>>,
1420 write_timeout: Cell<Option<Duration>>,
1421 read_timeouts: Rc<RefCell<Vec<Duration>>>,
1422 write_timeouts: Rc<RefCell<Vec<Duration>>>,
1423 }
1424
1425 type SlowPartialFixture = (
1426 SlowPartialIo,
1427 FakeClock,
1428 Rc<RefCell<Vec<Duration>>>,
1429 Rc<RefCell<Vec<Duration>>>,
1430 );
1431
1432 impl SlowPartialIo {
1433 fn spend(&self, timeout: Option<Duration>) -> std::io::Result<()> {
1434 let timeout = timeout.expect("deadline stream must configure a timeout");
1435 if self.step > timeout {
1436 self.clock.elapsed.set(self.clock.elapsed.get() + timeout);
1437 return Err(std::io::Error::new(
1438 std::io::ErrorKind::TimedOut,
1439 "scripted operation reached its timeout",
1440 ));
1441 }
1442 self.clock.elapsed.set(self.clock.elapsed.get() + self.step);
1443 Ok(())
1444 }
1445 }
1446
1447 impl SocketTimeouts for SlowPartialIo {
1448 fn set_read_timeout(&self, timeout: Option<Duration>) -> std::io::Result<()> {
1449 self.read_timeout.set(timeout);
1450 self.read_timeouts
1451 .borrow_mut()
1452 .push(timeout.expect("read timeout must be bounded"));
1453 Ok(())
1454 }
1455
1456 fn set_write_timeout(&self, timeout: Option<Duration>) -> std::io::Result<()> {
1457 self.write_timeout.set(timeout);
1458 self.write_timeouts
1459 .borrow_mut()
1460 .push(timeout.expect("write timeout must be bounded"));
1461 Ok(())
1462 }
1463 }
1464
1465 impl Read for SlowPartialIo {
1466 fn read(&mut self, buffer: &mut [u8]) -> std::io::Result<usize> {
1467 self.spend(self.read_timeout.get())?;
1468 let Some(byte) = self.input.pop_front() else {
1469 return Ok(0);
1470 };
1471 buffer[0] = byte;
1472 Ok(1)
1473 }
1474 }
1475
1476 impl Write for SlowPartialIo {
1477 fn write(&mut self, _buffer: &[u8]) -> std::io::Result<usize> {
1478 self.spend(self.write_timeout.get())?;
1479 Ok(1)
1480 }
1481
1482 fn flush(&mut self) -> std::io::Result<()> {
1483 self.spend(self.write_timeout.get())
1484 }
1485 }
1486
1487 fn slow_partial_io(input: Vec<u8>) -> SlowPartialFixture {
1488 let clock = FakeClock {
1489 origin: Instant::now(),
1490 elapsed: Rc::new(Cell::new(Duration::ZERO)),
1491 };
1492 let read_timeouts = Rc::new(RefCell::new(Vec::new()));
1493 let write_timeouts = Rc::new(RefCell::new(Vec::new()));
1494 (
1495 SlowPartialIo {
1496 clock: clock.clone(),
1497 step: Duration::from_millis(30),
1498 input: input.into(),
1499 read_timeout: Cell::new(None),
1500 write_timeout: Cell::new(None),
1501 read_timeouts: Rc::clone(&read_timeouts),
1502 write_timeouts: Rc::clone(&write_timeouts),
1503 },
1504 clock,
1505 read_timeouts,
1506 write_timeouts,
1507 )
1508 }
1509
1510 #[test]
1511 fn sync_frame_read_refreshes_timeout_against_one_absolute_deadline() {
1512 let mut input = 1_u32.to_be_bytes().to_vec();
1513 input.push(42);
1514 let (io, clock, read_timeouts, _) = slow_partial_io(input);
1515 let deadline = clock.now() + Duration::from_millis(100);
1516 let mut stream = DeadlineStream::new_with_clock(io, deadline, clock.clone());
1517
1518 assert!(read_frame_sync(&mut stream).is_err());
1519
1520 assert_eq!(clock.elapsed.get(), Duration::from_millis(100));
1521 assert_eq!(
1522 *read_timeouts.borrow(),
1523 [100, 70, 40, 10].map(Duration::from_millis)
1524 );
1525 }
1526
1527 #[test]
1528 fn sync_frame_write_refreshes_timeout_against_one_absolute_deadline() {
1529 let (io, clock, _, write_timeouts) = slow_partial_io(Vec::new());
1530 let deadline = clock.now() + Duration::from_millis(100);
1531 let mut stream = DeadlineStream::new_with_clock(io, deadline, clock.clone());
1532
1533 assert!(write_value_sync(&mut stream, &42_u64).is_err());
1534
1535 assert_eq!(clock.elapsed.get(), Duration::from_millis(100));
1536 assert_eq!(
1537 *write_timeouts.borrow(),
1538 [100, 70, 40, 10].map(Duration::from_millis)
1539 );
1540 }
1541
1542 #[test]
1543 fn legacy_client_bounds_partial_response_drip_by_sender_deadline() {
1544 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1545 let address = listener.local_addr().unwrap();
1546 let (advertised_tx, advertised_rx) = mpsc::channel();
1547 let server = thread::spawn(move || {
1548 let (mut stream, _) = listener.accept().unwrap();
1549 let hello: Hello = decode_exact(&read_frame_sync(&mut stream).unwrap()).unwrap();
1550 assert_eq!(hello.version, WIRE_VERSION);
1551 thread::sleep(Duration::from_millis(80));
1552 write_value_sync(
1553 &mut stream,
1554 &HelloReply::Accepted {
1555 version: WIRE_VERSION,
1556 recorder_id: "node-1".into(),
1557 },
1558 )
1559 .unwrap();
1560 let request: RequestFrame =
1561 decode_exact(&read_frame_sync(&mut stream).unwrap()).unwrap();
1562 advertised_tx.send(request.remaining_deadline_ms).unwrap();
1563 for byte in [0_u8, 0, 0, 1, 0] {
1564 thread::sleep(Duration::from_millis(120));
1565 if stream.write_all(&[byte]).is_err() {
1566 break;
1567 }
1568 }
1569 });
1570 let client = TcpPostcardRecorderClient::new_with_transport_and_timeout(
1571 address,
1572 "node-1",
1573 "node-2",
1574 "peer-token-2",
1575 7,
1576 ClientTransport::Plain,
1577 Duration::from_millis(400),
1578 )
1579 .unwrap();
1580
1581 let started = Instant::now();
1582 assert!(client.recorder_id().is_err());
1583 let elapsed = started.elapsed();
1584
1585 let advertised = advertised_rx.recv_timeout(Duration::from_secs(1)).unwrap();
1586 assert!(
1587 advertised > 0 && advertised <= 350,
1588 "advertised {advertised}ms"
1589 );
1590 assert!(
1591 elapsed < Duration::from_millis(550),
1592 "partial response exceeded the sender-owned deadline: {elapsed:?}"
1593 );
1594 server.join().unwrap();
1595 }
1596
1597 #[test]
1598 fn legacy_read_fence_uses_the_short_control_deadline() {
1599 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1600 let address = listener.local_addr().unwrap();
1601 let server = thread::spawn(move || {
1602 let (_stream, _) = listener.accept().unwrap();
1603 thread::sleep(Duration::from_secs(2));
1604 });
1605 let client = TcpPostcardRecorderClient::new_with_transport_and_timeout(
1606 address,
1607 "node-1",
1608 "node-2",
1609 "peer-token-2",
1610 7,
1611 ClientTransport::Plain,
1612 Duration::from_secs(5),
1613 )
1614 .unwrap();
1615
1616 let started = Instant::now();
1617 assert!(client
1618 .observe_read_fence(ReadFenceRequest {
1619 cluster_id: "cluster".into(),
1620 epoch: 1,
1621 config_id: 1,
1622 config_digest: LogHash::ZERO,
1623 slot: 1,
1624 })
1625 .is_err());
1626 assert!(started.elapsed() < Duration::from_millis(1_500));
1627 server.join().unwrap();
1628 }
1629
1630 #[test]
1631 fn legacy_record_transport_failure_releases_the_quorum_attempt_promptly() {
1632 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
1633 let address = listener.local_addr().unwrap();
1634 let server = thread::spawn(move || {
1635 let (_stream, _) = listener.accept().unwrap();
1636 thread::sleep(Duration::from_secs(2));
1637 });
1638 let client = TcpPostcardRecorderClient::new_with_transport_and_timeout(
1639 address,
1640 "node-1",
1641 "node-2",
1642 "peer-token-2",
1643 7,
1644 ClientTransport::Plain,
1645 Duration::from_secs(5),
1646 )
1647 .unwrap();
1648
1649 let started = Instant::now();
1650 let result = client.record(RecordRequest {
1651 cluster_id: "cluster".into(),
1652 epoch: 1,
1653 config_id: 1,
1654 config_digest: LogHash::ZERO,
1655 slot: 1,
1656 step: 1,
1657 proposal: rhiza_quepaxa::Proposal::nil(),
1658 command: None,
1659 });
1660
1661 assert!(matches!(result, Err(Error::ProposeFailed)));
1662 assert!(started.elapsed() < Duration::from_millis(1_500));
1663 server.join().unwrap();
1664 }
1665
1666 #[derive(Clone)]
1667 struct BlockingMutation {
1668 started: mpsc::Sender<()>,
1669 release: Arc<(Mutex<bool>, Condvar)>,
1670 completed: Arc<AtomicUsize>,
1671 }
1672
1673 #[derive(Clone)]
1674 struct CountingMutation {
1675 calls: Arc<AtomicUsize>,
1676 }
1677
1678 impl RecorderRpc for CountingMutation {
1679 fn recorder_id(&self) -> rhiza_quepaxa::Result<String> {
1680 Ok("node-1".into())
1681 }
1682
1683 fn store_command_for(
1684 &self,
1685 _cluster_id: String,
1686 _epoch: u64,
1687 _config_id: u64,
1688 _config_digest: LogHash,
1689 _command_hash: LogHash,
1690 _command: StoredCommand,
1691 ) -> rhiza_quepaxa::Result<()> {
1692 self.calls.fetch_add(1, Ordering::SeqCst);
1693 Ok(())
1694 }
1695 }
1696
1697 impl RecorderRpc for BlockingMutation {
1698 fn recorder_id(&self) -> rhiza_quepaxa::Result<String> {
1699 Ok("node-1".into())
1700 }
1701
1702 fn store_command_for(
1703 &self,
1704 _cluster_id: String,
1705 _epoch: u64,
1706 _config_id: u64,
1707 _config_digest: LogHash,
1708 _command_hash: LogHash,
1709 _command: StoredCommand,
1710 ) -> rhiza_quepaxa::Result<()> {
1711 self.started.send(()).unwrap();
1712 let (released, ready) = &*self.release;
1713 let mut released = released.lock().unwrap();
1714 while !*released {
1715 released = ready.wait(released).unwrap();
1716 }
1717 self.completed.fetch_add(1, Ordering::SeqCst);
1718 Ok(())
1719 }
1720 }
1721
1722 fn peers() -> Vec<PeerConfig> {
1723 (1..=3)
1724 .map(|index| {
1725 PeerConfig::new(
1726 format!("node-{index}"),
1727 format!("http://node-{index}:8081"),
1728 format!("peer-token-{index}"),
1729 )
1730 .unwrap()
1731 })
1732 .collect()
1733 }
1734
1735 #[tokio::test]
1736 async fn request_expired_before_dispatch_never_reaches_recorder() {
1737 let calls = Arc::new(AtomicUsize::new(0));
1738 let command = StoredCommand::new(rhiza_core::EntryType::Command, b"expired".to_vec());
1739 let permit = Arc::new(tokio::sync::Semaphore::new(1))
1740 .acquire_owned()
1741 .await
1742 .unwrap();
1743
1744 let response = dispatch_with_deadline(
1745 CountingMutation {
1746 calls: Arc::clone(&calls),
1747 },
1748 RecorderRequestBody::StoreCommand {
1749 cluster_id: "rhiza:sql:cluster-a".into(),
1750 epoch: 1,
1751 config_id: 1,
1752 config_digest: LogHash::ZERO,
1753 command_hash: command.hash(),
1754 command,
1755 },
1756 Operation::StoreCommand,
1757 permit,
1758 Instant::now() - Duration::from_millis(1),
1759 "node-1".into(),
1760 peers().into(),
1761 )
1762 .await;
1763
1764 assert_eq!(calls.load(Ordering::SeqCst), 0);
1765 assert!(matches!(
1766 response,
1767 RecorderResponseBody::StoreCommand(RpcResult::Error(message))
1768 if message.contains("deadline")
1769 ));
1770 }
1771
1772 #[tokio::test]
1773 async fn saturated_server_returns_overload_without_calling_recorder() {
1774 let calls = Arc::new(AtomicUsize::new(0));
1775 let slots = Arc::new(tokio::sync::Semaphore::new(1));
1776 let held = Arc::clone(&slots).acquire_owned().await.unwrap();
1777 let (mut client, server_stream) = tokio::io::duplex(4096);
1778 let server = tokio::spawn(serve_connection(
1779 server_stream,
1780 CountingMutation {
1781 calls: Arc::clone(&calls),
1782 },
1783 peers().into(),
1784 7,
1785 slots,
1786 ));
1787 write_value_async(
1788 &mut client,
1789 &Hello {
1790 version: WIRE_VERSION,
1791 node_id: "node-2".into(),
1792 recovery_generation: 7,
1793 token: "peer-token-2".into(),
1794 },
1795 )
1796 .await
1797 .unwrap();
1798 assert!(matches!(
1799 decode_exact::<HelloReply>(&read_frame_async(&mut client).await.unwrap()).unwrap(),
1800 HelloReply::Accepted { .. }
1801 ));
1802 let command = StoredCommand::new(rhiza_core::EntryType::Command, b"overloaded".to_vec());
1803 write_value_async(
1804 &mut client,
1805 &RequestFrame {
1806 version: WIRE_VERSION,
1807 request_id: 1,
1808 remaining_deadline_ms: 1_000,
1809 body: RecorderRequestBody::StoreCommand {
1810 cluster_id: "rhiza:sql:cluster-a".into(),
1811 epoch: 1,
1812 config_id: 1,
1813 config_digest: LogHash::ZERO,
1814 command_hash: command.hash(),
1815 command,
1816 },
1817 },
1818 )
1819 .await
1820 .unwrap();
1821
1822 let response: ResponseFrame =
1823 decode_exact(&read_frame_async(&mut client).await.unwrap()).unwrap();
1824 assert!(matches!(
1825 response.body,
1826 RecorderResponseBody::StoreCommand(RpcResult::Overloaded)
1827 ));
1828 assert_eq!(calls.load(Ordering::SeqCst), 0);
1829
1830 drop(client);
1831 drop(held);
1832 server.await.unwrap().unwrap();
1833 }
1834
1835 #[tokio::test(flavor = "multi_thread")]
1836 async fn server_deadline_returns_while_admitted_mutation_finishes_and_shutdown_drains_it() {
1837 let (started_tx, started_rx) = mpsc::channel();
1838 let release = Arc::new((Mutex::new(false), Condvar::new()));
1839 let completed = Arc::new(AtomicUsize::new(0));
1840 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1841 let address = listener.local_addr().unwrap();
1842 let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel();
1843 let server = tokio::spawn(serve_recorder_tcp(
1844 listener,
1845 BlockingMutation {
1846 started: started_tx,
1847 release: Arc::clone(&release),
1848 completed: Arc::clone(&completed),
1849 },
1850 peers(),
1851 7,
1852 async move {
1853 let _ = shutdown_rx.await;
1854 },
1855 ));
1856 let mut stream = tokio::net::TcpStream::connect(address).await.unwrap();
1857 write_value_async(
1858 &mut stream,
1859 &Hello {
1860 version: WIRE_VERSION,
1861 node_id: "node-2".into(),
1862 recovery_generation: 7,
1863 token: "peer-token-2".into(),
1864 },
1865 )
1866 .await
1867 .unwrap();
1868 assert!(matches!(
1869 decode_exact::<HelloReply>(&read_frame_async(&mut stream).await.unwrap()).unwrap(),
1870 HelloReply::Accepted { .. }
1871 ));
1872 let membership = Membership::new(["node-1", "node-2", "node-3"]).unwrap();
1873 let command = StoredCommand::new(rhiza_core::EntryType::Command, b"slow".to_vec());
1874 write_value_async(
1875 &mut stream,
1876 &RequestFrame {
1877 version: WIRE_VERSION,
1878 request_id: 1,
1879 remaining_deadline_ms: 50,
1880 body: RecorderRequestBody::StoreCommand {
1881 cluster_id: "rhiza:sql:cluster-a".into(),
1882 epoch: 1,
1883 config_id: 1,
1884 config_digest: membership.digest(),
1885 command_hash: command.hash(),
1886 command,
1887 },
1888 },
1889 )
1890 .await
1891 .unwrap();
1892 started_rx.recv_timeout(Duration::from_secs(1)).unwrap();
1893 let response =
1894 tokio::time::timeout(Duration::from_millis(300), read_frame_async(&mut stream)).await;
1895 shutdown_tx.send(()).unwrap();
1896 tokio::time::sleep(Duration::from_millis(20)).await;
1897 assert!(!server.is_finished());
1898 let (released, ready) = &*release;
1899 *released.lock().unwrap() = true;
1900 ready.notify_all();
1901 server.await.unwrap().unwrap();
1902 assert_eq!(completed.load(Ordering::SeqCst), 1);
1903 let response = response
1904 .expect("server must answer the advertised deadline")
1905 .unwrap();
1906 assert!(matches!(
1907 decode_exact::<ResponseFrame>(&response).unwrap().body,
1908 RecorderResponseBody::StoreCommand(RpcResult::Error(message))
1909 if message.contains("deadline")
1910 ));
1911 }
1912
1913 #[test]
1914 fn postcard_decoder_rejects_trailing_bytes_and_wrong_hello_version() {
1915 assert_eq!(WIRE_VERSION, 3);
1916 assert_eq!(RECORDER_TLS_ALPN, b"rhiza-recorder/3");
1917 let hello = Hello {
1918 version: WIRE_VERSION,
1919 node_id: "node-1".into(),
1920 recovery_generation: 7,
1921 token: "peer-token-1".into(),
1922 };
1923 let mut encoded = postcard::to_allocvec(&hello).unwrap();
1924 encoded.push(0);
1925 assert!(decode_exact::<Hello>(&encoded).is_err());
1926
1927 let wrong_version = Hello {
1928 version: WIRE_VERSION + 1,
1929 ..hello
1930 };
1931 assert!(!hello_authenticated(&wrong_version, &[], 7));
1932 }
1933
1934 #[test]
1935 fn recorder_tcp_endpoint_accepts_socket_and_dns_addresses_without_paths() {
1936 assert!(validate_recorder_tcp_endpoint("127.0.0.1:8082").is_ok());
1937 assert!(validate_recorder_tcp_endpoint("node-1.internal:8082").is_ok());
1938 assert!(validate_recorder_tcp_endpoint("[::1]:8082").is_ok());
1939 assert!(validate_recorder_tcp_endpoint("127.0.0.1").is_err());
1940 assert!(validate_recorder_tcp_endpoint("127.0.0.1:8082/path").is_err());
1941 }
1942
1943 #[tokio::test]
1944 async fn frame_reader_rejects_zero_oversize_and_truncated_frames() {
1945 for length in [0_u32, u32::try_from(MAX_HTTP_BODY_BYTES + 1).unwrap()] {
1946 let (mut writer, mut reader) = tokio::io::duplex(16);
1947 writer.write_all(&length.to_be_bytes()).await.unwrap();
1948 assert!(read_frame_async(&mut reader).await.is_err());
1949 }
1950
1951 let (mut writer, mut reader) = tokio::io::duplex(16);
1952 writer.write_all(&4_u32.to_be_bytes()).await.unwrap();
1953 writer.write_all(&[1, 2]).await.unwrap();
1954 drop(writer);
1955 assert!(read_frame_async(&mut reader).await.is_err());
1956 }
1957}