1use crate::builder::MessageBuilder;
30use crate::message::{Message, Question, Record};
31use crate::name::ToName;
32use crate::rdata::{ParseRdata, Soa};
33use crate::wire::OutBuf;
34use crate::{Class, Error, Flags, Opcode, Result, Rtype};
35
36pub fn build_query<B: OutBuf>(
44 b: &mut MessageBuilder<B>,
45 zone: impl ToName,
46 class: Class,
47 soa: Option<(u32, &Soa<'_>)>,
48) -> Result<()> {
49 if !b.is_empty() {
50 return Err(Error::SectionOrder);
51 }
52 let cp = b.checkpoint();
53 let flags = b.header().flags;
54 let res = (|| {
55 b.set_flags(Flags::default().with_opcode(Opcode::NOTIFY).with_aa(true));
56 let zone = zone.to_name();
57 b.push_question(zone, Rtype::SOA, class)?;
58 if let Some((ttl, soa)) = soa {
59 b.push_answer(zone, class, ttl, soa)?;
60 }
61 Ok(())
62 })();
63 if res.is_err() {
64 b.rollback(cp);
65 b.set_flags(flags);
66 }
67 res
68}
69
70pub fn build_response<B: OutBuf>(
73 b: &mut MessageBuilder<B>,
74 query: &NotifyMessage<'_>,
75) -> Result<()> {
76 if !b.is_empty() {
77 return Err(Error::SectionOrder);
78 }
79 let cp = b.checkpoint();
80 let (id, flags) = (b.header().id, b.header().flags);
81 b.set_id(query.message().id());
82 b.set_flags(
83 Flags::default()
84 .with_opcode(Opcode::NOTIFY)
85 .with_qr(true)
86 .with_aa(true),
87 );
88 let res = b.copy_question(&query.zone());
89 if res.is_err() {
90 b.rollback(cp);
91 b.set_id(id);
92 b.set_flags(flags);
93 }
94 res
95}
96
97#[derive(Clone, Copy, Debug)]
99pub struct NotifyMessage<'a> {
100 msg: Message<'a>,
101 zone: Question<'a>,
102}
103
104impl<'a> NotifyMessage<'a> {
105 pub fn new(msg: Message<'a>) -> Result<Self> {
110 if msg.flags().opcode() != Opcode::NOTIFY {
111 return Err(Error::WrongType);
112 }
113 if msg.header().qdcount != 1 {
114 return Err(Error::InvalidRdata);
115 }
116 let zone = msg.questions().next().ok_or(Error::UnexpectedEof)??;
117 Ok(NotifyMessage { msg, zone })
118 }
119
120 #[inline]
122 #[must_use]
123 pub const fn message(&self) -> Message<'a> {
124 self.msg
125 }
126
127 #[inline]
130 #[must_use]
131 pub const fn zone(&self) -> Question<'a> {
132 self.zone
133 }
134
135 #[inline]
137 #[must_use]
138 pub const fn is_response(&self) -> bool {
139 self.msg.flags().qr()
140 }
141
142 pub fn soa(&self) -> Result<Option<(Record<'a>, Soa<'a>)>> {
146 for rr in self.msg.answers() {
147 let rr = rr?;
148 if rr.rtype() == Soa::RTYPE && rr.name() == self.zone.name() {
149 return Ok(Some((rr, rr.data_as::<Soa<'a>>()?)));
150 }
151 }
152 Ok(None)
153 }
154
155 pub fn serial(&self) -> Result<Option<u32>> {
157 Ok(self.soa()?.map(|(_, soa)| soa.serial))
158 }
159}
160
161#[cfg(test)]
162mod tests {
163 use super::*;
164 use crate::{Name, NameBuf};
165 use std::vec::Vec;
166
167 fn soa(serial: u32) -> (NameBuf, NameBuf, u32) {
168 (
169 "ns1.example.com".parse().unwrap(),
170 "hostmaster.example.com".parse().unwrap(),
171 serial,
172 )
173 }
174
175 #[test]
176 fn query_with_hint() {
177 let zone: NameBuf = "example.com".parse().unwrap();
178 let (m, r, serial) = soa(2);
179 let data = Soa {
180 mname: m.as_name(),
181 rname: r.as_name(),
182 serial,
183 refresh: 7200,
184 retry: 3600,
185 expire: 1_209_600,
186 minimum: 300,
187 };
188 let mut buf = [0u8; 512];
189 let mut b = MessageBuilder::new(&mut buf).unwrap();
190 b.set_id(0x4ab7);
191 build_query(&mut b, &zone, Class::IN, Some((0, &data))).unwrap();
192 let wire = b.finish().to_vec();
193 assert_eq!(
196 wire,
197 crate::testutil::hex(
198 "4ab724000001000100000000076578616d706c6503636f6d0000060001c00c\
199 00060001000000000027036e7331c00c0a686f73746d6173746572c00c00000002\
200 00001c2000000e10001275000000012c"
201 )
202 );
203 let n = NotifyMessage::new(Message::parse_validated(&wire).unwrap()).unwrap();
204 assert!(!n.is_response());
205 assert!(n.message().flags().aa());
206 assert_eq!(n.zone().qtype(), Rtype::SOA);
207 assert_eq!(n.serial(), Ok(Some(2)));
208 let (rr, s) = n.soa().unwrap().unwrap();
209 assert_eq!(rr.ttl(), 0);
210 assert_eq!(s.mname, m.as_name());
211
212 let mut rbuf = [0u8; 512];
213 let mut rb = MessageBuilder::new(&mut rbuf).unwrap();
214 build_response(&mut rb, &n).unwrap();
215 let resp = rb.finish().to_vec();
216 let rn = NotifyMessage::new(Message::parse_validated(&resp).unwrap()).unwrap();
217 assert!(rn.is_response());
218 assert_eq!(rn.message().id(), 0x4ab7);
219 assert_eq!(rn.zone().name(), zone.as_name());
220 assert_eq!(rn.serial(), Ok(None));
221 }
222
223 #[test]
224 fn errors() {
225 let zone: NameBuf = "example.com".parse().unwrap();
226 let mut buf = [0u8; 512];
228 let mut b = MessageBuilder::new(&mut buf).unwrap();
229 b.push_question(&zone, Rtype::A, Class::IN).unwrap();
230 assert_eq!(
231 build_query(&mut b, &zone, Class::IN, None),
232 Err(Error::SectionOrder)
233 );
234 let mut small = [0u8; 20];
235 let mut b = MessageBuilder::new(&mut small).unwrap();
236 b.set_flags(Flags::default().with_rd(true));
237 assert_eq!(
238 build_query(&mut b, &zone, Class::IN, None),
239 Err(Error::BufferTooSmall)
240 );
241 assert_eq!(b.header().flags, Flags::default().with_rd(true));
242 assert!(b.is_empty());
243
244 let mut buf = [0u8; 512];
246 let mut b = MessageBuilder::new(&mut buf).unwrap();
247 b.push_question(&zone, Rtype::SOA, Class::IN).unwrap();
248 let q = b.finish().to_vec();
249 assert_eq!(
250 NotifyMessage::new(Message::parse(&q).unwrap()).err(),
251 Some(Error::WrongType)
252 );
253 let mut buf = [0u8; 512];
254 let mut b = MessageBuilder::new(&mut buf).unwrap();
255 b.set_flags(Flags::default().with_opcode(Opcode::NOTIFY));
256 let empty = b.finish().to_vec();
257 assert_eq!(
258 NotifyMessage::new(Message::parse(&empty).unwrap()).err(),
259 Some(Error::InvalidRdata)
260 );
261 let mut buf = [0u8; 512];
263 let mut b = MessageBuilder::new(&mut buf).unwrap();
264 let (m, r, serial) = soa(9);
265 let data = Soa {
266 mname: m.as_name(),
267 rname: r.as_name(),
268 serial,
269 refresh: 1,
270 retry: 1,
271 expire: 1,
272 minimum: 1,
273 };
274 build_query(&mut b, &zone, Class::IN, Some((60, &data))).unwrap();
275 let wire: Vec<u8> = b.finish().to_vec();
276 for end in 0..wire.len() {
277 if let Ok(m) = Message::parse(&wire[..end])
278 && let Ok(n) = NotifyMessage::new(m)
279 {
280 assert!(n.serial().is_err() || end == wire.len());
281 }
282 }
283 let mut buf = [0u8; 512];
285 let mut b = MessageBuilder::new(&mut buf).unwrap();
286 b.set_flags(Flags::default().with_opcode(Opcode::NOTIFY));
287 b.push_question(&zone, Rtype::SOA, Class::IN).unwrap();
288 b.push_answer(Name::ROOT, Class::IN, 0, &data).unwrap();
289 let other = b.finish().to_vec();
290 let n = NotifyMessage::new(Message::parse(&other).unwrap()).unwrap();
291 assert_eq!(n.serial(), Ok(None));
292 }
293}