1use std::collections::HashMap;
28use std::future::Future;
29use std::pin::Pin;
30use std::sync::{Arc, OnceLock};
31
32use parking_lot::RwLock;
33
34#[derive(Debug, Clone, PartialEq, Eq)]
36pub struct Endpoint {
37 pub host: String,
38 pub port: u16,
39}
40
41impl Endpoint {
42 pub fn new(host: impl Into<String>, port: u16) -> Self {
43 Self {
44 host: host.into(),
45 port,
46 }
47 }
48}
49
50type TargetKey = (String, String, u16);
52
53fn targets() -> &'static RwLock<HashMap<TargetKey, Endpoint>> {
54 static TARGETS: OnceLock<RwLock<HashMap<TargetKey, Endpoint>>> = OnceLock::new();
55 TARGETS.get_or_init(|| RwLock::new(HashMap::new()))
56}
57
58pub type ResolveFuture = Pin<Box<dyn Future<Output = Option<Endpoint>> + Send>>;
60
61pub trait InstanceResolver: Send + Sync {
63 fn resolve(
69 &self,
70 account_id: &str,
71 instance_id: &str,
72 port: u16,
73 source_groups: Vec<String>,
74 ) -> ResolveFuture;
75}
76
77fn instance_resolver() -> &'static RwLock<Option<Arc<dyn InstanceResolver>>> {
78 static RESOLVER: OnceLock<RwLock<Option<Arc<dyn InstanceResolver>>>> = OnceLock::new();
79 RESOLVER.get_or_init(|| RwLock::new(None))
80}
81
82pub fn set_instance_resolver(resolver: Arc<dyn InstanceResolver>) {
85 *instance_resolver().write() = Some(resolver);
86}
87
88pub fn register_target(account_id: &str, target_id: &str, port: u16, endpoint: Endpoint) {
90 targets().write().insert(
91 (account_id.to_string(), target_id.to_string(), port),
92 endpoint,
93 );
94}
95
96pub fn unregister_target(account_id: &str, target_id: &str) {
98 targets()
99 .write()
100 .retain(|(acct, id, _), _| !(acct == account_id && id == target_id));
101}
102
103pub fn registered_target(account_id: &str, target_id: &str, port: u16) -> Option<Endpoint> {
105 targets()
106 .read()
107 .get(&(account_id.to_string(), target_id.to_string(), port))
108 .cloned()
109}
110
111pub async fn resolve_target(
122 account_id: &str,
123 target_id: &str,
124 port: u16,
125 sibling_host: &str,
126 source_groups: &[String],
127) -> Endpoint {
128 if let Some(ep) = registered_target(account_id, target_id, port) {
129 return ep;
130 }
131 if target_id.starts_with("i-") {
132 let resolver = instance_resolver().read().clone();
133 if let Some(resolver) = resolver {
134 if let Some(ep) = resolver
135 .resolve(account_id, target_id, port, source_groups.to_vec())
136 .await
137 {
138 return ep;
139 }
140 }
141 }
142 Endpoint::new(fallback_host(target_id, sibling_host), port)
143}
144
145pub fn fallback_host(target_id: &str, sibling_host: &str) -> String {
147 if target_id.starts_with("i-") || target_id == "127.0.0.1" {
148 sibling_host.to_string()
149 } else {
150 target_id.to_string()
151 }
152}
153
154pub fn account_of_arn(arn: &str) -> Option<&str> {
156 arn.split(':').nth(4).filter(|a| !a.is_empty())
157}
158
159fn container_ports() -> &'static RwLock<HashMap<u16, u16>> {
160 static PORTS: OnceLock<RwLock<HashMap<u16, u16>>> = OnceLock::new();
161 PORTS.get_or_init(|| RwLock::new(HashMap::new()))
162}
163
164pub fn register_container_port(host_port: u16, container_port: u16) {
167 container_ports().write().insert(host_port, container_port);
168}
169
170pub fn unregister_container_port(host_port: u16) {
172 container_ports().write().remove(&host_port);
173}
174
175pub fn container_port_for(host_port: u16) -> Option<u16> {
178 container_ports().read().get(&host_port).copied()
179}
180
181#[cfg(test)]
182mod tests {
183 use super::*;
184
185 #[test]
186 fn fallback_host_keeps_historical_routing() {
187 assert_eq!(
188 fallback_host("i-0abc", "host.docker.internal"),
189 "host.docker.internal"
190 );
191 assert_eq!(fallback_host("127.0.0.1", "127.0.0.1"), "127.0.0.1");
192 assert_eq!(
193 fallback_host("10.0.4.7", "host.docker.internal"),
194 "10.0.4.7"
195 );
196 }
197
198 #[tokio::test]
199 async fn registered_endpoint_wins_over_the_verbatim_ip() {
200 let acct = "111111111111";
201 register_target(acct, "10.9.8.7", 80, Endpoint::new("127.0.0.1", 49153));
202 assert_eq!(
203 resolve_target(acct, "10.9.8.7", 80, "127.0.0.1", &[]).await,
204 Endpoint::new("127.0.0.1", 49153)
205 );
206 assert_eq!(
208 resolve_target(acct, "10.9.8.7", 81, "127.0.0.1", &[]).await,
209 Endpoint::new("10.9.8.7", 81)
210 );
211 assert_eq!(
212 resolve_target("222222222222", "10.9.8.7", 80, "127.0.0.1", &[]).await,
213 Endpoint::new("10.9.8.7", 80)
214 );
215 unregister_target(acct, "10.9.8.7");
216 assert_eq!(
217 resolve_target(acct, "10.9.8.7", 80, "127.0.0.1", &[]).await,
218 Endpoint::new("10.9.8.7", 80)
219 );
220 }
221
222 struct FixedResolver;
223 impl InstanceResolver for FixedResolver {
224 fn resolve(
225 &self,
226 _account: &str,
227 instance_id: &str,
228 port: u16,
229 source_groups: Vec<String>,
230 ) -> ResolveFuture {
231 let known = instance_id == "i-resolvable" && source_groups == ["sg-alb"];
233 Box::pin(async move { known.then(|| Endpoint::new("127.0.0.1", port + 1000)) })
234 }
235 }
236
237 #[tokio::test]
238 async fn instance_resolver_publishes_instance_ports() {
239 set_instance_resolver(Arc::new(FixedResolver));
240 assert_eq!(
241 resolve_target(
242 "333333333333",
243 "i-resolvable",
244 8080,
245 "host.docker.internal",
246 &["sg-alb".to_string()]
247 )
248 .await,
249 Endpoint::new("127.0.0.1", 9080)
250 );
251 assert_eq!(
253 resolve_target(
254 "333333333333",
255 "i-unknown",
256 8080,
257 "host.docker.internal",
258 &[]
259 )
260 .await,
261 Endpoint::new("host.docker.internal", 8080)
262 );
263 }
264
265 #[test]
266 fn account_is_read_from_the_arn() {
267 assert_eq!(
268 account_of_arn(
269 "arn:aws:elasticloadbalancing:us-east-1:123456789012:targetgroup/tg/abc"
270 ),
271 Some("123456789012")
272 );
273 assert_eq!(account_of_arn("not-an-arn"), None);
274 }
275
276 #[test]
277 fn container_port_map_round_trips() {
278 register_container_port(40001, 40002);
279 assert_eq!(container_port_for(40001), Some(40002));
280 unregister_container_port(40001);
281 assert_eq!(container_port_for(40001), None);
282 }
283}