use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, OnceLock};
use parking_lot::RwLock;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Endpoint {
pub host: String,
pub port: u16,
}
impl Endpoint {
pub fn new(host: impl Into<String>, port: u16) -> Self {
Self {
host: host.into(),
port,
}
}
}
type TargetKey = (String, String, u16);
fn targets() -> &'static RwLock<HashMap<TargetKey, Endpoint>> {
static TARGETS: OnceLock<RwLock<HashMap<TargetKey, Endpoint>>> = OnceLock::new();
TARGETS.get_or_init(|| RwLock::new(HashMap::new()))
}
pub type ResolveFuture = Pin<Box<dyn Future<Output = Option<Endpoint>> + Send>>;
pub trait InstanceResolver: Send + Sync {
fn resolve(
&self,
account_id: &str,
instance_id: &str,
port: u16,
source_groups: Vec<String>,
) -> ResolveFuture;
}
fn instance_resolver() -> &'static RwLock<Option<Arc<dyn InstanceResolver>>> {
static RESOLVER: OnceLock<RwLock<Option<Arc<dyn InstanceResolver>>>> = OnceLock::new();
RESOLVER.get_or_init(|| RwLock::new(None))
}
pub fn set_instance_resolver(resolver: Arc<dyn InstanceResolver>) {
*instance_resolver().write() = Some(resolver);
}
pub fn register_target(account_id: &str, target_id: &str, port: u16, endpoint: Endpoint) {
targets().write().insert(
(account_id.to_string(), target_id.to_string(), port),
endpoint,
);
}
pub fn unregister_target(account_id: &str, target_id: &str) {
targets()
.write()
.retain(|(acct, id, _), _| !(acct == account_id && id == target_id));
}
pub fn registered_target(account_id: &str, target_id: &str, port: u16) -> Option<Endpoint> {
targets()
.read()
.get(&(account_id.to_string(), target_id.to_string(), port))
.cloned()
}
pub async fn resolve_target(
account_id: &str,
target_id: &str,
port: u16,
sibling_host: &str,
source_groups: &[String],
) -> Endpoint {
if let Some(ep) = registered_target(account_id, target_id, port) {
return ep;
}
if target_id.starts_with("i-") {
let resolver = instance_resolver().read().clone();
if let Some(resolver) = resolver {
if let Some(ep) = resolver
.resolve(account_id, target_id, port, source_groups.to_vec())
.await
{
return ep;
}
}
}
Endpoint::new(fallback_host(target_id, sibling_host), port)
}
pub fn fallback_host(target_id: &str, sibling_host: &str) -> String {
if target_id.starts_with("i-") || target_id == "127.0.0.1" {
sibling_host.to_string()
} else {
target_id.to_string()
}
}
pub fn account_of_arn(arn: &str) -> Option<&str> {
arn.split(':').nth(4).filter(|a| !a.is_empty())
}
fn container_ports() -> &'static RwLock<HashMap<u16, u16>> {
static PORTS: OnceLock<RwLock<HashMap<u16, u16>>> = OnceLock::new();
PORTS.get_or_init(|| RwLock::new(HashMap::new()))
}
pub fn register_container_port(host_port: u16, container_port: u16) {
container_ports().write().insert(host_port, container_port);
}
pub fn unregister_container_port(host_port: u16) {
container_ports().write().remove(&host_port);
}
pub fn container_port_for(host_port: u16) -> Option<u16> {
container_ports().read().get(&host_port).copied()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fallback_host_keeps_historical_routing() {
assert_eq!(
fallback_host("i-0abc", "host.docker.internal"),
"host.docker.internal"
);
assert_eq!(fallback_host("127.0.0.1", "127.0.0.1"), "127.0.0.1");
assert_eq!(
fallback_host("10.0.4.7", "host.docker.internal"),
"10.0.4.7"
);
}
#[tokio::test]
async fn registered_endpoint_wins_over_the_verbatim_ip() {
let acct = "111111111111";
register_target(acct, "10.9.8.7", 80, Endpoint::new("127.0.0.1", 49153));
assert_eq!(
resolve_target(acct, "10.9.8.7", 80, "127.0.0.1", &[]).await,
Endpoint::new("127.0.0.1", 49153)
);
assert_eq!(
resolve_target(acct, "10.9.8.7", 81, "127.0.0.1", &[]).await,
Endpoint::new("10.9.8.7", 81)
);
assert_eq!(
resolve_target("222222222222", "10.9.8.7", 80, "127.0.0.1", &[]).await,
Endpoint::new("10.9.8.7", 80)
);
unregister_target(acct, "10.9.8.7");
assert_eq!(
resolve_target(acct, "10.9.8.7", 80, "127.0.0.1", &[]).await,
Endpoint::new("10.9.8.7", 80)
);
}
struct FixedResolver;
impl InstanceResolver for FixedResolver {
fn resolve(
&self,
_account: &str,
instance_id: &str,
port: u16,
source_groups: Vec<String>,
) -> ResolveFuture {
let known = instance_id == "i-resolvable" && source_groups == ["sg-alb"];
Box::pin(async move { known.then(|| Endpoint::new("127.0.0.1", port + 1000)) })
}
}
#[tokio::test]
async fn instance_resolver_publishes_instance_ports() {
set_instance_resolver(Arc::new(FixedResolver));
assert_eq!(
resolve_target(
"333333333333",
"i-resolvable",
8080,
"host.docker.internal",
&["sg-alb".to_string()]
)
.await,
Endpoint::new("127.0.0.1", 9080)
);
assert_eq!(
resolve_target(
"333333333333",
"i-unknown",
8080,
"host.docker.internal",
&[]
)
.await,
Endpoint::new("host.docker.internal", 8080)
);
}
#[test]
fn account_is_read_from_the_arn() {
assert_eq!(
account_of_arn(
"arn:aws:elasticloadbalancing:us-east-1:123456789012:targetgroup/tg/abc"
),
Some("123456789012")
);
assert_eq!(account_of_arn("not-an-arn"), None);
}
#[test]
fn container_port_map_round_trips() {
register_container_port(40001, 40002);
assert_eq!(container_port_for(40001), Some(40002));
unregister_container_port(40001);
assert_eq!(container_port_for(40001), None);
}
}