1#![warn(missing_docs)]
7
8use aws_smithy_runtime_api::client::dns::{DnsFuture, ResolveDns, ResolveDnsError};
40use std::collections::{HashMap, VecDeque};
41use std::error::Error;
42use std::fmt;
43use std::fmt::Write as _;
44use std::net::{IpAddr, SocketAddr};
45use std::sync::atomic::{AtomicU64, Ordering};
46use std::sync::{Arc, Mutex};
47use std::time::Duration;
48use tokio::io::{AsyncReadExt, AsyncWriteExt};
49use tokio::net::{TcpListener, TcpStream};
50use tokio::sync::watch;
51use tokio::task::{JoinHandle, JoinSet};
52
53const MAX_HTTP1_HEADER_BYTES: usize = 64 * 1024;
54const MAX_HTTP1_BODY_BYTES: usize = 8 * 1024 * 1024;
55const READ_CHUNK_SIZE: usize = 8 * 1024;
56
57#[derive(Clone, Debug, Eq, PartialEq)]
59pub struct HarnessError {
60 message: Arc<str>,
61}
62
63impl HarnessError {
64 fn new(message: impl Into<String>) -> Self {
65 Self {
66 message: Arc::from(message.into()),
67 }
68 }
69}
70
71impl fmt::Display for HarnessError {
72 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
73 f.write_str(&self.message)
74 }
75}
76
77impl Error for HarnessError {}
78
79#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
81pub struct ConnectionId(u64);
82
83impl ConnectionId {
84 pub fn as_u64(self) -> u64 {
86 self.0
87 }
88}
89
90impl fmt::Display for ConnectionId {
91 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
92 self.0.fmt(f)
93 }
94}
95
96#[non_exhaustive]
98#[derive(Clone, Copy, Debug, Eq, PartialEq)]
99pub enum ConnectionCloseReason {
100 ClientClosed,
102 ScriptCompleted,
104 Reset,
106 HarnessShutdown,
108 ScriptFailed,
110 ScriptedTransportAbort,
112}
113
114#[non_exhaustive]
116#[derive(Clone, Debug, Eq, PartialEq)]
117pub enum ConnectionEvent {
118 DnsLookup {
120 hostname: String,
122 },
123 TcpAccepted {
125 connection_id: ConnectionId,
127 endpoint_addr: SocketAddr,
129 },
130 Http1Request {
132 connection_id: ConnectionId,
134 endpoint_addr: SocketAddr,
136 method: String,
138 target: String,
140 host: Option<String>,
142 },
143 ConnectionClosed {
145 connection_id: ConnectionId,
147 reason: ConnectionCloseReason,
149 },
150}
151
152#[derive(Debug)]
153struct RecordedState {
154 events: Vec<ConnectionEvent>,
155 failures: Vec<HarnessError>,
156 generation: u64,
157}
158
159#[derive(Debug)]
162struct SharedState {
163 recorded: Mutex<RecordedState>,
164 changed: watch::Sender<u64>,
165}
166
167impl SharedState {
168 fn new() -> Self {
169 let (changed, _) = watch::channel(0);
170 Self {
171 recorded: Mutex::new(RecordedState {
172 events: Vec::new(),
173 failures: Vec::new(),
174 generation: 0,
175 }),
176 changed,
177 }
178 }
179
180 fn record_event(&self, event: ConnectionEvent) {
181 let generation = {
182 let mut state = self.recorded.lock().unwrap_or_else(|err| err.into_inner());
183 state.events.push(event);
184 state.generation += 1;
185 state.generation
186 };
187 self.changed.send_replace(generation);
188 }
189
190 fn record_failure(&self, failure: HarnessError) {
191 let generation = {
192 let mut state = self.recorded.lock().unwrap_or_else(|err| err.into_inner());
193 state.failures.push(failure);
194 state.generation += 1;
195 state.generation
196 };
197 self.changed.send_replace(generation);
198 }
199
200 fn events(&self) -> Vec<ConnectionEvent> {
201 self.recorded
202 .lock()
203 .unwrap_or_else(|err| err.into_inner())
204 .events
205 .clone()
206 }
207
208 fn failure(&self) -> Option<HarnessError> {
209 let state = self.recorded.lock().unwrap_or_else(|err| err.into_inner());
210 match state.failures.as_slice() {
211 [] => None,
212 [failure] => Some(failure.clone()),
213 failures => Some(HarnessError::new(format!(
214 "{} harness failures: {}",
215 failures.len(),
216 failures
217 .iter()
218 .map(ToString::to_string)
219 .collect::<Vec<_>>()
220 .join("; ")
221 ))),
222 }
223 }
224
225 async fn wait_for<F>(
226 &self,
227 description: &str,
228 timeout: Duration,
229 predicate: F,
230 ) -> Result<(), HarnessError>
231 where
232 F: Fn(&[ConnectionEvent]) -> bool,
233 {
234 let mut changed = self.changed.subscribe();
235 let wait = async {
236 loop {
237 {
238 let state = self.recorded.lock().unwrap_or_else(|err| err.into_inner());
239 if let Some(failure) = state.failures.first() {
240 return Err(failure.clone());
241 }
242 if predicate(&state.events) {
243 return Ok(());
244 }
245 }
246 changed.changed().await.map_err(|_| {
247 HarnessError::new(format!(
248 "event notification closed while waiting for {description}"
249 ))
250 })?;
251 }
252 };
253
254 tokio::time::timeout(timeout, wait).await.map_err(|_| {
255 HarnessError::new(format!(
256 "timed out after {timeout:?} waiting for {description}"
257 ))
258 })?
259 }
260}
261
262#[derive(Clone, Debug)]
267pub struct ManualGate {
268 state: Arc<GateState>,
269}
270
271#[derive(Debug)]
272struct GateState {
273 snapshot: watch::Sender<GateSnapshot>,
274}
275
276#[derive(Clone, Copy, Debug)]
277struct GateSnapshot {
278 arrivals: usize,
279 released: bool,
280}
281
282impl ManualGate {
283 pub fn new() -> Self {
285 let (snapshot, _) = watch::channel(GateSnapshot {
286 arrivals: 0,
287 released: false,
288 });
289 Self {
290 state: Arc::new(GateState { snapshot }),
291 }
292 }
293
294 pub fn waiter(&self) -> GateWaiter {
296 GateWaiter {
297 state: self.state.clone(),
298 }
299 }
300
301 pub fn arrivals(&self) -> usize {
305 self.state.snapshot.borrow().arrivals
306 }
307
308 pub async fn wait_until_reached(&self, timeout: Duration) -> Result<(), HarnessError> {
310 self.wait_for_arrivals(1, timeout).await
311 }
312
313 pub async fn wait_for_arrivals(
315 &self,
316 expected: usize,
317 timeout: Duration,
318 ) -> Result<(), HarnessError> {
319 let mut snapshot = self.state.snapshot.subscribe();
320 let wait = async {
321 loop {
322 if snapshot.borrow().arrivals >= expected {
323 return Ok(());
324 }
325 snapshot.changed().await.map_err(|_| {
326 HarnessError::new("gate notification closed while waiting for arrivals")
327 })?;
328 }
329 };
330
331 tokio::time::timeout(timeout, wait).await.map_err(|_| {
332 HarnessError::new(format!(
333 "timed out after {timeout:?} waiting for {expected} gate arrivals; observed {}",
334 self.arrivals()
335 ))
336 })?
337 }
338
339 pub fn release(&self) {
341 self.state
342 .snapshot
343 .send_modify(|snapshot| snapshot.released = true);
344 }
345}
346
347impl Default for ManualGate {
348 fn default() -> Self {
349 Self::new()
350 }
351}
352
353#[derive(Clone, Debug)]
358pub struct GateWaiter {
359 state: Arc<GateState>,
360}
361
362impl GateWaiter {
363 pub async fn wait(&self) -> Result<(), HarnessError> {
369 let mut snapshot = self.state.snapshot.subscribe();
370 self.state
371 .snapshot
372 .send_modify(|snapshot| snapshot.arrivals += 1);
373 loop {
374 if snapshot.borrow().released {
375 return Ok(());
376 }
377 snapshot
378 .changed()
379 .await
380 .map_err(|_| HarnessError::new("gate notification closed before release"))?;
381 }
382 }
383}
384
385#[non_exhaustive]
387#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
388pub enum Finish {
389 #[default]
391 AwaitClientClose,
392 Close,
394 Reset,
396}
397
398#[derive(Clone, Debug)]
404pub struct BodyPlan {
405 parts: Vec<BodyPart>,
406 length: usize,
407}
408
409#[derive(Clone, Debug)]
410enum BodyPart {
411 Bytes(Vec<u8>),
412 Wait(GateWaiter),
413}
414
415impl BodyPlan {
416 pub fn complete(body: impl AsRef<[u8]>) -> Self {
418 let body = body.as_ref().to_vec();
419 Self {
420 length: body.len(),
421 parts: vec![BodyPart::Bytes(body)],
422 }
423 }
424
425 pub fn split_at_gate(
427 before: impl AsRef<[u8]>,
428 gate: GateWaiter,
429 after: impl AsRef<[u8]>,
430 ) -> Self {
431 let before = before.as_ref().to_vec();
432 let after = after.as_ref().to_vec();
433 Self {
434 length: before.len() + after.len(),
435 parts: vec![
436 BodyPart::Bytes(before),
437 BodyPart::Wait(gate),
438 BodyPart::Bytes(after),
439 ],
440 }
441 }
442}
443
444impl Default for BodyPlan {
445 fn default() -> Self {
446 Self::complete([])
447 }
448}
449
450#[derive(Clone, Debug)]
452pub struct Http1Response {
453 status: u16,
454 headers: Vec<(String, String)>,
455 body: BodyPlan,
456 close: bool,
457}
458
459impl Http1Response {
460 pub fn ok() -> Self {
462 Self::new(200)
463 }
464
465 pub fn new(status: u16) -> Self {
467 Self {
468 status,
469 headers: Vec::new(),
470 body: BodyPlan::default(),
471 close: false,
472 }
473 }
474
475 pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
477 self.headers.push((name.into(), value.into()));
478 self
479 }
480
481 pub fn body(mut self, body: impl AsRef<[u8]>) -> Self {
483 self.body = BodyPlan::complete(body);
484 self
485 }
486
487 pub fn body_plan(mut self, body: BodyPlan) -> Self {
489 self.body = body;
490 self
491 }
492
493 pub fn connection_close(mut self) -> Self {
495 self.close = true;
496 self
497 }
498
499 fn validate(&self) -> Result<(), HarnessError> {
500 http_1x::StatusCode::from_u16(self.status)
501 .map_err(|_| HarnessError::new(format!("invalid HTTP status {}", self.status)))?;
502 for (name, value) in &self.headers {
503 if name.is_empty() || name.contains(['\r', '\n', ':']) || value.contains(['\r', '\n']) {
504 return Err(HarnessError::new(format!(
505 "invalid HTTP response header {name:?}: {value:?}"
506 )));
507 }
508 if name.eq_ignore_ascii_case("content-length")
509 || name.eq_ignore_ascii_case("connection")
510 {
511 return Err(HarnessError::new(format!(
512 "{name} is managed by Http1Response; use SocketScript for raw framing"
513 )));
514 }
515 }
516 Ok(())
517 }
518
519 fn actions(&self) -> Vec<Action> {
520 let reason = http_1x::StatusCode::from_u16(self.status)
521 .ok()
522 .and_then(|code| code.canonical_reason())
523 .unwrap_or("Response");
524 let mut head = String::new();
525 let _ = write!(
526 head,
527 "HTTP/1.1 {} {}\r\nContent-Length: {}\r\nConnection: {}\r\n",
528 self.status,
529 reason,
530 self.body.length,
531 if self.close { "close" } else { "keep-alive" }
532 );
533 for (name, value) in &self.headers {
534 let _ = write!(head, "{name}: {value}\r\n");
535 }
536 head.push_str("\r\n");
537
538 let mut actions = vec![Action::WriteAll(head.into_bytes())];
539 for part in &self.body.parts {
540 match part {
541 BodyPart::Bytes(bytes) if !bytes.is_empty() => {
542 actions.push(Action::WriteAll(bytes.clone()));
543 }
544 BodyPart::Bytes(_) => {}
545 BodyPart::Wait(waiter) => actions.push(Action::Wait(waiter.clone())),
546 }
547 }
548 if self.close {
549 actions.push(Action::Close);
550 }
551 actions
552 }
553}
554
555#[derive(Clone, Debug)]
563pub struct Http1Script {
564 responses: Http1Responses,
565 finish: Finish,
566}
567
568#[derive(Clone, Debug)]
569enum Http1Responses {
570 Finite(Vec<Http1Response>),
571 Repeated(Http1Response),
572}
573
574impl Http1Script {
575 pub fn new() -> Self {
580 Self {
581 responses: Http1Responses::Finite(Vec::new()),
582 finish: Finish::default(),
583 }
584 }
585
586 pub fn responses<I>(responses: I) -> Self
590 where
591 I: IntoIterator<Item = Http1Response>,
592 {
593 Self {
594 responses: Http1Responses::Finite(responses.into_iter().collect()),
595 finish: Finish::default(),
596 }
597 }
598
599 pub fn serve(response: Http1Response) -> Self {
601 Self {
602 responses: Http1Responses::Repeated(response),
603 finish: Finish::default(),
604 }
605 }
606
607 pub fn respond(mut self, response: Http1Response) -> Self {
614 match &mut self.responses {
615 Http1Responses::Finite(responses) => responses.push(response),
616 Http1Responses::Repeated(_) => panic!(
617 "cannot append a response to a repeating Http1Script (created with Http1Script::serve)"
618 ),
619 }
620 self
621 }
622
623 pub fn finish(mut self, finish: Finish) -> Self {
630 assert!(
631 !matches!(&self.responses, Http1Responses::Repeated(_)),
632 "cannot set a finite finish policy on a repeating Http1Script (created with Http1Script::serve)"
633 );
634 self.finish = finish;
635 self
636 }
637
638 fn validate(&self) -> Result<(), HarnessError> {
639 match &self.responses {
640 Http1Responses::Finite(responses) => {
641 for (index, response) in responses.iter().enumerate() {
642 response.validate()?;
643 if response.close && index + 1 != responses.len() {
644 return Err(HarnessError::new(
645 "a connection-closing response must be the final response",
646 ));
647 }
648 }
649 if responses.last().is_some_and(|response| response.close)
650 && self.finish != Finish::AwaitClientClose
651 {
652 return Err(HarnessError::new(
653 "a connection-closing response cannot also have a finish policy",
654 ));
655 }
656 }
657 Http1Responses::Repeated(response) => {
658 response.validate()?;
659 }
660 }
661 Ok(())
662 }
663}
664
665impl Default for Http1Script {
666 fn default() -> Self {
667 Self::new()
668 }
669}
670
671#[derive(Clone, Debug, Default)]
678pub struct SocketScript {
679 actions: Vec<Action>,
680}
681
682impl SocketScript {
683 pub fn new() -> Self {
685 Self::default()
686 }
687
688 pub fn read_http1_request(mut self) -> Self {
693 self.actions.push(Action::ReadHttp1Request);
694 self
695 }
696
697 pub fn read_until(mut self, delimiter: impl AsRef<[u8]>, limit: usize) -> Self {
701 self.actions.push(Action::ReadUntil {
702 delimiter: delimiter.as_ref().to_vec(),
703 limit,
704 });
705 self
706 }
707
708 pub fn read_exact(mut self, length: usize) -> Self {
710 self.actions.push(Action::ReadExact(length));
711 self
712 }
713
714 pub fn expect_bytes(mut self, expected: impl AsRef<[u8]>) -> Self {
716 self.actions
717 .push(Action::ExpectBytes(expected.as_ref().to_vec()));
718 self
719 }
720
721 pub fn write_all(mut self, bytes: impl AsRef<[u8]>) -> Self {
723 self.actions.push(Action::WriteAll(bytes.as_ref().to_vec()));
724 self
725 }
726
727 pub fn wait(mut self, gate: GateWaiter) -> Self {
729 self.actions.push(Action::Wait(gate));
730 self
731 }
732
733 pub fn delay(mut self, duration: Duration) -> Self {
737 self.actions.push(Action::Delay(duration));
738 self
739 }
740
741 pub fn shutdown_write(mut self) -> Self {
743 self.actions.push(Action::ShutdownWrite);
744 self
745 }
746
747 pub fn await_client_close(mut self) -> Self {
753 self.actions.push(Action::AwaitClientClose);
754 self
755 }
756
757 pub fn close(mut self) -> Self {
759 self.actions.push(Action::Close);
760 self
761 }
762
763 pub fn reset(mut self) -> Self {
765 self.actions.push(Action::Reset);
766 self
767 }
768
769 fn validate(&self) -> Result<(), HarnessError> {
770 for (index, action) in self.actions.iter().enumerate() {
771 if let Action::ReadUntil { delimiter, limit } = action {
772 if delimiter.is_empty() {
773 return Err(HarnessError::new(
774 "SocketScript::read_until delimiter must not be empty",
775 ));
776 }
777 if *limit < delimiter.len() {
778 return Err(HarnessError::new(
779 "SocketScript::read_until limit is shorter than its delimiter",
780 ));
781 }
782 }
783 if matches!(action, Action::AwaitClientClose) && index + 1 != self.actions.len() {
784 return Err(HarnessError::new(
785 "SocketScript::await_client_close must be the final action",
786 ));
787 }
788 if matches!(action, Action::Close | Action::Reset) && index + 1 != self.actions.len() {
789 return Err(HarnessError::new(
790 "SocketScript close and reset actions must be final",
791 ));
792 }
793 }
794 Ok(())
795 }
796}
797
798#[derive(Clone, Debug)]
799enum Action {
800 ReadHttp1Request,
801 ReadUntil { delimiter: Vec<u8>, limit: usize },
802 ReadExact(usize),
803 ExpectBytes(Vec<u8>),
804 WriteAll(Vec<u8>),
805 Wait(GateWaiter),
806 Delay(Duration),
807 ShutdownWrite,
808 AwaitClientClose,
809 Close,
810 Reset,
811}
812
813#[derive(Clone, Debug)]
818pub struct ConnectionScript {
819 kind: ConnectionScriptKind,
820}
821
822#[derive(Clone, Debug)]
823enum ConnectionScriptKind {
824 Http1(Http1Script),
825 Socket(SocketScript),
826}
827
828impl ConnectionScript {
829 pub fn http1(script: Http1Script) -> Self {
831 Self {
832 kind: ConnectionScriptKind::Http1(script),
833 }
834 }
835
836 pub fn socket(script: SocketScript) -> Self {
838 Self {
839 kind: ConnectionScriptKind::Socket(script),
840 }
841 }
842
843 fn validate(&self) -> Result<(), HarnessError> {
844 match &self.kind {
845 ConnectionScriptKind::Http1(script) => script.validate(),
846 ConnectionScriptKind::Socket(script) => script.validate(),
847 }
848 }
849}
850
851impl From<Http1Script> for ConnectionScript {
852 fn from(script: Http1Script) -> Self {
853 Self::http1(script)
854 }
855}
856
857impl From<SocketScript> for ConnectionScript {
858 fn from(script: SocketScript) -> Self {
859 Self::socket(script)
860 }
861}
862
863#[derive(Clone, Debug)]
870pub struct EndpointPlan {
871 kind: EndpointPlanKind,
872}
873
874#[derive(Clone, Debug)]
875enum EndpointPlanKind {
876 Queue(VecDeque<ConnectionScript>),
877 Repeat {
878 script: ConnectionScript,
879 remaining: Option<usize>,
880 },
881}
882
883impl EndpointPlan {
884 pub fn queue<I, S>(scripts: I) -> Self
886 where
887 I: IntoIterator<Item = S>,
888 S: Into<ConnectionScript>,
889 {
890 Self {
891 kind: EndpointPlanKind::Queue(scripts.into_iter().map(Into::into).collect()),
892 }
893 }
894
895 pub fn repeat_n(accepts: usize, script: impl Into<ConnectionScript>) -> Self {
897 Self {
898 kind: EndpointPlanKind::Repeat {
899 script: script.into(),
900 remaining: Some(accepts),
901 },
902 }
903 }
904
905 pub fn unbounded(script: impl Into<ConnectionScript>) -> Self {
907 Self {
908 kind: EndpointPlanKind::Repeat {
909 script: script.into(),
910 remaining: None,
911 },
912 }
913 }
914
915 fn next_script(&mut self) -> Option<ConnectionScript> {
916 match &mut self.kind {
917 EndpointPlanKind::Queue(scripts) => scripts.pop_front(),
918 EndpointPlanKind::Repeat { script, remaining } => match remaining {
919 Some(0) => None,
920 Some(remaining) => {
921 *remaining -= 1;
922 Some(script.clone())
923 }
924 None => Some(script.clone()),
925 },
926 }
927 }
928
929 fn validate(&self) -> Result<(), HarnessError> {
930 match &self.kind {
931 EndpointPlanKind::Queue(scripts) => {
932 for script in scripts {
933 script.validate()?;
934 }
935 }
936 EndpointPlanKind::Repeat { script, .. } => script.validate()?,
937 }
938 Ok(())
939 }
940}
941
942impl From<ConnectionScript> for EndpointPlan {
943 fn from(script: ConnectionScript) -> Self {
944 Self::queue([script])
945 }
946}
947
948impl From<Http1Script> for EndpointPlan {
949 fn from(script: Http1Script) -> Self {
950 ConnectionScript::from(script).into()
951 }
952}
953
954impl From<SocketScript> for EndpointPlan {
955 fn from(script: SocketScript) -> Self {
956 ConnectionScript::from(script).into()
957 }
958}
959
960#[derive(Debug)]
962pub struct TestEndpoint {
963 addr: SocketAddr,
964}
965
966impl TestEndpoint {
967 pub fn ip(&self) -> IpAddr {
969 self.addr.ip()
970 }
971
972 pub fn port(&self) -> u16 {
974 self.addr.port()
975 }
976
977 pub fn addr(&self) -> SocketAddr {
979 self.addr
980 }
981
982 pub fn endpoint_url(&self) -> String {
984 format!("http://{}/", self.addr)
985 }
986}
987
988#[derive(Clone, Debug)]
995pub struct MockDnsResolver {
996 entries: Arc<HashMap<String, Vec<IpAddr>>>,
997 state: Arc<SharedState>,
998}
999
1000impl ResolveDns for MockDnsResolver {
1001 fn resolve_dns<'a>(&'a self, name: &'a str) -> DnsFuture<'a> {
1002 self.state.record_event(ConnectionEvent::DnsLookup {
1003 hostname: name.to_owned(),
1004 });
1005 match self.entries.get(name) {
1006 Some(addrs) => DnsFuture::ready(Ok(addrs.clone())),
1007 None => DnsFuture::ready(Err(ResolveDnsError::new(std::io::Error::other(format!(
1008 "no DNS entry configured for {name:?}"
1009 ))))),
1010 }
1011 }
1012}
1013
1014#[derive(Debug, Default)]
1019pub struct HarnessBuilder {
1020 endpoints: Vec<EndpointConfig>,
1021 dns: Vec<DnsConfig>,
1022}
1023
1024#[derive(Debug)]
1025struct EndpointConfig {
1026 ip: IpAddr,
1027 plan: EndpointPlan,
1028}
1029
1030#[derive(Debug)]
1031enum DnsConfig {
1032 Explicit(String, Vec<IpAddr>),
1033 All(String),
1034}
1035
1036impl HarnessBuilder {
1037 pub fn endpoint(mut self, ip: IpAddr, plan: impl Into<EndpointPlan>) -> Self {
1041 self.endpoints.push(EndpointConfig {
1042 ip,
1043 plan: plan.into(),
1044 });
1045 self
1046 }
1047
1048 pub fn dns<I>(mut self, hostname: impl Into<String>, ips: I) -> Self
1052 where
1053 I: IntoIterator<Item = IpAddr>,
1054 {
1055 self.dns.push(DnsConfig::Explicit(
1056 hostname.into(),
1057 ips.into_iter().collect(),
1058 ));
1059 self
1060 }
1061
1062 pub fn dns_all(mut self, hostname: impl Into<String>) -> Self {
1066 self.dns.push(DnsConfig::All(hostname.into()));
1067 self
1068 }
1069
1070 pub async fn build(self) -> Result<ConnectionTestHarness, HarnessError> {
1072 if self.endpoints.is_empty() {
1073 return Err(HarnessError::new(
1074 "a connection test harness requires at least one endpoint",
1075 ));
1076 }
1077 for config in &self.endpoints {
1078 config.plan.validate()?;
1079 }
1080
1081 let mut bound = Vec::with_capacity(self.endpoints.len());
1082 let mut port = 0;
1083 for config in self.endpoints {
1084 let requested = SocketAddr::new(config.ip, port);
1085 let listener = TcpListener::bind(requested).await.map_err(|err| {
1086 HarnessError::new(format!("failed to bind endpoint {requested}: {err}"))
1087 })?;
1088 let addr = listener.local_addr().map_err(|err| {
1089 HarnessError::new(format!("failed to read endpoint address: {err}"))
1090 })?;
1091 if port == 0 {
1092 port = addr.port();
1093 }
1094 bound.push((listener, addr, config.plan));
1095 }
1096
1097 let state = Arc::new(SharedState::new());
1098 let next_connection_id = Arc::new(AtomicU64::new(1));
1099 let (shutdown, _) = watch::channel(false);
1100 let mut endpoints = Vec::with_capacity(bound.len());
1101 let mut endpoint_tasks = Vec::with_capacity(bound.len());
1102 for (listener, addr, plan) in bound {
1103 endpoints.push(TestEndpoint { addr });
1104 endpoint_tasks.push(tokio::spawn(run_endpoint(
1105 listener,
1106 addr,
1107 plan,
1108 state.clone(),
1109 next_connection_id.clone(),
1110 shutdown.subscribe(),
1111 )));
1112 }
1113
1114 let all_ips = endpoints.iter().map(TestEndpoint::ip).collect::<Vec<_>>();
1115 let mut dns_entries = HashMap::new();
1116 for config in self.dns {
1117 match config {
1118 DnsConfig::Explicit(hostname, ips) => {
1119 dns_entries.insert(hostname, ips);
1120 }
1121 DnsConfig::All(hostname) => {
1122 dns_entries.insert(hostname, all_ips.clone());
1123 }
1124 }
1125 }
1126 let dns_resolver = MockDnsResolver {
1127 entries: Arc::new(dns_entries),
1128 state: state.clone(),
1129 };
1130
1131 Ok(ConnectionTestHarness {
1132 endpoints,
1133 state,
1134 dns_resolver,
1135 shutdown,
1136 endpoint_tasks,
1137 })
1138 }
1139}
1140
1141#[derive(Debug)]
1148pub struct ConnectionTestHarness {
1149 endpoints: Vec<TestEndpoint>,
1150 state: Arc<SharedState>,
1151 dns_resolver: MockDnsResolver,
1152 shutdown: watch::Sender<bool>,
1153 endpoint_tasks: Vec<JoinHandle<()>>,
1154}
1155
1156impl ConnectionTestHarness {
1157 pub fn builder() -> HarnessBuilder {
1159 HarnessBuilder::default()
1160 }
1161
1162 pub fn endpoints(&self) -> &[TestEndpoint] {
1164 &self.endpoints
1165 }
1166
1167 pub fn endpoint(&self, index: usize) -> Option<&TestEndpoint> {
1169 self.endpoints.get(index)
1170 }
1171
1172 pub fn port(&self) -> u16 {
1174 self.endpoints[0].port()
1175 }
1176
1177 pub fn endpoint_url(&self) -> String {
1179 self.endpoints[0].endpoint_url()
1180 }
1181
1182 pub fn dns_resolver(&self) -> MockDnsResolver {
1184 self.dns_resolver.clone()
1185 }
1186
1187 pub fn events(&self) -> Vec<ConnectionEvent> {
1191 self.state.events()
1192 }
1193
1194 pub fn tcp_accepted_count(&self) -> usize {
1196 self.events()
1197 .iter()
1198 .filter(|event| matches!(event, ConnectionEvent::TcpAccepted { .. }))
1199 .count()
1200 }
1201
1202 pub fn tcp_accepted_by(&self, ip: IpAddr) -> usize {
1204 self.events()
1205 .iter()
1206 .filter(|event| {
1207 matches!(
1208 event,
1209 ConnectionEvent::TcpAccepted { endpoint_addr, .. }
1210 if endpoint_addr.ip() == ip
1211 )
1212 })
1213 .count()
1214 }
1215
1216 pub fn dns_lookup_count(&self) -> usize {
1218 self.events()
1219 .iter()
1220 .filter(|event| matches!(event, ConnectionEvent::DnsLookup { .. }))
1221 .count()
1222 }
1223
1224 pub fn http_requests(&self) -> Vec<(String, Option<String>)> {
1226 self.events()
1227 .into_iter()
1228 .filter_map(|event| match event {
1229 ConnectionEvent::Http1Request { target, host, .. } => Some((target, host)),
1230 _ => None,
1231 })
1232 .collect()
1233 }
1234
1235 pub async fn wait_for_tcp_accepts(
1239 &self,
1240 expected: usize,
1241 timeout: Duration,
1242 ) -> Result<(), HarnessError> {
1243 self.state
1244 .wait_for("TCP accepts", timeout, |events| {
1245 events
1246 .iter()
1247 .filter(|event| matches!(event, ConnectionEvent::TcpAccepted { .. }))
1248 .count()
1249 >= expected
1250 })
1251 .await
1252 }
1253
1254 pub async fn wait_for_http_requests(
1258 &self,
1259 expected: usize,
1260 timeout: Duration,
1261 ) -> Result<(), HarnessError> {
1262 self.state
1263 .wait_for("HTTP/1 requests", timeout, |events| {
1264 events
1265 .iter()
1266 .filter(|event| matches!(event, ConnectionEvent::Http1Request { .. }))
1267 .count()
1268 >= expected
1269 })
1270 .await
1271 }
1272
1273 pub async fn wait_for_event<F>(
1277 &self,
1278 timeout: Duration,
1279 predicate: F,
1280 ) -> Result<(), HarnessError>
1281 where
1282 F: Fn(&ConnectionEvent) -> bool,
1283 {
1284 self.state
1285 .wait_for("matching event", timeout, |events| {
1286 events.iter().any(&predicate)
1287 })
1288 .await
1289 }
1290
1291 pub async fn shutdown(mut self) -> Result<(), HarnessError> {
1302 self.shutdown.send_replace(true);
1303 for task in self.endpoint_tasks.drain(..) {
1304 if let Err(err) = task.await {
1305 self.state.record_failure(HarnessError::new(format!(
1306 "endpoint task failed while shutting down: {err}"
1307 )));
1308 }
1309 }
1310 match self.state.failure() {
1311 Some(failure) => Err(failure),
1312 None => Ok(()),
1313 }
1314 }
1315}
1316
1317impl Drop for ConnectionTestHarness {
1318 fn drop(&mut self) {
1319 if std::thread::panicking() {
1323 if let Some(failure) = self.state.failure() {
1324 eprintln!(
1325 "\n[ConnectionTestHarness] background failure during panic:\n {failure}\n"
1326 );
1327 }
1328 }
1329 self.shutdown.send_replace(true);
1330 for task in &self.endpoint_tasks {
1331 task.abort();
1332 }
1333 }
1334}
1335
1336async fn run_endpoint(
1337 listener: TcpListener,
1338 endpoint_addr: SocketAddr,
1339 mut plan: EndpointPlan,
1340 state: Arc<SharedState>,
1341 next_connection_id: Arc<AtomicU64>,
1342 mut shutdown: watch::Receiver<bool>,
1343) {
1344 let mut connections = JoinSet::new();
1347 loop {
1348 tokio::select! {
1349 biased;
1350 _ = wait_for_shutdown(&mut shutdown) => break,
1351 completed = connections.join_next(), if !connections.is_empty() => {
1352 if let Some(Err(err)) = completed {
1353 state.record_failure(HarnessError::new(format!(
1354 "connection task at {endpoint_addr} failed: {err}"
1355 )));
1356 }
1357 }
1358 accepted = listener.accept() => {
1359 let (stream, _) = match accepted {
1360 Ok(accepted) => accepted,
1361 Err(err) => {
1362 state.record_failure(HarnessError::new(format!(
1363 "failed to accept a connection at {endpoint_addr}: {err}"
1364 )));
1365 break;
1366 }
1367 };
1368 let connection_id =
1369 ConnectionId(next_connection_id.fetch_add(1, Ordering::Relaxed));
1370 state.record_event(ConnectionEvent::TcpAccepted {
1371 connection_id,
1372 endpoint_addr,
1373 });
1374 let Some(script) = plan.next_script() else {
1375 state.record_failure(HarnessError::new(format!(
1376 "endpoint {endpoint_addr} accepted connection {connection_id} after its plan was exhausted"
1377 )));
1378 drop(stream);
1379 continue;
1380 };
1381
1382 let state = state.clone();
1383 let connection_shutdown = shutdown.clone();
1384 connections.spawn(async move {
1385 run_connection_task(
1386 stream,
1387 script,
1388 connection_id,
1389 endpoint_addr,
1390 state,
1391 connection_shutdown,
1392 )
1393 .await;
1394 });
1395 }
1396 }
1397 }
1398
1399 while let Some(result) = connections.join_next().await {
1400 if let Err(err) = result {
1401 state.record_failure(HarnessError::new(format!(
1402 "connection task at {endpoint_addr} failed while shutting down: {err}"
1403 )));
1404 }
1405 }
1406}
1407
1408async fn wait_for_shutdown(shutdown: &mut watch::Receiver<bool>) {
1409 loop {
1410 if *shutdown.borrow() {
1411 return;
1412 }
1413 if shutdown.changed().await.is_err() {
1414 return;
1415 }
1416 }
1417}
1418
1419async fn run_connection_task(
1420 stream: TcpStream,
1421 script: ConnectionScript,
1422 connection_id: ConnectionId,
1423 endpoint_addr: SocketAddr,
1424 state: Arc<SharedState>,
1425 mut shutdown: watch::Receiver<bool>,
1426) {
1427 let result = tokio::select! {
1428 biased;
1429 _ = wait_for_shutdown(&mut shutdown) => Ok(ConnectionCloseReason::HarnessShutdown),
1430 result = run_connection(stream, script, connection_id, endpoint_addr, &state) => result,
1431 };
1432 let reason = match result {
1433 Ok(reason) => reason,
1434 Err(err) => {
1435 state.record_failure(HarnessError::new(format!(
1436 "connection {connection_id} at {endpoint_addr}: {err}"
1437 )));
1438 ConnectionCloseReason::ScriptFailed
1439 }
1440 };
1441 state.record_event(ConnectionEvent::ConnectionClosed {
1442 connection_id,
1443 reason,
1444 });
1445}
1446
1447async fn run_connection(
1448 stream: TcpStream,
1449 script: ConnectionScript,
1450 connection_id: ConnectionId,
1451 endpoint_addr: SocketAddr,
1452 state: &SharedState,
1453) -> Result<ConnectionCloseReason, HarnessError> {
1454 let mut executor = ScriptExecutor {
1455 stream,
1456 pending: Vec::new(),
1457 connection_id,
1458 endpoint_addr,
1459 state,
1460 };
1461 match script.kind {
1462 ConnectionScriptKind::Socket(script) => Ok(executor
1463 .execute(&script.actions)
1464 .await?
1465 .unwrap_or(ConnectionCloseReason::ScriptCompleted)),
1466 ConnectionScriptKind::Http1(script) => match script.responses {
1467 Http1Responses::Finite(responses) => {
1468 let mut actions = Vec::new();
1469 for response in responses {
1470 actions.push(Action::ReadHttp1Request);
1471 actions.extend(response.actions());
1472 }
1473 if !actions
1474 .last()
1475 .is_some_and(|action| matches!(action, Action::Close | Action::Reset))
1476 {
1477 actions.push(match script.finish {
1478 Finish::AwaitClientClose => Action::AwaitClientClose,
1479 Finish::Close => Action::Close,
1480 Finish::Reset => Action::Reset,
1481 });
1482 }
1483 Ok(executor
1484 .execute(&actions)
1485 .await?
1486 .unwrap_or(ConnectionCloseReason::ScriptCompleted))
1487 }
1488 Http1Responses::Repeated(response) => loop {
1489 match executor.read_http1_request().await {
1490 Ok(request) => executor.record_request(request),
1491 Err(ReadRequestError::ClientClosed) => {
1492 return Ok(ConnectionCloseReason::ClientClosed);
1493 }
1494 Err(ReadRequestError::Failed(err)) => return Err(err),
1495 }
1496 if let Some(reason) = executor.execute(&response.actions()).await? {
1497 return Ok(reason);
1498 }
1499 },
1500 },
1501 }
1502}
1503
1504struct ScriptExecutor<'a> {
1505 stream: TcpStream,
1506 pending: Vec<u8>,
1507 connection_id: ConnectionId,
1508 endpoint_addr: SocketAddr,
1509 state: &'a SharedState,
1510}
1511
1512impl ScriptExecutor<'_> {
1513 async fn execute(
1514 &mut self,
1515 actions: &[Action],
1516 ) -> Result<Option<ConnectionCloseReason>, HarnessError> {
1517 for action in actions {
1518 match action {
1519 Action::ReadHttp1Request => {
1520 let request = self.read_http1_request().await.map_err(|err| match err {
1521 ReadRequestError::ClientClosed => {
1522 HarnessError::new("client closed before the expected HTTP/1 request")
1523 }
1524 ReadRequestError::Failed(err) => err,
1525 })?;
1526 self.record_request(request);
1527 }
1528 Action::ReadUntil { delimiter, limit } => {
1529 self.read_until(delimiter, *limit).await?;
1530 }
1531 Action::ReadExact(length) => {
1532 self.fill_pending(*length).await?;
1533 self.pending.drain(..*length);
1534 }
1535 Action::ExpectBytes(expected) => {
1536 self.fill_pending(expected.len()).await?;
1537 if self.pending[..expected.len()] != expected[..] {
1538 return Err(HarnessError::new(format!(
1539 "socket bytes differed: expected {expected:?}, got {:?}",
1540 &self.pending[..expected.len()]
1541 )));
1542 }
1543 self.pending.drain(..expected.len());
1544 }
1545 Action::WriteAll(bytes) => {
1546 self.stream
1547 .write_all(bytes)
1548 .await
1549 .map_err(|err| HarnessError::new(format!("failed to write: {err}")))?;
1550 }
1551 Action::Wait(gate) => gate.wait().await?,
1552 Action::Delay(duration) => tokio::time::sleep(*duration).await,
1553 Action::ShutdownWrite => {
1554 self.stream
1555 .shutdown()
1556 .await
1557 .map_err(|err| HarnessError::new(format!("failed to shut down: {err}")))?;
1558 }
1559 Action::AwaitClientClose => {
1560 if !self.pending.is_empty() {
1561 return Err(HarnessError::new(
1562 "client sent bytes after the scripted HTTP/1 responses were exhausted",
1563 ));
1564 }
1565 let mut byte = [0u8; 1];
1566 return match self.stream.read(&mut byte).await {
1567 Ok(0) => Ok(Some(ConnectionCloseReason::ClientClosed)),
1568 Ok(_) => Err(HarnessError::new(
1569 "client sent another request after the HTTP/1 script was exhausted",
1570 )),
1571 Err(err) if peer_close_error(&err) => {
1572 Ok(Some(ConnectionCloseReason::ClientClosed))
1573 }
1574 Err(err) => Err(HarnessError::new(format!(
1575 "failed while waiting for the client to close: {err}"
1576 ))),
1577 };
1578 }
1579 Action::Close => {
1580 return Ok(Some(ConnectionCloseReason::ScriptCompleted));
1581 }
1582 Action::Reset => {
1583 socket2::SockRef::from(&self.stream)
1584 .set_linger(Some(Duration::ZERO))
1585 .map_err(|err| {
1586 HarnessError::new(format!("failed to configure TCP reset: {err}"))
1587 })?;
1588 return Ok(Some(ConnectionCloseReason::Reset));
1589 }
1590 }
1591 }
1592 Ok(None)
1593 }
1594
1595 async fn read_until(&mut self, delimiter: &[u8], limit: usize) -> Result<(), HarnessError> {
1596 loop {
1597 if let Some(index) = find_bytes(&self.pending, delimiter) {
1598 let consumed = index + delimiter.len();
1599 if consumed > limit {
1600 return Err(HarnessError::new(format!(
1601 "read_until exceeded its {limit}-byte limit"
1602 )));
1603 }
1604 self.pending.drain(..consumed);
1605 return Ok(());
1606 }
1607 if self.pending.len() >= limit {
1608 return Err(HarnessError::new(format!(
1609 "read_until did not find its delimiter within {limit} bytes"
1610 )));
1611 }
1612 self.read_more().await?;
1613 }
1614 }
1615
1616 async fn fill_pending(&mut self, length: usize) -> Result<(), HarnessError> {
1617 while self.pending.len() < length {
1618 self.read_more().await?;
1619 }
1620 Ok(())
1621 }
1622
1623 async fn read_more(&mut self) -> Result<(), HarnessError> {
1624 let mut chunk = [0u8; READ_CHUNK_SIZE];
1625 match self.stream.read(&mut chunk).await {
1626 Ok(0) => Err(HarnessError::new(
1627 "client closed while the script was reading",
1628 )),
1629 Ok(read) => {
1630 self.pending.extend_from_slice(&chunk[..read]);
1631 Ok(())
1632 }
1633 Err(err) => Err(HarnessError::new(format!(
1634 "failed to read from client: {err}"
1635 ))),
1636 }
1637 }
1638
1639 async fn read_http1_request(&mut self) -> Result<ParsedRequest, ReadRequestError> {
1640 loop {
1641 let parsed = parse_request_head(&self.pending).map_err(ReadRequestError::Failed)?;
1642 if let Some(mut request) = parsed {
1643 let total_length = request
1644 .header_length
1645 .checked_add(request.body_length)
1646 .ok_or_else(|| {
1647 ReadRequestError::Failed(HarnessError::new(
1648 "HTTP/1 request length overflow",
1649 ))
1650 })?;
1651 if request.body_length > MAX_HTTP1_BODY_BYTES {
1652 return Err(ReadRequestError::Failed(HarnessError::new(format!(
1653 "HTTP/1 request body exceeds {MAX_HTTP1_BODY_BYTES} bytes"
1654 ))));
1655 }
1656 while self.pending.len() < total_length {
1657 self.read_more().await.map_err(ReadRequestError::Failed)?;
1658 }
1659 self.pending.drain(..total_length);
1660 request.header_length = 0;
1661 request.body_length = 0;
1662 return Ok(request);
1663 }
1664 if self.pending.len() >= MAX_HTTP1_HEADER_BYTES {
1665 return Err(ReadRequestError::Failed(HarnessError::new(format!(
1666 "HTTP/1 request headers exceed {MAX_HTTP1_HEADER_BYTES} bytes"
1667 ))));
1668 }
1669
1670 let mut chunk = [0u8; READ_CHUNK_SIZE];
1671 match self.stream.read(&mut chunk).await {
1672 Ok(0) if self.pending.is_empty() => return Err(ReadRequestError::ClientClosed),
1673 Ok(0) => {
1674 return Err(ReadRequestError::Failed(HarnessError::new(
1675 "client closed during HTTP/1 request headers",
1676 )))
1677 }
1678 Ok(read) => self.pending.extend_from_slice(&chunk[..read]),
1679 Err(err) if self.pending.is_empty() && peer_close_error(&err) => {
1680 return Err(ReadRequestError::ClientClosed)
1681 }
1682 Err(err) => {
1683 return Err(ReadRequestError::Failed(HarnessError::new(format!(
1684 "failed to read HTTP/1 request: {err}"
1685 ))))
1686 }
1687 }
1688 }
1689 }
1690
1691 fn record_request(&self, request: ParsedRequest) {
1692 self.state.record_event(ConnectionEvent::Http1Request {
1693 connection_id: self.connection_id,
1694 endpoint_addr: self.endpoint_addr,
1695 method: request.method,
1696 target: request.target,
1697 host: request.host,
1698 });
1699 }
1700}
1701
1702enum ReadRequestError {
1703 ClientClosed,
1704 Failed(HarnessError),
1705}
1706
1707struct ParsedRequest {
1708 method: String,
1709 target: String,
1710 host: Option<String>,
1711 header_length: usize,
1712 body_length: usize,
1713}
1714
1715fn parse_request_head(bytes: &[u8]) -> Result<Option<ParsedRequest>, HarnessError> {
1716 let mut headers = [httparse::EMPTY_HEADER; 64];
1717 let mut request = httparse::Request::new(&mut headers);
1718 let header_length = match request
1719 .parse(bytes)
1720 .map_err(|err| HarnessError::new(format!("invalid HTTP/1 request: {err}")))?
1721 {
1722 httparse::Status::Partial => return Ok(None),
1723 httparse::Status::Complete(length) => length,
1724 };
1725 if header_length > MAX_HTTP1_HEADER_BYTES {
1726 return Err(HarnessError::new(format!(
1727 "HTTP/1 request headers exceed {MAX_HTTP1_HEADER_BYTES} bytes"
1728 )));
1729 }
1730 let method = request
1731 .method
1732 .ok_or_else(|| HarnessError::new("HTTP/1 request has no method"))?
1733 .to_owned();
1734 let target = request
1735 .path
1736 .ok_or_else(|| HarnessError::new("HTTP/1 request has no target"))?
1737 .to_owned();
1738 let mut host = None;
1739 let mut content_length = None;
1740 for header in request.headers.iter() {
1741 if header.name.eq_ignore_ascii_case("host") {
1742 host = Some(
1743 std::str::from_utf8(header.value)
1744 .map_err(|_| HarnessError::new("Host header is not valid UTF-8"))?
1745 .trim()
1746 .to_owned(),
1747 );
1748 } else if header.name.eq_ignore_ascii_case("content-length") {
1749 if content_length.is_some() {
1750 return Err(HarnessError::new(
1751 "multiple Content-Length headers are not supported",
1752 ));
1753 }
1754 let value = std::str::from_utf8(header.value)
1755 .map_err(|_| HarnessError::new("Content-Length is not valid ASCII"))?
1756 .trim();
1757 content_length = Some(
1758 value
1759 .parse::<usize>()
1760 .map_err(|_| HarnessError::new(format!("invalid Content-Length {value:?}")))?,
1761 );
1762 } else if header.name.eq_ignore_ascii_case("transfer-encoding") {
1763 return Err(HarnessError::new(
1764 "Transfer-Encoding is not supported by read_http1_request; use raw socket actions",
1765 ));
1766 }
1767 }
1768
1769 Ok(Some(ParsedRequest {
1770 method,
1771 target,
1772 host,
1773 header_length,
1774 body_length: content_length.unwrap_or(0),
1775 }))
1776}
1777
1778fn find_bytes(haystack: &[u8], needle: &[u8]) -> Option<usize> {
1779 haystack
1780 .windows(needle.len())
1781 .position(|window| window == needle)
1782}
1783
1784fn peer_close_error(err: &std::io::Error) -> bool {
1785 matches!(
1786 err.kind(),
1787 std::io::ErrorKind::ConnectionAborted
1788 | std::io::ErrorKind::ConnectionReset
1789 | std::io::ErrorKind::BrokenPipe
1790 | std::io::ErrorKind::NotConnected
1791 )
1792}