use crate::utils::ChainError;
use std::fmt;
use std::net::{IpAddr, Ipv4Addr};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum ListenOn {
All,
#[default]
Localhost,
Address(IpAddr),
}
impl ListenOn {
#[must_use]
pub fn ip(&self) -> IpAddr {
match self {
ListenOn::All => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
ListenOn::Localhost => IpAddr::V4(Ipv4Addr::LOCALHOST),
ListenOn::Address(address) => *address,
}
}
#[must_use]
pub fn is_public(&self) -> bool {
!self.ip().is_loopback()
}
pub fn parse(raw: &str, field: &str) -> Result<Self, ChainError> {
match raw.trim().to_ascii_lowercase().as_str() {
"localhost" => return Ok(ListenOn::Localhost),
"all" => return Ok(ListenOn::All),
_ => {}
}
let address: IpAddr = raw.trim().parse().map_err(|_| ChainError::Validation {
field: field.to_string(),
reason: format!(
"must be an IP address, or `localhost` or `all`, got {:?}",
raw.trim()
),
})?;
Ok(match address {
address if address == IpAddr::V4(Ipv4Addr::UNSPECIFIED) => ListenOn::All,
address if address == IpAddr::V4(Ipv4Addr::LOCALHOST) => ListenOn::Localhost,
address => ListenOn::Address(address),
})
}
}
impl From<ListenOn> for String {
fn from(listen_on: ListenOn) -> Self {
listen_on.ip().to_string()
}
}
impl fmt::Display for ListenOn {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.ip())
}
}
pub const BIND_ADDRESS_VAR: &str = "OCS_BIND_ADDRESS";
pub const PORT_VAR: &str = "OCS_PORT";
pub const DEFAULT_BIND_ADDRESS: ListenOn = ListenOn::Localhost;
pub const DEFAULT_PORT: u16 = 7070;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ServerConfig {
pub address: ListenOn,
pub port: u16,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
address: DEFAULT_BIND_ADDRESS,
port: DEFAULT_PORT,
}
}
}
impl ServerConfig {
pub fn from_env() -> Result<Self, ChainError> {
let address = match super::read_var(BIND_ADDRESS_VAR) {
Some(raw) => ListenOn::parse(&raw, BIND_ADDRESS_VAR)?,
None => DEFAULT_BIND_ADDRESS,
};
let port = match super::read_var(PORT_VAR) {
Some(raw) => raw
.parse::<u16>()
.ok()
.filter(|port| *port > 0)
.ok_or_else(|| ChainError::Validation {
field: PORT_VAR.to_string(),
reason: format!("must be an integer between 1 and 65535, got {raw:?}"),
})?,
None => DEFAULT_PORT,
};
Ok(Self { address, port })
}
}
#[cfg(test)]
mod tests {
use super::*;
use once_cell::sync::Lazy;
use std::net::{Ipv6Addr, SocketAddr};
use std::sync::Mutex;
static ENV_MUTEX: Lazy<Mutex<()>> = Lazy::new(|| Mutex::new(()));
fn set_var(name: &str, value: &str) {
#[allow(unused_unsafe)]
unsafe {
std::env::set_var(name, value);
}
}
fn remove_var(name: &str) {
#[allow(unused_unsafe)]
unsafe {
std::env::remove_var(name);
}
}
fn clear() {
remove_var(BIND_ADDRESS_VAR);
remove_var(PORT_VAR);
}
#[test]
fn test_neither_variable_binds_loopback_on_the_default_port() {
let _guard = match ENV_MUTEX.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
clear();
let config = match ServerConfig::from_env() {
Ok(config) => config,
Err(error) => panic!("an empty environment must resolve: {error}"),
};
assert_eq!(config.address, ListenOn::Localhost);
assert_eq!(config.address.ip(), IpAddr::V4(Ipv4Addr::LOCALHOST));
assert_eq!(config.port, 7070);
assert!(
!config.address.is_public(),
"an unauthenticated service must not default to reachable"
);
}
#[test]
fn test_both_variables_are_honoured() {
let _guard = match ENV_MUTEX.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
for (raw, expected) in [
("0.0.0.0", ListenOn::All),
("all", ListenOn::All),
("127.0.0.1", ListenOn::Localhost),
("localhost", ListenOn::Localhost),
(
"10.1.2.3",
ListenOn::Address(IpAddr::V4(Ipv4Addr::new(10, 1, 2, 3))),
),
] {
set_var(BIND_ADDRESS_VAR, raw);
set_var(PORT_VAR, "9001");
match ServerConfig::from_env() {
Ok(config) => {
assert_eq!(config.address, expected, "for {raw:?}");
assert_eq!(config.port, 9001, "for {raw:?}");
}
Err(error) => panic!("{raw:?} must resolve: {error}"),
}
}
clear();
}
#[test]
fn test_blank_variables_fall_back_to_the_defaults() {
let _guard = match ENV_MUTEX.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
set_var(BIND_ADDRESS_VAR, " ");
set_var(PORT_VAR, "");
match ServerConfig::from_env() {
Ok(config) => assert_eq!(config, ServerConfig::default()),
Err(error) => panic!("a blank value is unset, not invalid: {error}"),
}
clear();
}
#[test]
fn test_an_unusable_address_names_the_variable() {
let _guard = match ENV_MUTEX.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
remove_var(PORT_VAR);
for raw in ["not-an-address", "999.1.1.1", "127.0.0.1:7070"] {
set_var(BIND_ADDRESS_VAR, raw);
match ServerConfig::from_env() {
Ok(config) => panic!("{raw:?} must not resolve, got {config:?}"),
Err(ChainError::Validation { field, reason }) => {
assert_eq!(field, BIND_ADDRESS_VAR, "for {raw:?}");
assert!(
reason.contains(raw.trim()),
"the reason must quote it: {reason}"
);
}
Err(error) => panic!("expected a validation failure, got {error:?}"),
}
}
clear();
}
#[test]
fn test_an_unusable_port_names_the_variable() {
let _guard = match ENV_MUTEX.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
remove_var(BIND_ADDRESS_VAR);
for raw in ["0", "-1", "70000", "http", "7070.5"] {
set_var(PORT_VAR, raw);
match ServerConfig::from_env() {
Ok(config) => panic!("{raw:?} must not resolve, got {config:?}"),
Err(ChainError::Validation { field, .. }) => {
assert_eq!(field, PORT_VAR, "for {raw:?}");
}
Err(error) => panic!("expected a validation failure, got {error:?}"),
}
}
clear();
}
#[test]
fn test_every_variant_renders_as_the_address_it_binds() {
for (listen_on, expected) in [
(ListenOn::All, "0.0.0.0"),
(ListenOn::Localhost, "127.0.0.1"),
(
ListenOn::Address(IpAddr::V4(Ipv4Addr::new(10, 1, 2, 3))),
"10.1.2.3",
),
(
ListenOn::Address(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1))),
"::1",
),
] {
assert_eq!(listen_on.ip().to_string(), expected);
assert_eq!(listen_on.to_string(), expected, "Display for {listen_on:?}");
assert_eq!(
String::from(listen_on),
expected,
"String::from for {listen_on:?}"
);
}
}
#[test]
fn test_the_resolved_configuration_produces_a_usable_bind_address() {
for (address, port, expected) in [
(ListenOn::All, 7070_u16, "0.0.0.0:7070"),
(ListenOn::Localhost, 9001, "127.0.0.1:9001"),
(
ListenOn::Address(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1))),
9001,
"[::1]:9001",
),
] {
let config = ServerConfig { address, port };
assert_eq!(
SocketAddr::new(config.address.ip(), config.port).to_string(),
expected
);
}
}
#[test]
fn test_only_loopback_is_not_public() {
assert!(!ListenOn::Localhost.is_public());
assert!(!ListenOn::Address(IpAddr::V6(Ipv6Addr::LOCALHOST)).is_public());
assert!(ListenOn::All.is_public());
assert!(ListenOn::Address(IpAddr::V4(Ipv4Addr::new(10, 1, 2, 3))).is_public());
}
#[test]
fn test_an_ipv6_address_is_accepted() {
match ListenOn::parse("::1", BIND_ADDRESS_VAR) {
Ok(listen_on) => {
assert_eq!(
listen_on,
ListenOn::Address(IpAddr::V6(Ipv6Addr::LOCALHOST))
);
assert!(!listen_on.is_public());
}
Err(error) => panic!("an IPv6 literal must parse: {error}"),
}
}
#[test]
fn test_two_instances_resolve_to_different_ports() {
let _guard = match ENV_MUTEX.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
set_var(PORT_VAR, "7071");
let first = ServerConfig::from_env();
set_var(PORT_VAR, "7072");
let second = ServerConfig::from_env();
clear();
match (first, second) {
(Ok(first), Ok(second)) => {
assert_eq!(first.port, 7071);
assert_eq!(second.port, 7072);
assert_ne!(first, second, "two shards must not fight over one port");
}
other => panic!("both must resolve, got {other:?}"),
}
}
}