Skip to content

Commit 8a7b2f6

Browse files
authored
fix(network): preserve DNS answer order
Signed-off-by: John Myers <9696606+johntmyers@users.noreply.github.com>
1 parent 4cb6955 commit 8a7b2f6

2 files changed

Lines changed: 56 additions & 6 deletions

File tree

crates/openshell-supervisor-network/src/policy_dns/resolver.rs

Lines changed: 26 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@ use super::name::NormalizedName;
1111
use hickory_proto::op::{Message, MessageType, OpCode, Query, ResponseCode};
1212
use hickory_proto::rr::{Name, RData, RecordType};
1313
use openshell_core::net::connect_tcp_nodelay_best_effort;
14-
use std::collections::{BTreeMap, BTreeSet};
14+
use std::collections::{BTreeMap, BTreeSet, HashSet};
1515
use std::net::{IpAddr, SocketAddr};
1616
use std::time::Duration;
1717
use tokio::io::{AsyncReadExt, AsyncWriteExt};
@@ -204,8 +204,7 @@ impl TrustedResolver for SocketTrustedResolver {
204204
.iter()
205205
.map(|(address, _)| *address)
206206
.collect::<Vec<_>>();
207-
addresses.sort_unstable();
208-
addresses.dedup();
207+
retain_first_addresses(&mut addresses);
209208
addresses.truncate(MAX_RETAINED_ADDRESSES);
210209
let address_ttl = records.iter().map(|(_, ttl)| *ttl).min().unwrap_or(1);
211210
return Ok(TrustedAnswer {
@@ -239,6 +238,11 @@ impl TrustedResolver for SocketTrustedResolver {
239238
}
240239
}
241240

241+
fn retain_first_addresses(addresses: &mut Vec<IpAddr>) {
242+
let mut seen = HashSet::new();
243+
addresses.retain(|address| seen.insert(*address));
244+
}
245+
242246
fn parse_response(wire: &[u8], id: u16, query: &Query) -> Result<Message, ResolveError> {
243247
if wire.len() > MAX_DNS_MESSAGE_BYTES {
244248
return Err(ResolveError::Oversized);
@@ -311,6 +315,25 @@ mod tests {
311315
use openshell_core::net::set_tcp_nodelay_best_effort;
312316
use tokio::net::TcpListener;
313317

318+
#[test]
319+
fn resolver_address_deduplication_preserves_answer_order() {
320+
let mut addresses = vec![
321+
"203.0.113.20".parse().unwrap(),
322+
"203.0.113.10".parse().unwrap(),
323+
"203.0.113.20".parse().unwrap(),
324+
];
325+
326+
retain_first_addresses(&mut addresses);
327+
328+
assert_eq!(
329+
addresses,
330+
vec![
331+
"203.0.113.20".parse::<IpAddr>().unwrap(),
332+
"203.0.113.10".parse::<IpAddr>().unwrap(),
333+
]
334+
);
335+
}
336+
314337
#[test]
315338
fn answer_parser_keeps_only_requested_family_and_bounds_are_constants() {
316339
let owner = Name::from_ascii("db.example.").unwrap();

crates/openshell-supervisor-network/src/policy_dns/store.rs

Lines changed: 30 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ use crate::proxy::destination::{
99
DestinationRequest, DestinationValidationPlan, UpstreamConnector, build_pinned_validation_plan,
1010
validate_destination,
1111
};
12-
use std::collections::{BTreeMap, BTreeSet};
12+
use std::collections::{BTreeMap, BTreeSet, HashSet};
1313
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
1414
use std::ops::RangeInclusive;
1515
use std::sync::RwLock;
@@ -89,14 +89,14 @@ impl MappingLookup {
8989
&self,
9090
endpoint_id: &PolicyEndpointId,
9191
) -> Result<UpstreamConnector, MappingLookupError> {
92+
let mut seen = HashSet::new();
9293
let addresses = self
9394
.record
9495
.contracts
9596
.iter()
9697
.filter(|contract| contract.port == self.port && &contract.endpoint_id == endpoint_id)
9798
.flat_map(|contract| contract.pinned_addresses.iter().copied())
98-
.collect::<BTreeSet<_>>()
99-
.into_iter()
99+
.filter(|address| seen.insert(*address))
100100
.collect::<Vec<_>>();
101101
if addresses.is_empty() {
102102
return Err(MappingLookupError::EndpointMismatch);
@@ -655,4 +655,31 @@ mod tests {
655655
let connector = lookup.connector_for(&endpoint).await.unwrap();
656656
assert_eq!(connector.addrs(), &["203.0.113.8:5432".parse().unwrap()]);
657657
}
658+
659+
#[tokio::test]
660+
async fn connector_preserves_resolver_address_order_while_deduplicating() {
661+
let store = store(1);
662+
let now = Instant::now();
663+
let mut request = request("must-not-resolve.invalid", 1, Duration::from_secs(5));
664+
request.contracts[0].pinned_addresses = vec![
665+
"203.0.113.20".parse().unwrap(),
666+
"203.0.113.10".parse().unwrap(),
667+
"203.0.113.20".parse().unwrap(),
668+
];
669+
let record = store.publish(request, 1, now).unwrap();
670+
let lookup = store
671+
.lookup(record.synthetic_address, 5432, 1, now)
672+
.unwrap();
673+
let endpoint = lookup.endpoint_ids().next().unwrap().clone();
674+
675+
let connector = lookup.connector_for(&endpoint).await.unwrap();
676+
677+
assert_eq!(
678+
connector.addrs(),
679+
&[
680+
"203.0.113.20:5432".parse().unwrap(),
681+
"203.0.113.10:5432".parse().unwrap(),
682+
]
683+
);
684+
}
658685
}

0 commit comments

Comments
 (0)