1#![cfg(all(feature = "tsig", feature = "unstable-client-transport"))]
47#![warn(missing_docs)]
48#![warn(clippy::missing_docs_in_private_items)]
49
50use core::ops::DerefMut;
51
52use alloc::boxed::Box;
53use alloc::sync::Arc;
54use alloc::vec::Vec;
55use core::fmt::{Debug, Formatter};
56use core::future::Future;
57use core::pin::Pin;
58
59use bytes::Bytes;
60use octseq::Octets;
61use tracing::trace;
62
63use crate::base::Message;
64use crate::base::StaticCompressor;
65use crate::base::message::CopyRecordsError;
66use crate::base::message_builder::AdditionalBuilder;
67use crate::base::wire::Composer;
68use crate::net::client::request::{
69 ComposeRequest, ComposeRequestMulti, Error, GetResponse,
70 GetResponseMulti, SendRequest, SendRequestMulti,
71};
72use crate::rdata::tsig::Time48;
73use crate::tsig::{ClientSequence, ClientTransaction, Key};
74
75#[derive(Clone, Debug)]
83enum TsigClient<K> {
84 Transaction(ClientTransaction<K>),
87
88 Sequence(ClientSequence<K>),
91}
92
93impl<K> TsigClient<K>
94where
95 K: AsRef<Key>,
96{
97 pub fn answer<Octs>(
101 &mut self,
102 message: &mut Message<Octs>,
103 now: Time48,
104 ) -> Result<(), Error>
105 where
106 Octs: Octets + AsMut<[u8]> + ?Sized,
107 {
108 match self {
109 TsigClient::Transaction(client) => client.answer(message, now),
110 TsigClient::Sequence(client) => client.answer(message, now),
111 }
112 .map_err(Error::Authentication)
113 }
114
115 fn done(self) -> Result<(), Error> {
122 match self {
123 TsigClient::Transaction(_) => {
124 Ok(())
126 }
127 TsigClient::Sequence(client) => {
128 client.done().map_err(Error::Authentication)
129 }
130 }
131 }
132}
133
134#[derive(Clone)]
142pub struct Connection<Upstream, K> {
143 upstream: Arc<Upstream>,
148
149 key: K,
151}
152
153impl<Upstream, K> Connection<Upstream, K> {
154 pub fn new(key: K, upstream: Upstream) -> Self {
160 Self {
161 upstream: Arc::new(upstream),
162 key,
163 }
164 }
165}
166
167impl<CR, Upstream, K> SendRequest<CR> for Connection<Upstream, K>
170where
171 CR: ComposeRequest + 'static,
172 Upstream: SendRequest<RequestMessage<CR, K>> + Send + Sync + 'static,
173 K: Clone + AsRef<Key> + Send + Sync + 'static,
174{
175 fn send_request(
176 &self,
177 request_msg: CR,
178 ) -> Box<dyn GetResponse + Send + Sync> {
179 Box::new(Request::<CR, Upstream, K>::new(
180 request_msg,
181 self.key.clone(),
182 self.upstream.clone(),
183 ))
184 }
185}
186
187impl<CR, Upstream, K> SendRequestMulti<CR> for Connection<Upstream, K>
190where
191 CR: ComposeRequestMulti + 'static,
192 Upstream: SendRequestMulti<RequestMessage<CR, K>> + Send + Sync + 'static,
193 K: Clone + AsRef<Key> + Send + Sync + 'static,
194{
195 fn send_request(
196 &self,
197 request_msg: CR,
198 ) -> Box<dyn GetResponseMulti + Send + Sync> {
199 Box::new(Request::<CR, Upstream, K>::new_multi(
200 request_msg,
201 self.key.clone(),
202 self.upstream.clone(),
203 ))
204 }
205}
206
207type Forwarder<Upstream, CR, K> = fn(
216 &Upstream,
217 RequestMessage<CR, K>,
218 Arc<std::sync::Mutex<Option<TsigClient<K>>>>,
219) -> RequestState<K>;
220
221fn forwarder<CR, K, Upstream>(
227 upstream: &Upstream,
228 msg: RequestMessage<CR, K>,
229 tsig_client: Arc<std::sync::Mutex<Option<TsigClient<K>>>>,
230) -> RequestState<K>
231where
232 CR: ComposeRequest,
233 Upstream: SendRequest<RequestMessage<CR, K>> + Send + Sync,
234{
235 RequestState::GetResponse(upstream.send_request(msg), tsig_client)
236}
237
238fn forwarder_multi<CR, K, Upstream>(
244 upstream: &Upstream,
245 msg: RequestMessage<CR, K>,
246 tsig_client: Arc<std::sync::Mutex<Option<TsigClient<K>>>>,
247) -> RequestState<K>
248where
249 CR: ComposeRequestMulti,
250 Upstream: SendRequestMulti<RequestMessage<CR, K>> + Send + Sync,
251{
252 RequestState::GetResponseMulti(upstream.send_request(msg), tsig_client)
253}
254
255struct Request<CR, Upstream, K> {
259 state: RequestState<K>,
261
262 request_msg: Option<CR>,
266
267 key: K,
269
270 upstream: Arc<Upstream>,
272}
273
274impl<CR, Upstream, K> Request<CR, Upstream, K>
275where
276 CR: ComposeRequest,
277 Upstream: SendRequest<RequestMessage<CR, K>> + Send + Sync,
278 K: Clone + AsRef<Key>,
279 Self: GetResponse,
280{
281 fn new(request_msg: CR, key: K, upstream: Arc<Upstream>) -> Self {
283 Self {
284 state: RequestState::Init,
285 request_msg: Some(request_msg),
286 key,
287 upstream,
288 }
289 }
290}
291
292impl<CR, Upstream, K> Request<CR, Upstream, K>
293where
294 CR: Sync + Send,
295 K: Clone + AsRef<Key>,
296{
297 fn new_multi(request_msg: CR, key: K, upstream: Arc<Upstream>) -> Self {
299 Self {
300 state: RequestState::Init,
301 request_msg: Some(request_msg),
302 key,
303 upstream,
304 }
305 }
306
307 async fn get_response_impl(
311 &mut self,
312 upstream_sender: Forwarder<Upstream, CR, K>,
313 ) -> Result<Option<Message<Bytes>>, Error> {
314 let (response, tsig_client) = loop {
315 match &mut self.state {
316 RequestState::Init => {
317 let tsig_client = Arc::new(std::sync::Mutex::new(None));
318
319 let msg = RequestMessage::new(
320 self.request_msg.take().unwrap(),
321 self.key.clone(),
322 tsig_client.clone(),
323 );
324
325 trace!("Sending request upstream...");
326 self.state =
327 upstream_sender(&self.upstream, msg, tsig_client);
328 continue;
329 }
330
331 RequestState::GetResponse(request, tsig_client) => {
332 let response = request.get_response().await?;
333 break (Some(response), tsig_client);
334 }
335
336 RequestState::GetResponseMulti(request, tsig_client) => {
337 let response = request.get_response().await?;
338 break (response, tsig_client);
339 }
340
341 RequestState::Complete => {
342 return Err(Error::StreamReceiveError);
343 }
344 }
345 };
346
347 let res = Self::validate_response(response, tsig_client)?;
348
349 if res.is_none() {
350 self.state = RequestState::Complete;
351 }
352
353 Ok(res)
354 }
355
356 fn validate_response(
377 response: Option<Message<Bytes>>,
378 tsig_client: &mut Arc<std::sync::Mutex<Option<TsigClient<K>>>>,
379 ) -> Result<Option<Message<Bytes>>, Error> {
380 let res = match response {
381 None => {
382 let client = tsig_client.lock().unwrap().take().unwrap();
383 client.done()?;
384 None
385 }
386
387 Some(msg) => {
388 let mut modifiable_msg =
389 Message::from_octets(msg.as_slice().to_vec())?;
390
391 if let Some(client) = tsig_client.lock().unwrap().deref_mut()
392 {
393 trace!("Validating TSIG for sequence reply");
394 client.answer(&mut modifiable_msg, Time48::now())?;
395 }
396
397 let out_vec = modifiable_msg.into_octets();
398 let out_bytes = Bytes::from(out_vec);
399 let out_msg = Message::<Bytes>::from_octets(out_bytes)?;
400 Some(out_msg)
401 }
402 };
403
404 Ok(res)
405 }
406}
407
408impl<CR, Upstream, K> Debug for Request<CR, Upstream, K> {
411 fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), core::fmt::Error> {
412 f.debug_struct("Request").finish()
413 }
414}
415
416impl<CR, Upstream, K> GetResponse for Request<CR, Upstream, K>
419where
420 CR: ComposeRequest,
421 Upstream: SendRequest<RequestMessage<CR, K>> + Send + Sync,
422 K: Clone + AsRef<Key> + Send + Sync,
423{
424 fn get_response(
425 &mut self,
426 ) -> Pin<
427 Box<
428 dyn Future<Output = Result<Message<Bytes>, Error>>
429 + Send
430 + Sync
431 + '_,
432 >,
433 > {
434 Box::pin(async move {
435 self.get_response_impl(forwarder).await.map(|v| v.unwrap())
438 })
439 }
440}
441
442impl<CR, Upstream, K> GetResponseMulti for Request<CR, Upstream, K>
445where
446 CR: ComposeRequestMulti,
447 Upstream: SendRequestMulti<RequestMessage<CR, K>> + Send + Sync,
448 K: Clone + AsRef<Key> + Send + Sync,
449{
450 fn get_response(
451 &mut self,
452 ) -> Pin<
453 Box<
454 dyn Future<Output = Result<Option<Message<Bytes>>, Error>>
455 + Send
456 + Sync
457 + '_,
458 >,
459 > {
460 Box::pin(self.get_response_impl(forwarder_multi))
461 }
462}
463
464enum RequestState<K> {
472 Init,
474
475 GetResponse(
477 Box<dyn GetResponse + Send + Sync>,
478 Arc<std::sync::Mutex<Option<TsigClient<K>>>>,
479 ),
480
481 GetResponseMulti(
483 Box<dyn GetResponseMulti + Send + Sync>,
484 Arc<std::sync::Mutex<Option<TsigClient<K>>>>,
485 ),
486
487 Complete,
494}
495
496#[derive(Clone, Debug)]
514pub struct RequestMessage<CR, K>
515where
516 CR: Send + Sync,
517{
518 request: CR,
520
521 key: K,
523
524 signer: Arc<std::sync::Mutex<Option<TsigClient<K>>>>,
538}
539
540impl<CR, K> RequestMessage<CR, K>
541where
542 CR: Send + Sync,
543{
544 fn new(
546 request: CR,
547 key: K,
548 signer: Arc<std::sync::Mutex<Option<TsigClient<K>>>>,
549 ) -> Self
550 where
551 CR: Sync + Send,
552 K: Clone + AsRef<Key>,
553 {
554 Self {
555 request,
556 key,
557 signer,
558 }
559 }
560}
561
562impl<CR, K> ComposeRequest for RequestMessage<CR, K>
563where
564 CR: ComposeRequest,
565 K: Clone + Debug + Send + Sync + AsRef<Key>,
566{
567 fn append_message<Target: Composer>(
569 &self,
570 target: Target,
571 ) -> Result<AdditionalBuilder<Target>, CopyRecordsError> {
572 let mut target = self.request.append_message(target)?;
573
574 let client = {
575 trace!(
576 "Signing single request transaction with key '{}'",
577 self.key.as_ref().name()
578 );
579 TsigClient::Transaction(
580 ClientTransaction::request(
581 self.key.clone(),
582 &mut target,
583 Time48::now(),
584 )
585 .unwrap(),
586 )
587 };
588
589 *self.signer.lock().unwrap() = Some(client);
590
591 Ok(target)
592 }
593
594 fn to_vec(&self) -> Result<Vec<u8>, Error> {
595 let msg = self.to_message()?;
596 Ok(msg.as_octets().clone())
597 }
598
599 fn to_message(&self) -> Result<Message<Vec<u8>>, Error> {
600 let mut target = StaticCompressor::new(Vec::new());
601
602 self.append_message(&mut target)?;
603
604 let msg = Message::from_octets(target.into_target()).expect(
608 "Message should be able to parse output from MessageBuilder",
609 );
610 Ok(msg)
611 }
612
613 fn header(&self) -> &crate::base::Header {
614 self.request.header()
615 }
616
617 fn header_mut(&mut self) -> &mut crate::base::Header {
618 self.request.header_mut()
619 }
620
621 fn set_udp_payload_size(&mut self, value: u16) {
622 self.request.set_udp_payload_size(value)
623 }
624
625 fn set_dnssec_ok(&mut self, value: bool) {
626 self.request.set_dnssec_ok(value)
627 }
628
629 fn add_opt(
630 &mut self,
631 opt: &impl crate::base::opt::ComposeOptData,
632 ) -> Result<(), crate::base::opt::LongOptData> {
633 self.request.add_opt(opt)
634 }
635
636 fn is_answer(&self, answer: &Message<[u8]>) -> bool {
637 self.request.is_answer(answer)
638 }
639
640 fn dnssec_ok(&self) -> bool {
641 self.request.dnssec_ok()
642 }
643}
644
645impl<CR, K> ComposeRequestMulti for RequestMessage<CR, K>
646where
647 CR: ComposeRequestMulti,
648 K: Clone + Debug + Send + Sync + AsRef<Key>,
649{
650 fn append_message<Target: Composer>(
652 &self,
653 target: Target,
654 ) -> Result<AdditionalBuilder<Target>, CopyRecordsError> {
655 let mut target = self.request.append_message(target)?;
656
657 trace!(
658 "Signing streaming request sequence with key '{}'",
659 self.key.as_ref().name()
660 );
661 let client = TsigClient::Sequence(
662 ClientSequence::request(
663 self.key.clone(),
664 &mut target,
665 Time48::now(),
666 )
667 .unwrap(),
668 );
669
670 *self.signer.lock().unwrap() = Some(client);
671
672 Ok(target)
673 }
674
675 fn to_message(&self) -> Result<Message<Vec<u8>>, Error> {
676 let mut target = StaticCompressor::new(Vec::new());
677
678 self.append_message(&mut target)?;
679
680 let msg = Message::from_octets(target.into_target()).expect(
684 "Message should be able to parse output from MessageBuilder",
685 );
686 Ok(msg)
687 }
688
689 fn header(&self) -> &crate::base::Header {
690 self.request.header()
691 }
692
693 fn header_mut(&mut self) -> &mut crate::base::Header {
694 self.request.header_mut()
695 }
696
697 fn set_udp_payload_size(&mut self, value: u16) {
698 self.request.set_udp_payload_size(value)
699 }
700
701 fn set_dnssec_ok(&mut self, value: bool) {
702 self.request.set_dnssec_ok(value)
703 }
704
705 fn add_opt(
706 &mut self,
707 opt: &impl crate::base::opt::ComposeOptData,
708 ) -> Result<(), crate::base::opt::LongOptData> {
709 self.request.add_opt(opt)
710 }
711
712 fn is_answer(&self, answer: &Message<[u8]>) -> bool {
713 self.request.is_answer(answer)
714 }
715
716 fn dnssec_ok(&self) -> bool {
717 self.request.dnssec_ok()
718 }
719}
720
721#[cfg(test)]
722mod tests {
723 use super::*;
724 use crate::base::iana::Rcode;
725 use crate::base::message_builder::QuestionBuilder;
726 use crate::base::{MessageBuilder, Name, Rtype};
727 use crate::tsig::{
728 Algorithm, KeyName, KeyStore, ServerSequence, ServerTransaction,
729 ValidationError,
730 };
731 use core::future::ready;
732 use core::str::FromStr;
733 use std::eprintln;
734
735 #[tokio::test]
736 async fn single_signed_valid_response() {
737 do_single_response(false).await;
738 }
739
740 #[tokio::test]
741 async fn single_signed_invalid_response() {
742 do_single_response(true).await;
743 }
744
745 async fn do_single_response(invalidate_signature: bool) {
746 let msg = mk_request_msg(Rtype::A);
749
750 let req =
753 crate::net::client::request::RequestMessage::new(msg).unwrap();
754
755 let key = mk_tsig_key();
757
758 let upstream =
761 Arc::new(MockUpstream::new(key.clone(), invalidate_signature));
762
763 let mut req = Request::new(req, key, upstream);
766
767 let res = req.get_response().await;
769
770 assert_eq!(res.is_err(), invalidate_signature);
771
772 if let Ok(res) = res {
773 assert_eq!(
776 res.header_counts().arcount(),
777 0,
778 "TSIG RR should have been removed from the additional section during response processing"
779 );
780 }
781 }
782
783 #[tokio::test]
784 async fn multiple_signed_valid_responses() {
785 do_multiple_responses(false, false).await
786 }
787
788 #[tokio::test]
789 async fn multiple_signed_responses_with_one_invalid() {
790 do_multiple_responses(true, false).await
791 }
792
793 #[tokio::test]
794 async fn multiple_signed_valid_responses_and_a_final_unsigned_response() {
795 do_multiple_responses(false, true).await
796 }
797
798 async fn do_multiple_responses(
799 invalidate_signature: bool,
800 dont_sign_last_response: bool,
801 ) {
802 let msg = mk_request_msg(Rtype::AXFR);
805
806 let req = crate::net::client::request::RequestMessageMulti::new(msg)
809 .unwrap();
810
811 let key = mk_tsig_key();
813
814 let upstream = Arc::new(MockUpstreamMulti::new(
817 key.clone(),
818 invalidate_signature,
819 dont_sign_last_response,
820 ));
821
822 let mut req = Request::new_multi(req, key, upstream);
825
826 let res = req
828 .get_response()
829 .await
830 .unwrap()
831 .expect("First response is missing");
832
833 assert_eq!(
836 res.header_counts().arcount(),
837 0,
838 "TSIG RR should have been removed from the additional section during response processing"
839 );
840
841 let res = req.get_response().await;
844
845 if invalidate_signature {
846 assert!(
847 matches!(
848 res,
849 Err(Error::Authentication(ValidationError::BadSig))
850 ),
851 "Expected error BadSig but the result was: {res:?}"
852 );
853 } else {
854 assert!(res.is_ok(), "Unexpected error message: {res:?}");
855 }
856
857 if let Ok(res) = res {
858 let res = res.expect("Second response is missing");
859
860 assert_eq!(
863 res.header_counts().arcount(),
864 0,
865 "TSIG RR should have been removed from the additional section during response processing"
866 );
867
868 let res = req
875 .get_response()
876 .await
877 .unwrap()
878 .expect("Third response is missing");
879
880 if dont_sign_last_response {
884 assert_eq!(
885 res.header_counts().arcount(),
886 0,
887 "TSIG RR should never have been added to the additional section during response generation"
888 );
889 } else {
890 assert_eq!(
891 res.header_counts().arcount(),
892 0,
893 "TSIG RR should have been removed from the additional section during response processing"
894 );
895 }
896
897 if dont_sign_last_response {
898 assert!(
901 matches!(
902 req.get_response().await,
903 Err(Error::Authentication(
904 ValidationError::TooManyUnsigned
905 ))
906 ),
907 "Receiving another response should have failed because the last response should have lacked a signature"
908 );
909 } else {
910 assert!(
913 req.get_response().await.unwrap().is_none(),
914 "There should not be a fourth response"
915 );
916 }
917 }
918 }
919
920 fn mk_request_msg(rtype: Rtype) -> QuestionBuilder<Vec<u8>> {
922 let mut msg = MessageBuilder::new_vec();
923 msg.header_mut().set_rd(true);
924 msg.header_mut().set_ad(true);
925 let mut msg = msg.question();
926 msg.push((Name::vec_from_str("example.com").unwrap(), rtype))
927 .unwrap();
928 msg
929 }
930
931 fn mk_tsig_key() -> Arc<Key> {
933 let key_name = KeyName::from_str("demo-key").unwrap();
935 let secret = crate::utils::base64::decode::<Vec<u8>>(
936 "zlCZbVJPIhobIs1gJNQfrsS3xCxxsR9pMUrGwG8OgG8=",
937 )
938 .unwrap();
939 Arc::new(
940 Key::new(Algorithm::Sha256, &secret, key_name, None, None)
941 .unwrap(),
942 )
943 }
944
945 #[derive(Debug)]
948 struct MockGetResponse<CR, KS> {
949 request_msg: CR,
950 key_store: KS,
951 invalidate_signature: bool,
952 }
953
954 impl<CR, KS> MockGetResponse<CR, KS> {
955 fn new(
956 request_msg: CR,
957 key_store: KS,
958 invalidate_signature: bool,
959 ) -> Self {
960 Self {
961 request_msg,
962 key_store,
963 invalidate_signature,
964 }
965 }
966 }
967
968 impl<CR: ComposeRequest + Debug, KS: Debug + KeyStore> GetResponse
971 for MockGetResponse<CR, KS>
972 {
973 fn get_response(
974 &mut self,
975 ) -> Pin<
976 Box<
977 dyn Future<Output = Result<Message<Bytes>, Error>>
978 + Send
979 + Sync
980 + '_,
981 >,
982 > {
983 let mut req = self.request_msg.to_message().unwrap();
984
985 let tsig = ServerTransaction::request(
988 &self.key_store,
989 &mut req,
990 Time48::now(),
991 )
992 .unwrap()
993 .unwrap();
994
995 let builder = MessageBuilder::new_bytes();
997 let builder = builder.start_answer(&req, Rcode::NOERROR).unwrap();
998 let mut builder = builder.additional();
999
1000 tsig.answer(&mut builder, Time48::now()).unwrap();
1002
1003 if self.invalidate_signature {
1004 builder.header_mut().set_rcode(Rcode::SERVFAIL);
1006 }
1007
1008 let res = builder.into_message();
1011 assert_eq!(
1012 res.header_counts().arcount(),
1013 1,
1014 "Constructed response lacks a TSIG RR in the additional section"
1015 );
1016 Box::pin(ready(Ok(res)))
1017 }
1018 }
1019
1020 #[derive(Debug)]
1023 struct MockGetResponseMulti<CR, KS> {
1024 request_msg: CR,
1025 key_store: KS,
1026 sent_request: Option<Message<Vec<u8>>>,
1027 num_responses_generated: usize,
1028 signer: Option<ServerSequence<KS>>,
1029 invalidate_signature: bool,
1030 dont_sign_last_response: bool,
1031 }
1032
1033 impl<CR, KS> MockGetResponseMulti<CR, KS> {
1034 fn new(
1035 request_msg: CR,
1036 key_store: KS,
1037 invalidate_signature: bool,
1038 dont_sign_last_response: bool,
1039 ) -> Self {
1040 Self {
1041 request_msg,
1042 key_store,
1043 sent_request: None,
1044 num_responses_generated: 0,
1045 signer: None,
1046 invalidate_signature,
1047 dont_sign_last_response,
1048 }
1049 }
1050 }
1051
1052 impl<CR, KS> GetResponseMulti for MockGetResponseMulti<CR, KS>
1055 where
1056 CR: ComposeRequestMulti + Debug,
1057 KS: Debug + KeyStore<Key = KS> + AsRef<Key>,
1058 {
1059 fn get_response(
1060 &mut self,
1061 ) -> Pin<
1062 Box<
1063 dyn Future<Output = Result<Option<Message<Bytes>>, Error>>
1064 + Send
1065 + Sync
1066 + '_,
1067 >,
1068 > {
1069 if self.num_responses_generated == 3 {
1071 return Box::pin(ready(Ok(None)));
1072 }
1073
1074 self.num_responses_generated += 1;
1075
1076 let mut tsig = match self.signer.take() {
1079 Some(tsig) => tsig,
1080 None => {
1081 let mut req = self.request_msg.to_message().unwrap();
1082
1083 let tsig = ServerSequence::request(
1084 &self.key_store,
1085 &mut req,
1086 Time48::now(),
1087 )
1088 .unwrap()
1089 .unwrap();
1090
1091 self.sent_request = Some(req);
1094
1095 tsig
1096 }
1097 };
1098
1099 let req = self.sent_request.as_ref().unwrap();
1101 let builder = MessageBuilder::new_bytes();
1102 let builder = builder.start_answer(req, Rcode::NOERROR).unwrap();
1103 let mut builder = builder.additional();
1104
1105 let (sign, invalidate) = match self.num_responses_generated {
1107 1 => (true, false),
1108 2 => (true, self.invalidate_signature),
1109 3 => (!self.dont_sign_last_response, false),
1110 _ => unreachable!(),
1111 };
1112
1113 eprintln!(
1114 "Response {}: sign={}, invalidate={}",
1115 self.num_responses_generated, sign, invalidate
1116 );
1117
1118 if sign {
1120 tsig.answer(&mut builder, Time48::now()).unwrap();
1121 }
1122
1123 self.signer = Some(tsig);
1125
1126 if invalidate {
1127 builder.header_mut().set_rcode(Rcode::SERVFAIL);
1129 }
1130
1131 let res = builder.into_message();
1134 if sign {
1135 assert_eq!(
1136 res.header_counts().arcount(),
1137 1,
1138 "Constructed response lacks a TSIG RR in the additional section"
1139 );
1140 let rec = res.additional().unwrap().next().unwrap().unwrap();
1141 assert_eq!(rec.rtype(), Rtype::TSIG);
1142 }
1143 Box::pin(ready(Ok(Some(res))))
1144 }
1145 }
1146
1147 struct MockUpstream {
1150 key: Arc<Key>,
1151 invalidate_signature: bool,
1152 }
1153
1154 impl MockUpstream {
1155 fn new(key: Arc<Key>, invalidate_signature: bool) -> Self {
1156 Self {
1157 key,
1158 invalidate_signature,
1159 }
1160 }
1161 }
1162
1163 impl<CR: ComposeRequest + Debug + Send + Sync + 'static> SendRequest<CR>
1166 for MockUpstream
1167 {
1168 fn send_request(
1169 &self,
1170 request_msg: CR,
1171 ) -> Box<dyn GetResponse + Send + Sync> {
1172 Box::new(MockGetResponse::new(
1173 request_msg,
1174 self.key.clone(),
1175 self.invalidate_signature,
1176 ))
1177 }
1178 }
1179
1180 struct MockUpstreamMulti {
1183 key: Arc<Key>,
1184 invalidate_signature: bool,
1185 dont_sign_last_response: bool,
1186 }
1187 impl MockUpstreamMulti {
1188 fn new(
1189 key: Arc<Key>,
1190 invalidate_signature: bool,
1191 dont_sign_last_response: bool,
1192 ) -> Self {
1193 Self {
1194 key,
1195 invalidate_signature,
1196 dont_sign_last_response,
1197 }
1198 }
1199 }
1200
1201 impl<CR> SendRequestMulti<CR> for MockUpstreamMulti
1202 where
1203 CR: ComposeRequestMulti + Debug + Send + Sync + 'static,
1204 {
1205 fn send_request(
1206 &self,
1207 request_msg: CR,
1208 ) -> Box<dyn GetResponseMulti + Send + Sync> {
1209 Box::new(MockGetResponseMulti::new(
1210 request_msg,
1211 self.key.clone(),
1212 self.invalidate_signature,
1213 self.dont_sign_last_response,
1214 ))
1215 }
1216 }
1217}