client.rs (44981B)
1 //! Bounded HTTP/1.1 administration client over Unix sockets. 2 3 use core::fmt; 4 use serde::{Deserialize, Deserializer, Serialize, de, de::DeserializeOwned}; 5 use std::collections::BTreeSet; 6 use std::error::Error; 7 use std::io; 8 use std::path::{Path, PathBuf}; 9 use std::pin::Pin; 10 use std::task::{Context, Poll}; 11 12 use bytes::Bytes; 13 use http::header::{CONTENT_LENGTH, CONTENT_TYPE, HOST, HeaderMap, HeaderValue}; 14 use http::{Method, Request, StatusCode, Uri, Version}; 15 use http_body_util::{BodyExt, Full, LengthLimitError, Limited}; 16 use hyper::client::conn::http1; 17 use hyper_util::rt::TokioIo; 18 use tokio::io::{AsyncRead, AsyncWrite, ReadBuf}; 19 use tokio::task::JoinHandle; 20 21 #[cfg(test)] 22 use super::test_support; 23 use super::{ 24 ADMIN_CONTRACT_VERSION, ADMIN_ROUTE_PATH_MAX_UTF8_BYTES, AdminCorrelationId, 25 AdminFailureResponse, AdminHttpMethod, AdminOperationId, AdminSuccessResponse, 26 AdminTransportLimits, 27 }; 28 29 const JSON_CONTENT_TYPE: &str = "application/json"; 30 const ADMIN_CLIENT_TARGET_MAX_UTF8_BYTES: usize = 32 * 1024; 31 const HTTP_MINIMUM_MAX_BUFFER_SIZE: usize = 8 * 1024; 32 33 #[derive(Clone, Copy, Debug, PartialEq, Eq)] 34 pub enum AdminClientTargetError { 35 Empty, 36 TooLong, 37 InvalidUri, 38 AuthorityForbidden, 39 WrongVersionPrefix, 40 PathTooLong, 41 EmptySegment, 42 PatternForbidden, 43 InvalidPercentEncoding, 44 } 45 46 impl fmt::Display for AdminClientTargetError { 47 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { 48 formatter.write_str("admin client target is not a bounded canonical v1 target") 49 } 50 } 51 52 impl Error for AdminClientTargetError {} 53 54 /// Validated relative v1 request target. Query contents are redacted from `Debug`. 55 #[derive(Clone, PartialEq, Eq)] 56 pub struct AdminClientTarget { 57 uri: Uri, 58 } 59 60 impl AdminClientTarget { 61 pub fn new(value: impl AsRef<str>) -> Result<Self, AdminClientTargetError> { 62 let value = value.as_ref(); 63 if value.is_empty() { 64 return Err(AdminClientTargetError::Empty); 65 } 66 if value.len() > ADMIN_CLIENT_TARGET_MAX_UTF8_BYTES { 67 return Err(AdminClientTargetError::TooLong); 68 } 69 let uri = value 70 .parse::<Uri>() 71 .map_err(|_| AdminClientTargetError::InvalidUri)?; 72 if uri.scheme().is_some() || uri.authority().is_some() { 73 return Err(AdminClientTargetError::AuthorityForbidden); 74 } 75 let path = uri.path(); 76 if !path.starts_with("/v1/") || path.ends_with('/') { 77 return Err(AdminClientTargetError::WrongVersionPrefix); 78 } 79 if path.len() > ADMIN_ROUTE_PATH_MAX_UTF8_BYTES { 80 return Err(AdminClientTargetError::PathTooLong); 81 } 82 if path.contains("//") { 83 return Err(AdminClientTargetError::EmptySegment); 84 } 85 if path.contains(['{', '}']) { 86 return Err(AdminClientTargetError::PatternForbidden); 87 } 88 if !valid_percent_encoding(value.as_bytes()) { 89 return Err(AdminClientTargetError::InvalidPercentEncoding); 90 } 91 Ok(Self { uri }) 92 } 93 94 #[must_use] 95 pub fn as_str(&self) -> &str { 96 self.uri.path_and_query().map_or("", |value| value.as_str()) 97 } 98 99 #[must_use] 100 pub fn path(&self) -> &str { 101 self.uri.path() 102 } 103 104 #[must_use] 105 pub fn query(&self) -> Option<&str> { 106 self.uri.query() 107 } 108 } 109 110 impl fmt::Debug for AdminClientTarget { 111 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { 112 formatter 113 .debug_struct("AdminClientTarget") 114 .field("path", &"[redacted]") 115 .field("has_query", &self.uri.query().is_some()) 116 .finish() 117 } 118 } 119 120 #[derive(Clone, Copy, Debug, PartialEq, Eq)] 121 pub enum AdminClientErrorKind { 122 SocketPath, 123 QueryLimit, 124 RequestEncoding, 125 RequestLimit, 126 Connect, 127 Transport, 128 Deadline, 129 ResponseHeaders, 130 ResponseLimit, 131 ResponseContentType, 132 ResponseHttpVersion, 133 MalformedResponse, 134 UnsupportedContractVersion, 135 ServerFailure, 136 } 137 138 pub struct AdminClientError { 139 kind: AdminClientErrorKind, 140 io_kind: Option<io::ErrorKind>, 141 failure: Option<AdminFailureResponse>, 142 source: Option<Box<dyn Error + Send + Sync>>, 143 } 144 145 impl AdminClientError { 146 fn simple(kind: AdminClientErrorKind) -> Self { 147 Self { 148 kind, 149 io_kind: None, 150 failure: None, 151 source: None, 152 } 153 } 154 155 fn sourced<E>(kind: AdminClientErrorKind, source: E) -> Self 156 where 157 E: Error + Send + Sync + 'static, 158 { 159 Self { 160 kind, 161 io_kind: None, 162 failure: None, 163 source: Some(Box::new(source)), 164 } 165 } 166 167 fn connect(source: io::Error) -> Self { 168 Self { 169 kind: AdminClientErrorKind::Connect, 170 io_kind: Some(source.kind()), 171 failure: None, 172 source: Some(Box::new(source)), 173 } 174 } 175 176 fn boxed(kind: AdminClientErrorKind, source: Box<dyn Error + Send + Sync>) -> Self { 177 Self { 178 kind, 179 io_kind: None, 180 failure: None, 181 source: Some(source), 182 } 183 } 184 185 fn server(failure: AdminFailureResponse) -> Self { 186 Self { 187 kind: AdminClientErrorKind::ServerFailure, 188 io_kind: None, 189 failure: Some(failure), 190 source: None, 191 } 192 } 193 194 #[must_use] 195 pub const fn kind(&self) -> AdminClientErrorKind { 196 self.kind 197 } 198 199 #[must_use] 200 pub const fn io_kind(&self) -> Option<io::ErrorKind> { 201 self.io_kind 202 } 203 204 #[must_use] 205 pub const fn failure(&self) -> Option<&AdminFailureResponse> { 206 self.failure.as_ref() 207 } 208 } 209 210 impl fmt::Debug for AdminClientError { 211 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { 212 formatter 213 .debug_struct("AdminClientError") 214 .field("kind", &self.kind) 215 .field("io_kind", &self.io_kind) 216 .field( 217 "failure_code", 218 &self.failure.as_ref().map(|failure| failure.error().code()), 219 ) 220 .field("source", &self.source.as_ref().map(|_| "<redacted>")) 221 .finish() 222 } 223 } 224 225 impl fmt::Display for AdminClientError { 226 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { 227 formatter.write_str("bounded local admin request failed") 228 } 229 } 230 231 impl Error for AdminClientError { 232 fn source(&self) -> Option<&(dyn Error + 'static)> { 233 self.source.as_deref().map(|source| source as _) 234 } 235 } 236 237 /// Bounded client for one service-owned Unix administration socket. 238 pub struct AdminClient { 239 socket_path: PathBuf, 240 limits: AdminTransportLimits, 241 } 242 243 impl AdminClient { 244 pub fn new( 245 socket_path: impl Into<PathBuf>, 246 limits: AdminTransportLimits, 247 ) -> Result<Self, AdminClientError> { 248 let socket_path = socket_path.into(); 249 if !socket_path.is_absolute() || socket_path.file_name().is_none() { 250 return Err(AdminClientError::simple(AdminClientErrorKind::SocketPath)); 251 } 252 Ok(Self { 253 socket_path, 254 limits, 255 }) 256 } 257 258 #[must_use] 259 pub fn socket_path(&self) -> &Path { 260 &self.socket_path 261 } 262 263 #[must_use] 264 pub const fn limits(&self) -> AdminTransportLimits { 265 self.limits 266 } 267 268 pub async fn get<T>( 269 &self, 270 target: &AdminClientTarget, 271 ) -> Result<AdminSuccessResponse<T>, AdminClientError> 272 where 273 T: DeserializeOwned + Serialize, 274 { 275 self.execute(AdminHttpMethod::Get, target, Bytes::new()) 276 .await 277 } 278 279 pub async fn mutate<RequestBody, ResponseBody>( 280 &self, 281 target: &AdminClientTarget, 282 operation_id: AdminOperationId, 283 correlation_id: Option<AdminCorrelationId>, 284 request: RequestBody, 285 ) -> Result<AdminSuccessResponse<ResponseBody>, AdminClientError> 286 where 287 RequestBody: Serialize, 288 ResponseBody: DeserializeOwned + Serialize, 289 { 290 let envelope = ClientMutationEnvelope { 291 contract_version: ADMIN_CONTRACT_VERSION, 292 operation_id: &operation_id, 293 correlation_id: correlation_id.as_ref(), 294 request: &request, 295 }; 296 let body = encode_bounded(&envelope, self.limits.request_body_utf8_bytes() as usize) 297 .map_err(|error| match error { 298 ClientEncodingError::Limit => { 299 AdminClientError::simple(AdminClientErrorKind::RequestLimit) 300 } 301 ClientEncodingError::Encoding(error) => { 302 AdminClientError::sourced(AdminClientErrorKind::RequestEncoding, error) 303 } 304 })?; 305 let _: StrictJsonPayload = serde_json::from_slice(&body).map_err(|error| { 306 AdminClientError::sourced(AdminClientErrorKind::RequestEncoding, error) 307 })?; 308 self.execute(AdminHttpMethod::Post, target, Bytes::from(body)) 309 .await 310 } 311 312 async fn execute<T>( 313 &self, 314 method: AdminHttpMethod, 315 target: &AdminClientTarget, 316 body: Bytes, 317 ) -> Result<AdminSuccessResponse<T>, AdminClientError> 318 where 319 T: DeserializeOwned + Serialize, 320 { 321 if query_item_count(target.query()) > self.limits.query_items() as usize { 322 return Err(AdminClientError::simple(AdminClientErrorKind::QueryLimit)); 323 } 324 let deadline = self.limits.request_deadline(); 325 tokio::time::timeout(deadline, self.execute_inner(method, target, body)) 326 .await 327 .map_err(|error| AdminClientError::sourced(AdminClientErrorKind::Deadline, error))? 328 } 329 330 async fn execute_inner<T>( 331 &self, 332 method: AdminHttpMethod, 333 target: &AdminClientTarget, 334 body: Bytes, 335 ) -> Result<AdminSuccessResponse<T>, AdminClientError> 336 where 337 T: DeserializeOwned + Serialize, 338 { 339 let stream = tokio::net::UnixStream::connect(&self.socket_path) 340 .await 341 .map_err(AdminClientError::connect)?; 342 let mut connection_builder = http1::Builder::new(); 343 connection_builder 344 .max_headers(self.limits.header_count() as usize) 345 .max_buf_size((self.limits.header_bytes() as usize).max(HTTP_MINIMUM_MAX_BUFFER_SIZE)); 346 let (mut sender, connection) = connection_builder 347 .handshake(TokioIo::new(ClientUnixStream(stream))) 348 .await 349 .map_err(|error| AdminClientError::sourced(AdminClientErrorKind::Transport, error))?; 350 let driver = ConnectionDriver::new(tokio::spawn(connection)); 351 352 let http_method = match method { 353 AdminHttpMethod::Get => Method::GET, 354 AdminHttpMethod::Post => Method::POST, 355 }; 356 let mut builder = Request::builder() 357 .version(Version::HTTP_11) 358 .method(http_method) 359 .uri(target.uri.clone()) 360 .header(HOST, HeaderValue::from_static("localhost")); 361 if method == AdminHttpMethod::Post { 362 builder = builder.header(CONTENT_TYPE, HeaderValue::from_static(JSON_CONTENT_TYPE)); 363 } 364 let request = builder 365 .body(Full::new(body)) 366 .expect("validated target and static headers must build"); 367 let response = sender 368 .send_request(request) 369 .await 370 .map_err(|error| AdminClientError::sourced(AdminClientErrorKind::Transport, error))?; 371 372 if response.version() != Version::HTTP_11 { 373 return Err(AdminClientError::simple( 374 AdminClientErrorKind::ResponseHttpVersion, 375 )); 376 } 377 if header_bytes(response.headers()) > u64::from(self.limits.header_bytes()) 378 || response.headers().len() > self.limits.header_count() as usize 379 { 380 return Err(AdminClientError::simple( 381 AdminClientErrorKind::ResponseHeaders, 382 )); 383 } 384 if !is_json_content_type(response.headers()) { 385 return Err(AdminClientError::simple( 386 AdminClientErrorKind::ResponseContentType, 387 )); 388 } 389 let response_limit = self.limits.response_body_utf8_bytes() as usize; 390 if content_length(response.headers()).is_some_and(|length| length > response_limit as u64) { 391 return Err(AdminClientError::simple( 392 AdminClientErrorKind::ResponseLimit, 393 )); 394 } 395 let status = response.status(); 396 let body = Limited::new(response.into_body(), response_limit) 397 .collect() 398 .await 399 .map_err(|error| { 400 if error.downcast_ref::<LengthLimitError>().is_some() { 401 AdminClientError::simple(AdminClientErrorKind::ResponseLimit) 402 } else { 403 AdminClientError::boxed(AdminClientErrorKind::Transport, error) 404 } 405 })? 406 .to_bytes(); 407 driver.finish().await?; 408 drop(sender); 409 decode_response(status, &body) 410 } 411 } 412 413 impl fmt::Debug for AdminClient { 414 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { 415 formatter 416 .debug_struct("AdminClient") 417 .field("socket_path", &"[redacted]") 418 .field("limits", &self.limits) 419 .finish() 420 } 421 } 422 423 #[derive(Serialize)] 424 struct ClientMutationEnvelope<'a, T> { 425 contract_version: u32, 426 operation_id: &'a AdminOperationId, 427 #[serde(skip_serializing_if = "Option::is_none")] 428 correlation_id: Option<&'a AdminCorrelationId>, 429 request: &'a T, 430 } 431 432 struct ConnectionDriver { 433 handle: JoinHandle<Result<(), hyper::Error>>, 434 completed: bool, 435 } 436 437 struct ClientUnixStream(tokio::net::UnixStream); 438 439 impl AsyncRead for ClientUnixStream { 440 #[cfg_attr(coverage_nightly, coverage(off))] 441 fn poll_read( 442 mut self: Pin<&mut Self>, 443 context: &mut Context<'_>, 444 buffer: &mut ReadBuf<'_>, 445 ) -> Poll<io::Result<()>> { 446 Pin::new(&mut self.0).poll_read(context, buffer) 447 } 448 } 449 450 impl AsyncWrite for ClientUnixStream { 451 #[cfg_attr(coverage_nightly, coverage(off))] 452 fn poll_write( 453 mut self: Pin<&mut Self>, 454 context: &mut Context<'_>, 455 buffer: &[u8], 456 ) -> Poll<Result<usize, io::Error>> { 457 Pin::new(&mut self.0).poll_write(context, buffer) 458 } 459 460 #[cfg_attr(coverage_nightly, coverage(off))] 461 fn poll_flush( 462 mut self: Pin<&mut Self>, 463 context: &mut Context<'_>, 464 ) -> Poll<Result<(), io::Error>> { 465 Pin::new(&mut self.0).poll_flush(context) 466 } 467 468 fn poll_shutdown( 469 mut self: Pin<&mut Self>, 470 context: &mut Context<'_>, 471 ) -> Poll<Result<(), io::Error>> { 472 match Pin::new(&mut self.0).poll_shutdown(context) { 473 Poll::Ready(Err(error)) 474 if matches!( 475 error.kind(), 476 io::ErrorKind::BrokenPipe | io::ErrorKind::NotConnected 477 ) => 478 { 479 Poll::Ready(Ok(())) 480 } 481 result => result, 482 } 483 } 484 } 485 486 impl ConnectionDriver { 487 fn new(handle: JoinHandle<Result<(), hyper::Error>>) -> Self { 488 Self { 489 handle, 490 completed: false, 491 } 492 } 493 494 async fn finish(mut self) -> Result<(), AdminClientError> { 495 let result = match (&mut self.handle).await { 496 Ok(Ok(())) => Ok(()), 497 Ok(Err(error)) => Err(AdminClientError::sourced( 498 AdminClientErrorKind::Transport, 499 error, 500 )), 501 Err(error) => Err(AdminClientError::sourced( 502 AdminClientErrorKind::Transport, 503 error, 504 )), 505 }; 506 self.completed = true; 507 result 508 } 509 } 510 511 impl Drop for ConnectionDriver { 512 fn drop(&mut self) { 513 if !self.completed { 514 self.handle.abort(); 515 } 516 } 517 } 518 519 fn decode_response<T>( 520 status: StatusCode, 521 body: &[u8], 522 ) -> Result<AdminSuccessResponse<T>, AdminClientError> 523 where 524 T: DeserializeOwned + Serialize, 525 { 526 let strict = serde_json::from_slice::<StrictResponseEnvelope>(body).map_err(|error| { 527 AdminClientError::sourced(AdminClientErrorKind::MalformedResponse, error) 528 })?; 529 if strict.contract_version != u64::from(ADMIN_CONTRACT_VERSION) { 530 return Err(AdminClientError::simple( 531 AdminClientErrorKind::UnsupportedContractVersion, 532 )); 533 } 534 if status.is_success() { 535 serde_json::from_slice::<AdminSuccessResponse<T>>(body).map_err(|error| { 536 AdminClientError::sourced(AdminClientErrorKind::MalformedResponse, error) 537 }) 538 } else { 539 let failure = serde_json::from_slice::<AdminFailureResponse>(body).map_err(|error| { 540 AdminClientError::sourced(AdminClientErrorKind::MalformedResponse, error) 541 })?; 542 Err(AdminClientError::server(failure)) 543 } 544 } 545 546 fn header_bytes(headers: &HeaderMap) -> u64 { 547 headers.iter().fold(0_u64, |total, (name, value)| { 548 total 549 .saturating_add(name.as_str().len() as u64) 550 .saturating_add(value.as_bytes().len() as u64) 551 }) 552 } 553 554 fn query_item_count(query: Option<&str>) -> usize { 555 match query { 556 None | Some("") => 0, 557 Some(query) => query.split('&').count(), 558 } 559 } 560 561 fn valid_percent_encoding(value: &[u8]) -> bool { 562 let mut index = 0; 563 while index < value.len() { 564 if value[index] == b'%' { 565 let Some(high) = value.get(index + 1) else { 566 return false; 567 }; 568 let Some(low) = value.get(index + 2) else { 569 return false; 570 }; 571 if !high.is_ascii_hexdigit() || !low.is_ascii_hexdigit() { 572 return false; 573 } 574 index += 3; 575 } else { 576 index += 1; 577 } 578 } 579 true 580 } 581 582 fn is_json_content_type(headers: &HeaderMap) -> bool { 583 headers 584 .get(CONTENT_TYPE) 585 .and_then(|value| value.to_str().ok()) 586 .is_some_and(|value| { 587 value == JSON_CONTENT_TYPE || value == "application/json; charset=utf-8" 588 }) 589 } 590 591 fn content_length(headers: &HeaderMap) -> Option<u64> { 592 headers 593 .get(CONTENT_LENGTH) 594 .and_then(|value| value.to_str().ok()) 595 .and_then(|value| value.parse().ok()) 596 } 597 598 enum ClientEncodingError { 599 Limit, 600 Encoding(serde_json::Error), 601 } 602 603 struct CappedWriter { 604 bytes: Vec<u8>, 605 limit: usize, 606 exceeded: bool, 607 } 608 609 impl CappedWriter { 610 fn new(limit: usize) -> Self { 611 Self { 612 bytes: Vec::with_capacity(limit.min(4096)), 613 limit, 614 exceeded: false, 615 } 616 } 617 } 618 619 impl io::Write for CappedWriter { 620 fn write(&mut self, bytes: &[u8]) -> io::Result<usize> { 621 if bytes.len() > self.limit.saturating_sub(self.bytes.len()) { 622 self.exceeded = true; 623 return Err(io::Error::other("bounded JSON output limit exceeded")); 624 } 625 self.bytes.extend_from_slice(bytes); 626 Ok(bytes.len()) 627 } 628 629 fn flush(&mut self) -> io::Result<()> { 630 Ok(()) 631 } 632 } 633 634 fn encode_bounded<T>(value: &T, limit: usize) -> Result<Vec<u8>, ClientEncodingError> 635 where 636 T: Serialize, 637 { 638 let mut writer = CappedWriter::new(limit); 639 match serde_json::to_writer(&mut writer, value) { 640 Ok(()) => Ok(writer.bytes), 641 Err(_) if writer.exceeded => Err(ClientEncodingError::Limit), 642 Err(error) => Err(ClientEncodingError::Encoding(error)), 643 } 644 } 645 646 #[derive(Clone, Copy)] 647 struct StrictJsonPayload; 648 649 impl<'de> Deserialize<'de> for StrictJsonPayload { 650 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error> 651 where 652 D: Deserializer<'de>, 653 { 654 deserializer.deserialize_any(StrictJsonVisitor) 655 } 656 } 657 658 struct StrictJsonVisitor; 659 660 impl<'de> de::Visitor<'de> for StrictJsonVisitor { 661 type Value = StrictJsonPayload; 662 663 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { 664 formatter.write_str("JSON without duplicate object keys or null values") 665 } 666 667 fn visit_bool<E>(self, _value: bool) -> Result<Self::Value, E> { 668 Ok(StrictJsonPayload) 669 } 670 671 fn visit_i64<E>(self, _value: i64) -> Result<Self::Value, E> { 672 Ok(StrictJsonPayload) 673 } 674 675 fn visit_u64<E>(self, _value: u64) -> Result<Self::Value, E> { 676 Ok(StrictJsonPayload) 677 } 678 679 fn visit_f64<E>(self, value: f64) -> Result<Self::Value, E> 680 where 681 E: de::Error, 682 { 683 if value.is_finite() { 684 Ok(StrictJsonPayload) 685 } else { 686 Err(E::custom("non-finite JSON number")) 687 } 688 } 689 690 fn visit_str<E>(self, _value: &str) -> Result<Self::Value, E> { 691 Ok(StrictJsonPayload) 692 } 693 694 fn visit_borrowed_str<E>(self, value: &'de str) -> Result<Self::Value, E> 695 where 696 E: de::Error, 697 { 698 self.visit_str(value) 699 } 700 701 fn visit_string<E>(self, _value: String) -> Result<Self::Value, E> { 702 Ok(StrictJsonPayload) 703 } 704 705 fn visit_none<E>(self) -> Result<Self::Value, E> 706 where 707 E: de::Error, 708 { 709 Err(E::custom("JSON null is forbidden")) 710 } 711 712 fn visit_unit<E>(self) -> Result<Self::Value, E> 713 where 714 E: de::Error, 715 { 716 Err(E::custom("JSON null is forbidden")) 717 } 718 719 fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error> 720 where 721 A: de::SeqAccess<'de>, 722 { 723 while sequence.next_element::<StrictJsonPayload>()?.is_some() {} 724 Ok(StrictJsonPayload) 725 } 726 727 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error> 728 where 729 A: de::MapAccess<'de>, 730 { 731 let mut keys = BTreeSet::new(); 732 while let Some(key) = map.next_key::<String>()? { 733 if !keys.insert(key) { 734 return Err(de::Error::custom("duplicate JSON object key")); 735 } 736 map.next_value::<StrictJsonPayload>()?; 737 } 738 Ok(StrictJsonPayload) 739 } 740 } 741 742 struct StrictResponseEnvelope { 743 contract_version: u64, 744 } 745 746 impl<'de> Deserialize<'de> for StrictResponseEnvelope { 747 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error> 748 where 749 D: Deserializer<'de>, 750 { 751 deserializer.deserialize_map(StrictResponseVisitor) 752 } 753 } 754 755 struct StrictResponseVisitor; 756 757 impl<'de> de::Visitor<'de> for StrictResponseVisitor { 758 type Value = StrictResponseEnvelope; 759 760 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { 761 formatter.write_str("an admin response envelope without duplicate keys or null values") 762 } 763 764 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error> 765 where 766 A: de::MapAccess<'de>, 767 { 768 let mut keys = BTreeSet::new(); 769 let mut contract_version = None; 770 while let Some(key) = map.next_key::<String>()? { 771 if !keys.insert(key.clone()) { 772 return Err(de::Error::custom("duplicate JSON object key")); 773 } 774 if key == "contract_version" { 775 contract_version = Some(map.next_value::<u64>()?); 776 } else { 777 map.next_value::<StrictJsonPayload>()?; 778 } 779 } 780 Ok(StrictResponseEnvelope { 781 contract_version: contract_version 782 .ok_or_else(|| de::Error::missing_field("contract_version"))?, 783 }) 784 } 785 } 786 787 #[cfg(test)] 788 mod tests { 789 use super::*; 790 use crate::{ 791 AdminError, AdminErrorCode, AdminErrorMessage, AdminMutationRequest, AdminRouteFailure, 792 AdminRouteFailureStatus, AdminRouteOutcome, AdminRouter, AdminServer, CancellationToken, 793 EntropyError, EntropySource, UnixAdminSocketBinding, UnixAdminSocketWriterAuthority, 794 }; 795 use serde::{Deserialize, Serialize}; 796 use std::sync::Arc; 797 use tokio::io::{AsyncReadExt, AsyncWriteExt}; 798 use tokio::sync::Notify; 799 800 #[derive(Clone, Copy)] 801 struct FixedEntropy; 802 803 impl EntropySource for FixedEntropy { 804 fn fill_bytes(&self, destination: &mut [u8]) -> Result<(), EntropyError> { 805 destination.fill(0x44); 806 Ok(()) 807 } 808 } 809 810 #[derive(Debug, Deserialize, Serialize)] 811 #[serde(deny_unknown_fields)] 812 struct EchoRequest { 813 value: String, 814 } 815 816 #[derive(Debug, Deserialize, Serialize, PartialEq, Eq)] 817 #[serde(deny_unknown_fields)] 818 struct EchoResponse { 819 value: String, 820 } 821 822 fn known_error(code: &'static str, message: &'static str) -> AdminError { 823 AdminError::new( 824 AdminErrorCode::new(code).expect("error code"), 825 AdminErrorMessage::new(message).expect("error message"), 826 ) 827 } 828 829 fn echo_router() -> AdminRouter { 830 let mut router = AdminRouter::new(); 831 router 832 .route(AdminHttpMethod::Post, "/v1/echo", |request| async move { 833 match request.decode_json::<AdminMutationRequest<EchoRequest>>() { 834 Ok(envelope) => request 835 .success(&EchoResponse { 836 value: envelope.into_request().value, 837 }) 838 .expect("echo response"), 839 Err(_) => AdminRouteOutcome::failure(AdminRouteFailure::new( 840 AdminRouteFailureStatus::BadRequest, 841 known_error("invalid_echo", "echo request is invalid"), 842 )), 843 } 844 }) 845 .expect("echo route"); 846 router 847 } 848 849 async fn binding(directory: &tempfile::TempDir) -> (PathBuf, UnixAdminSocketBinding) { 850 let socket = directory.path().join("admin.sock"); 851 let authority = 852 UnixAdminSocketWriterAuthority::acquire(directory.path()).expect("writer authority"); 853 let binding = UnixAdminSocketBinding::bind(authority, &socket) 854 .await 855 .expect("socket binding"); 856 (socket, binding) 857 } 858 859 async fn fake_server( 860 directory: &tempfile::TempDir, 861 name: &str, 862 response: Vec<u8>, 863 delay: std::time::Duration, 864 ) -> (PathBuf, JoinHandle<()>) { 865 let socket = directory.path().join(name); 866 let listener = tokio::net::UnixListener::bind(&socket).expect("fake listener"); 867 let task = tokio::spawn(async move { 868 let (mut stream, _) = listener.accept().await.expect("fake accept"); 869 let mut request = [0_u8; 4096]; 870 let _ = stream.read(&mut request).await; 871 tokio::time::sleep(delay).await; 872 let _ = stream.write_all(&response).await; 873 let _ = stream.shutdown().await; 874 }); 875 (socket, task) 876 } 877 878 fn raw_response(version: &str, status: &str, body: &str) -> Vec<u8> { 879 format!( 880 "{version} {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", 881 body.len() 882 ) 883 .into_bytes() 884 } 885 886 fn limits_with( 887 request_body: u32, 888 response_body: u32, 889 query_items: u32, 890 deadline: std::time::Duration, 891 ) -> AdminTransportLimits { 892 let mut values = AdminTransportLimits::DEFAULT.values(); 893 values.request_body_utf8_bytes = request_body; 894 values.response_body_utf8_bytes = response_body; 895 values.query_items = query_items; 896 values.request_deadline = deadline; 897 AdminTransportLimits::new(values).expect("client limits") 898 } 899 900 #[tokio::test] 901 async fn server_client_round_trip_preserves_version_and_correlation() { 902 let directory = super::test_support::short_tempdir(); 903 let (socket, binding) = binding(&directory).await; 904 let server = AdminServer::new(echo_router(), AdminTransportLimits::DEFAULT, FixedEntropy) 905 .expect("admin server"); 906 let cancellation = CancellationToken::new(); 907 let server_task = tokio::spawn(server.serve(binding, cancellation.clone())); 908 let client = AdminClient::new(&socket, AdminTransportLimits::DEFAULT).expect("client"); 909 let target = AdminClientTarget::new("/v1/echo").expect("target"); 910 let correlation = AdminCorrelationId::new("client-round-trip").expect("correlation"); 911 912 let response = client 913 .mutate::<_, EchoResponse>( 914 &target, 915 AdminOperationId::new("echo-01").expect("operation"), 916 Some(correlation.clone()), 917 EchoRequest { 918 value: "bounded".to_owned(), 919 }, 920 ) 921 .await 922 .unwrap_or_else(|error| { 923 panic!( 924 "round trip: {error:?}; source={:?}", 925 error.source().map(ToString::to_string) 926 ) 927 }); 928 assert_eq!(response.correlation_id(), &correlation); 929 assert_eq!( 930 response.into_result(), 931 EchoResponse { 932 value: "bounded".to_owned() 933 } 934 ); 935 936 let missing = AdminClientTarget::new("/v1/missing").expect("missing target"); 937 let error = client 938 .get::<EchoResponse>(&missing) 939 .await 940 .expect_err("missing route"); 941 assert_eq!(error.kind(), AdminClientErrorKind::ServerFailure); 942 assert_eq!( 943 error 944 .failure() 945 .expect("server failure") 946 .error() 947 .code() 948 .as_str(), 949 "route_not_found" 950 ); 951 952 cancellation.cancel(); 953 server_task 954 .await 955 .expect("server task") 956 .expect("server shutdown"); 957 } 958 959 #[tokio::test] 960 async fn unavailable_socket_and_deadline_are_typed_and_safe() { 961 let directory = super::test_support::short_tempdir(); 962 let missing = directory.path().join("missing.sock"); 963 let client = AdminClient::new(&missing, AdminTransportLimits::DEFAULT).expect("client"); 964 let target = AdminClientTarget::new("/v1/status").expect("target"); 965 let error = client 966 .get::<EchoResponse>(&target) 967 .await 968 .expect_err("unavailable socket"); 969 assert_eq!(error.kind(), AdminClientErrorKind::Connect); 970 assert!(error.io_kind().is_some()); 971 assert!(!format!("{error:?}").contains("missing.sock")); 972 973 let body = 974 r#"{"contract_version":1,"ok":true,"correlation_id":"late","result":{"value":"late"}}"#; 975 let (socket, task) = fake_server( 976 &directory, 977 "slow.sock", 978 raw_response("HTTP/1.1", "200 OK", body), 979 std::time::Duration::from_millis(100), 980 ) 981 .await; 982 let client = AdminClient::new( 983 socket, 984 limits_with(1024, 1024, 10, std::time::Duration::from_millis(20)), 985 ) 986 .expect("slow client"); 987 let error = client 988 .get::<EchoResponse>(&target) 989 .await 990 .expect_err("deadline"); 991 assert_eq!(error.kind(), AdminClientErrorKind::Deadline); 992 task.await.expect("fake task"); 993 } 994 995 #[tokio::test] 996 async fn version_malformed_duplicate_and_oversized_responses_fail_closed() { 997 let directory = super::test_support::short_tempdir(); 998 let target = AdminClientTarget::new("/v1/status").expect("target"); 999 let cases = [ 1000 ( 1001 "version.sock", 1002 raw_response( 1003 "HTTP/1.1", 1004 "200 OK", 1005 r#"{"contract_version":2,"ok":true,"correlation_id":"future","result":{"value":"future"}}"#, 1006 ), 1007 AdminClientErrorKind::UnsupportedContractVersion, 1008 ), 1009 ( 1010 "malformed.sock", 1011 raw_response("HTTP/1.1", "200 OK", "{"), 1012 AdminClientErrorKind::MalformedResponse, 1013 ), 1014 ( 1015 "duplicate.sock", 1016 raw_response( 1017 "HTTP/1.1", 1018 "200 OK", 1019 r#"{"contract_version":1,"contract_version":1,"ok":true,"correlation_id":"duplicate","result":{"value":"duplicate"}}"#, 1020 ), 1021 AdminClientErrorKind::MalformedResponse, 1022 ), 1023 ( 1024 "null.sock", 1025 raw_response( 1026 "HTTP/1.1", 1027 "200 OK", 1028 r#"{"contract_version":1,"ok":true,"correlation_id":"null-result","result":{"value":null}}"#, 1029 ), 1030 AdminClientErrorKind::MalformedResponse, 1031 ), 1032 ( 1033 "http10.sock", 1034 raw_response( 1035 "HTTP/1.0", 1036 "200 OK", 1037 r#"{"contract_version":1,"ok":true,"correlation_id":"old-http","result":{"value":"old"}}"#, 1038 ), 1039 AdminClientErrorKind::ResponseHttpVersion, 1040 ), 1041 ]; 1042 for (name, response, expected) in cases { 1043 let (socket, task) = 1044 fake_server(&directory, name, response, std::time::Duration::ZERO).await; 1045 let client = AdminClient::new(socket, AdminTransportLimits::DEFAULT).expect("client"); 1046 let error = client 1047 .get::<EchoResponse>(&target) 1048 .await 1049 .expect_err("invalid response"); 1050 assert_eq!( 1051 error.kind(), 1052 expected, 1053 "source={:?}", 1054 error.source().map(ToString::to_string) 1055 ); 1056 task.await.expect("fake task"); 1057 } 1058 1059 let oversized = "x".repeat(513); 1060 let (socket, task) = fake_server( 1061 &directory, 1062 "oversized.sock", 1063 raw_response("HTTP/1.1", "200 OK", &oversized), 1064 std::time::Duration::ZERO, 1065 ) 1066 .await; 1067 let client = AdminClient::new( 1068 socket, 1069 limits_with(1024, 512, 10, std::time::Duration::from_secs(1)), 1070 ) 1071 .expect("client"); 1072 let error = client 1073 .get::<EchoResponse>(&target) 1074 .await 1075 .expect_err("oversized response"); 1076 assert_eq!(error.kind(), AdminClientErrorKind::ResponseLimit); 1077 task.await.expect("fake task"); 1078 1079 let (socket, task) = fake_server( 1080 &directory, 1081 "bad.sock", 1082 b"not-http\r\n\r\n".to_vec(), 1083 std::time::Duration::ZERO, 1084 ) 1085 .await; 1086 let client = AdminClient::new(socket, AdminTransportLimits::DEFAULT).expect("client"); 1087 let error = client 1088 .get::<EchoResponse>(&target) 1089 .await 1090 .expect_err("invalid HTTP response"); 1091 assert_eq!(error.kind(), AdminClientErrorKind::Transport); 1092 task.await.expect("fake task"); 1093 } 1094 1095 #[tokio::test] 1096 async fn request_and_query_limits_fail_before_socket_access() { 1097 let directory = super::test_support::short_tempdir(); 1098 let missing = directory.path().join("missing.sock"); 1099 let client = AdminClient::new( 1100 &missing, 1101 limits_with(32, 512, 1, std::time::Duration::from_secs(1)), 1102 ) 1103 .expect("client"); 1104 let target = AdminClientTarget::new("/v1/mutate").expect("target"); 1105 let error = client 1106 .mutate::<_, EchoResponse>( 1107 &target, 1108 AdminOperationId::new("oversized-01").expect("operation"), 1109 None, 1110 EchoRequest { 1111 value: "x".repeat(128), 1112 }, 1113 ) 1114 .await 1115 .expect_err("request limit"); 1116 assert_eq!(error.kind(), AdminClientErrorKind::RequestLimit); 1117 1118 let null_client = 1119 AdminClient::new(&missing, AdminTransportLimits::DEFAULT).expect("null client"); 1120 let error = null_client 1121 .mutate::<_, EchoResponse>( 1122 &target, 1123 AdminOperationId::new("null-01").expect("operation"), 1124 None, 1125 Option::<EchoRequest>::None, 1126 ) 1127 .await 1128 .expect_err("null request"); 1129 assert_eq!(error.kind(), AdminClientErrorKind::RequestEncoding); 1130 1131 let query = AdminClientTarget::new("/v1/status?a=1&b=2").expect("query target"); 1132 let error = client 1133 .get::<EchoResponse>(&query) 1134 .await 1135 .expect_err("query limit"); 1136 assert_eq!(error.kind(), AdminClientErrorKind::QueryLimit); 1137 } 1138 1139 #[test] 1140 fn target_client_and_error_debug_bound_and_redact_authority() { 1141 for invalid in [ 1142 "", 1143 "http://localhost/v1/status", 1144 "/v2/status", 1145 "/v1/{route}", 1146 "/v1/status/", 1147 "/v1/items/%2", 1148 ] { 1149 assert!(AdminClientTarget::new(invalid).is_err(), "{invalid}"); 1150 } 1151 let target = 1152 AdminClientTarget::new("/v1/items/secret-id?token=protected").expect("redacted target"); 1153 let debug = format!("{target:?}"); 1154 assert!(!debug.contains("secret-id")); 1155 assert!(!debug.contains("protected")); 1156 1157 let directory = super::test_support::short_tempdir(); 1158 let socket = directory.path().join("protected-admin.sock"); 1159 let client = AdminClient::new(&socket, AdminTransportLimits::DEFAULT).expect("client"); 1160 assert!(!format!("{client:?}").contains("protected-admin.sock")); 1161 assert_eq!(client.socket_path(), socket); 1162 assert_eq!(client.limits(), AdminTransportLimits::DEFAULT); 1163 assert_eq!(target.path(), "/v1/items/secret-id"); 1164 assert_eq!(target.query(), Some("token=protected")); 1165 } 1166 1167 #[test] 1168 fn client_source_has_no_runtime_tcp_or_process_authority() { 1169 let source = include_str!("client.rs"); 1170 for forbidden in [ 1171 concat!("Tcp", "Stream"), 1172 concat!("Runtime", "::new"), 1173 concat!("process", "::exit"), 1174 concat!("tokio", "::signal"), 1175 ] { 1176 assert!( 1177 !source.contains(forbidden), 1178 "forbidden client authority: {forbidden}" 1179 ); 1180 } 1181 assert!(source.contains("UnixStream::connect")); 1182 } 1183 1184 #[test] 1185 fn shared_client_handles_are_thread_safe() { 1186 fn assert_send_sync<T: Send + Sync>() {} 1187 assert_send_sync::<AdminClient>(); 1188 assert_send_sync::<Arc<AdminClient>>(); 1189 } 1190 1191 #[tokio::test] 1192 async fn cancelling_driver_finish_aborts_and_drops_the_connection_task() { 1193 struct DropNotify(Arc<Notify>); 1194 1195 impl Drop for DropNotify { 1196 fn drop(&mut self) { 1197 self.0.notify_one(); 1198 } 1199 } 1200 1201 let dropped = Arc::new(Notify::new()); 1202 let task_dropped = Arc::clone(&dropped); 1203 let handle = tokio::spawn(async move { 1204 let _drop_notify = DropNotify(task_dropped); 1205 std::future::pending::<()>().await; 1206 Ok::<(), hyper::Error>(()) 1207 }); 1208 let driver = ConnectionDriver::new(handle); 1209 assert!( 1210 tokio::time::timeout(std::time::Duration::from_millis(10), driver.finish()) 1211 .await 1212 .is_err() 1213 ); 1214 tokio::time::timeout(std::time::Duration::from_secs(1), dropped.notified()) 1215 .await 1216 .expect("aborted driver task must drop"); 1217 } 1218 1219 #[test] 1220 fn strict_response_target_and_error_helpers_cover_the_full_value_surface() { 1221 for document in [ 1222 "true", 1223 "-3", 1224 "4", 1225 "2.5", 1226 r#""text""#, 1227 "[true,2]", 1228 r#"{"value":3}"#, 1229 ] { 1230 serde_json::from_str::<StrictJsonPayload>(document).unwrap(); 1231 } 1232 for rejected in ["null", "[1,null]", r#"{"same":1,"same":2}"#] { 1233 assert!(serde_json::from_str::<StrictJsonPayload>(rejected).is_err()); 1234 } 1235 1236 let target = AdminClientTarget::new("/v1/items/value%2D1?page=1&limit=2").unwrap(); 1237 assert_eq!(target.as_str(), "/v1/items/value%2D1?page=1&limit=2"); 1238 assert_eq!(target.path(), "/v1/items/value%2D1"); 1239 assert_eq!(target.query(), Some("page=1&limit=2")); 1240 assert_eq!(query_item_count(None), 0); 1241 assert_eq!(query_item_count(Some("")), 0); 1242 assert_eq!(query_item_count(target.query()), 2); 1243 assert!(valid_percent_encoding(b"/v1/items/%2d")); 1244 assert!(!valid_percent_encoding(b"/v1/items/%")); 1245 assert!(!valid_percent_encoding(b"/v1/items/%2")); 1246 assert!(!valid_percent_encoding(b"/v1/items/%GG")); 1247 assert!(!valid_percent_encoding(b"/v1/items/%G0")); 1248 assert!(!valid_percent_encoding(b"/v1/items/%0G")); 1249 1250 let invalid_targets = [ 1251 (String::new(), AdminClientTargetError::Empty), 1252 ( 1253 "x".repeat(ADMIN_CLIENT_TARGET_MAX_UTF8_BYTES + 1), 1254 AdminClientTargetError::TooLong, 1255 ), 1256 ( 1257 "http://[invalid".to_owned(), 1258 AdminClientTargetError::InvalidUri, 1259 ), 1260 ( 1261 "http://localhost/v1/status".to_owned(), 1262 AdminClientTargetError::AuthorityForbidden, 1263 ), 1264 ( 1265 "/v2/status".to_owned(), 1266 AdminClientTargetError::WrongVersionPrefix, 1267 ), 1268 ( 1269 format!("/v1/{}", "x".repeat(ADMIN_ROUTE_PATH_MAX_UTF8_BYTES)), 1270 AdminClientTargetError::PathTooLong, 1271 ), 1272 ( 1273 "/v1//status".to_owned(), 1274 AdminClientTargetError::EmptySegment, 1275 ), 1276 ( 1277 "/v1/{status}".to_owned(), 1278 AdminClientTargetError::PatternForbidden, 1279 ), 1280 ( 1281 "/v1/items/%GG".to_owned(), 1282 AdminClientTargetError::InvalidPercentEncoding, 1283 ), 1284 ]; 1285 for (target, expected) in invalid_targets { 1286 assert_eq!(AdminClientTarget::new(target).unwrap_err(), expected); 1287 } 1288 assert!(!AdminClientTargetError::Empty.to_string().is_empty()); 1289 assert_eq!( 1290 AdminClient::new("relative.sock", AdminTransportLimits::DEFAULT) 1291 .unwrap_err() 1292 .kind(), 1293 AdminClientErrorKind::SocketPath 1294 ); 1295 assert_eq!( 1296 AdminClient::new("/", AdminTransportLimits::DEFAULT) 1297 .unwrap_err() 1298 .kind(), 1299 AdminClientErrorKind::SocketPath 1300 ); 1301 1302 let mut headers = HeaderMap::new(); 1303 assert!(!is_json_content_type(&headers)); 1304 assert_eq!(content_length(&headers), None); 1305 headers.insert(CONTENT_TYPE, HeaderValue::from_static(JSON_CONTENT_TYPE)); 1306 headers.insert(CONTENT_LENGTH, HeaderValue::from_static("12")); 1307 assert!(is_json_content_type(&headers)); 1308 assert_eq!(content_length(&headers), Some(12)); 1309 headers.insert( 1310 CONTENT_TYPE, 1311 HeaderValue::from_static("application/json; charset=utf-8"), 1312 ); 1313 headers.insert(CONTENT_LENGTH, HeaderValue::from_static("invalid")); 1314 assert!(is_json_content_type(&headers)); 1315 assert_eq!(content_length(&headers), None); 1316 assert!(header_bytes(&headers) > 0); 1317 1318 let failure = AdminFailureResponse::new( 1319 AdminCorrelationId::new("safe-correlation").unwrap(), 1320 known_error("known_failure", "known failure"), 1321 ); 1322 let server = AdminClientError::server(failure.clone()); 1323 assert_eq!(server.kind(), AdminClientErrorKind::ServerFailure); 1324 assert_eq!(server.failure(), Some(&failure)); 1325 assert_eq!(server.io_kind(), None); 1326 assert!(server.source().is_none()); 1327 assert!(!server.to_string().is_empty()); 1328 1329 let connect = AdminClientError::connect(io::Error::new( 1330 io::ErrorKind::ConnectionRefused, 1331 "sensitive socket", 1332 )); 1333 assert_eq!(connect.kind(), AdminClientErrorKind::Connect); 1334 assert_eq!(connect.io_kind(), Some(io::ErrorKind::ConnectionRefused)); 1335 assert!(connect.source().is_some()); 1336 assert!(!format!("{connect:?}").contains("sensitive socket")); 1337 1338 let malformed = decode_response::<EchoResponse>(StatusCode::OK, b"[]").unwrap_err(); 1339 assert_eq!(malformed.kind(), AdminClientErrorKind::MalformedResponse); 1340 let unsupported = decode_response::<EchoResponse>( 1341 StatusCode::OK, 1342 br#"{"contract_version":2,"ok":true,"correlation_id":"safe","result":{"value":"x"}}"#, 1343 ) 1344 .unwrap_err(); 1345 assert_eq!( 1346 unsupported.kind(), 1347 AdminClientErrorKind::UnsupportedContractVersion 1348 ); 1349 1350 use std::io::Write as _; 1351 let mut writer = CappedWriter::new(4); 1352 writer.flush().unwrap(); 1353 assert_eq!(writer.write(b"four").unwrap(), 4); 1354 assert!(writer.write(b"x").is_err()); 1355 } 1356 }