use std::error::Error;
use std::time::Duration;
use argh::FromArgs;
use config::{Config, Environment, File, FileFormat};
use serde::Deserialize;
use tephra::log::set::SegmentConfig;
use tephra::read::ReadConfig;
use tephra::writer::WriterConfig;
use tephra_proto::DEFAULT_MAX_FRAME_LEN;
use tephra_server::ServerConfig;
#[derive(Debug, FromArgs)]
pub struct Args {
#[argh(option, short = 'c')]
pub config: Option<String>,
#[argh(option, short = 'b')]
pub bind: Option<String>,
#[argh(option, short = 'd')]
pub data_dir: Option<String>,
#[argh(option, short = 'l')]
pub log: Option<String>,
#[argh(switch)]
pub healthcheck: bool,
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct Settings {
pub bind: String,
pub data_dir: String,
pub log: Option<String>,
pub segment: SegmentSettings,
pub writer: WriterSettings,
pub read: ReadSettings,
pub server: ServerSettings,
pub metrics: MetricsSettings,
pub tls: TlsSettings,
pub auth: AuthSettings,
}
impl Default for Settings {
fn default() -> Self {
Settings {
bind: "127.0.0.1:9000".to_string(),
data_dir: "tephra-data".to_string(),
log: None,
segment: SegmentSettings::default(),
writer: WriterSettings::default(),
read: ReadSettings::default(),
server: ServerSettings::default(),
metrics: MetricsSettings::default(),
tls: TlsSettings::default(),
auth: AuthSettings::default(),
}
}
}
#[derive(Debug, Default, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct MetricsSettings {
pub bind: Option<String>,
}
#[derive(Debug, Default, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct TlsSettings {
pub cert: Option<String>,
pub key: Option<String>,
}
#[derive(Debug, Default, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct AuthSettings {
pub tokens: Vec<TokenSettings>,
pub allow_insecure: bool,
}
#[derive(Debug, Default, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct TokenSettings {
pub token: String,
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct SegmentSettings {
pub size: usize,
}
impl Default for SegmentSettings {
fn default() -> Self {
SegmentSettings {
size: 256 * 1024 * 1024,
}
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct WriterSettings {
pub queue_capacity: usize,
pub max_batch_records: usize,
pub max_batch_bytes: usize,
pub tips_window: u64,
pub condition_force_scan: bool,
}
impl Default for WriterSettings {
fn default() -> Self {
WriterSettings {
queue_capacity: 16384,
max_batch_records: 2048,
max_batch_bytes: 8 * 1024 * 1024,
tips_window: 1_000_000,
condition_force_scan: false,
}
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct ReadSettings {
pub scan_bias: u32,
}
impl Default for ReadSettings {
fn default() -> Self {
ReadSettings { scan_bias: 4 }
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct ServerSettings {
pub max_frame_len: u32,
pub reads: ReadsSettings,
pub subscriptions: SubscriptionsSettings,
pub backpressure: BackpressureSettings,
pub limits: LimitsSettings,
pub keepalive: KeepaliveSettings,
pub timeouts: TimeoutsSettings,
}
impl Default for ServerSettings {
fn default() -> Self {
ServerSettings {
max_frame_len: DEFAULT_MAX_FRAME_LEN,
reads: ReadsSettings::default(),
subscriptions: SubscriptionsSettings::default(),
backpressure: BackpressureSettings::default(),
limits: LimitsSettings::default(),
keepalive: KeepaliveSettings::default(),
timeouts: TimeoutsSettings::default(),
}
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct ReadsSettings {
pub batch_events: usize,
pub batch_bytes: usize,
pub worker_threads: usize,
}
impl Default for ReadsSettings {
fn default() -> Self {
ReadsSettings {
batch_events: 1024,
batch_bytes: 512 * 1024,
worker_threads: 0,
}
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct SubscriptionsSettings {
pub wait_tick_ms: u64,
pub max_concurrent: usize,
}
impl Default for SubscriptionsSettings {
fn default() -> Self {
SubscriptionsSettings {
wait_tick_ms: 250,
max_concurrent: 64,
}
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct BackpressureSettings {
pub max_inflight_per_conn: usize,
pub frame_queue_depth: usize,
}
impl Default for BackpressureSettings {
fn default() -> Self {
BackpressureSettings {
max_inflight_per_conn: 256,
frame_queue_depth: 256,
}
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct LimitsSettings {
pub max_connections: usize,
}
impl Default for LimitsSettings {
fn default() -> Self {
LimitsSettings {
max_connections: 1024,
}
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct KeepaliveSettings {
pub idle_secs: u64,
pub interval_secs: u64,
}
impl Default for KeepaliveSettings {
fn default() -> Self {
KeepaliveSettings {
idle_secs: 60,
interval_secs: 15,
}
}
}
#[derive(Debug, PartialEq, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct TimeoutsSettings {
pub incomplete_frame_secs: u64,
pub handshake_secs: u64,
pub idle_secs: u64,
}
impl Default for TimeoutsSettings {
fn default() -> Self {
TimeoutsSettings {
incomplete_frame_secs: 30,
handshake_secs: 0,
idle_secs: 0,
}
}
}
impl Settings {
pub fn segment_config(&self) -> SegmentConfig {
SegmentConfig::new(self.segment.size)
}
pub fn writer_config(&self) -> WriterConfig {
WriterConfig {
queue_capacity: self.writer.queue_capacity,
max_batch_records: self.writer.max_batch_records,
max_batch_bytes: self.writer.max_batch_bytes,
tips_window: self.writer.tips_window,
verify_tips: false,
condition_force_scan: self.writer.condition_force_scan,
read: ReadConfig {
scan_bias: self.read.scan_bias,
},
}
}
pub fn server_config(&self) -> ServerConfig {
let server = &self.server;
ServerConfig {
max_frame_len: server.max_frame_len,
read_batch_events: server.reads.batch_events,
read_batch_bytes: server.reads.batch_bytes,
subscribe_wait_tick: Duration::from_millis(server.subscriptions.wait_tick_ms),
max_inflight_requests_per_conn: server.backpressure.max_inflight_per_conn,
max_concurrent_subscriptions: server.subscriptions.max_concurrent,
read_worker_threads: server.reads.worker_threads,
frame_queue_depth: server.backpressure.frame_queue_depth,
keepalive_idle: Duration::from_secs(server.keepalive.idle_secs),
keepalive_interval: Duration::from_secs(server.keepalive.interval_secs),
max_connections: server.limits.max_connections,
incomplete_frame_timeout: Duration::from_secs(server.timeouts.incomplete_frame_secs),
handshake_timeout: Duration::from_secs(server.timeouts.handshake_secs),
idle_timeout: Duration::from_secs(server.timeouts.idle_secs),
}
}
fn validate(&self) -> Result<(), String> {
if self.writer.queue_capacity == 0 {
return Err("writer.queue_capacity must be at least 1".to_string());
}
if self.writer.max_batch_records == 0 {
return Err("writer.max_batch_records must be at least 1".to_string());
}
if self.server.subscriptions.wait_tick_ms == 0 {
return Err("server.subscriptions.wait_tick_ms must be at least 1".to_string());
}
if self.server.backpressure.max_inflight_per_conn == 0 {
return Err("server.backpressure.max_inflight_per_conn must be at least 1".to_string());
}
if self.server.subscriptions.max_concurrent == 0 {
return Err("server.subscriptions.max_concurrent must be at least 1".to_string());
}
if self.server.backpressure.frame_queue_depth == 0 {
return Err("server.backpressure.frame_queue_depth must be at least 1".to_string());
}
if self.server.keepalive.idle_secs == 0 {
return Err("server.keepalive.idle_secs must be at least 1".to_string());
}
if self.server.keepalive.interval_secs == 0 {
return Err("server.keepalive.interval_secs must be at least 1".to_string());
}
if self.tls.cert.is_some() != self.tls.key.is_some() {
return Err("tls.cert and tls.key must be set together".to_string());
}
if self.auth.tokens.iter().any(|t| t.token.is_empty()) {
return Err("auth.tokens entries must have a non-empty token".to_string());
}
let tls_enabled = self.tls.cert.is_some() && self.tls.key.is_some();
if !self.auth.tokens.is_empty() && !tls_enabled && !self.auth.allow_insecure {
return Err(
"auth.tokens require tls; set tls.cert and tls.key, or auth.allow_insecure = true"
.to_string(),
);
}
Ok(())
}
pub fn auth_tokens(&self) -> Option<Vec<String>> {
if self.auth.tokens.is_empty() {
return None;
}
Some(self.auth.tokens.iter().map(|t| t.token.clone()).collect())
}
pub fn first_auth_token(&self) -> Option<&str> {
self.auth.tokens.first().map(|t| t.token.as_str())
}
}
pub fn load(args: &Args) -> Result<Settings, Box<dyn Error>> {
let mut builder = Config::builder();
if let Some(path) = &args.config {
builder = builder.add_source(File::new(path, FileFormat::Toml).required(true));
}
builder = builder.add_source(
Environment::with_prefix("TEPHRA")
.prefix_separator("__")
.separator("__")
.try_parsing(true),
);
let mut settings: Settings = builder.build()?.try_deserialize()?;
if let Some(bind) = &args.bind {
settings.bind = bind.clone();
}
if let Some(data_dir) = &args.data_dir {
settings.data_dir = data_dir.clone();
}
if args.log.is_some() {
settings.log = args.log.clone();
}
settings.validate()?;
Ok(settings)
}
#[cfg(test)]
mod tests {
use super::*;
fn no_args() -> Args {
Args {
config: None,
bind: None,
data_dir: None,
log: None,
healthcheck: false,
}
}
#[test]
fn defaults_match_the_library_defaults() {
let settings = load(&no_args()).unwrap();
let writer = settings.writer_config();
let library_default = WriterConfig::default();
assert_eq!(writer.queue_capacity, library_default.queue_capacity);
assert_eq!(writer.max_batch_records, library_default.max_batch_records);
assert_eq!(writer.max_batch_bytes, library_default.max_batch_bytes);
assert_eq!(writer.tips_window, library_default.tips_window);
assert!(!writer.verify_tips);
assert_eq!(writer.read.scan_bias, ReadConfig::default().scan_bias);
let server = settings.server_config();
let server_default = ServerConfig::default();
assert_eq!(server.max_frame_len, server_default.max_frame_len);
assert_eq!(server.read_batch_events, server_default.read_batch_events);
assert_eq!(server.read_batch_bytes, server_default.read_batch_bytes);
assert_eq!(
server.subscribe_wait_tick,
server_default.subscribe_wait_tick
);
assert_eq!(
server.max_inflight_requests_per_conn,
server_default.max_inflight_requests_per_conn
);
assert_eq!(
server.max_concurrent_subscriptions,
server_default.max_concurrent_subscriptions
);
assert_eq!(
server.read_worker_threads,
server_default.read_worker_threads
);
assert_eq!(server.frame_queue_depth, server_default.frame_queue_depth);
assert_eq!(server.keepalive_idle, server_default.keepalive_idle);
assert_eq!(server.keepalive_interval, server_default.keepalive_interval);
assert_eq!(server.max_connections, server_default.max_connections);
assert_eq!(
server.incomplete_frame_timeout,
server_default.incomplete_frame_timeout
);
assert_eq!(server.handshake_timeout, server_default.handshake_timeout);
assert_eq!(server.idle_timeout, server_default.idle_timeout);
assert_eq!(settings.bind, "127.0.0.1:9000");
assert_eq!(settings.data_dir, "tephra-data");
}
#[test]
fn example_toml_mirrors_the_defaults() {
let path = concat!(env!("CARGO_MANIFEST_DIR"), "/tephra.example.toml");
let settings: Settings = Config::builder()
.add_source(File::new(path, FileFormat::Toml).required(true))
.build()
.unwrap()
.try_deserialize()
.unwrap();
assert_eq!(settings, Settings::default());
}
#[test]
fn cli_overrides_win() {
let args = Args {
config: None,
bind: Some("0.0.0.0:7000".to_string()),
data_dir: Some("/var/lib/tephra".to_string()),
log: Some("tephra=debug".to_string()),
healthcheck: false,
};
let settings = load(&args).unwrap();
assert_eq!(settings.bind, "0.0.0.0:7000");
assert_eq!(settings.data_dir, "/var/lib/tephra");
assert_eq!(settings.log.as_deref(), Some("tephra=debug"));
}
#[test]
fn zero_writer_counts_are_rejected_not_panicked() {
let mut settings = Settings::default();
settings.writer.queue_capacity = 0;
assert!(settings.validate().is_err());
let mut settings = Settings::default();
settings.writer.max_batch_records = 0;
assert!(settings.validate().is_err());
}
#[test]
fn zero_server_durations_are_rejected() {
let mut settings = Settings::default();
settings.server.subscriptions.wait_tick_ms = 0;
assert!(settings.validate().is_err());
let mut settings = Settings::default();
settings.server.keepalive.idle_secs = 0;
assert!(settings.validate().is_err());
let mut settings = Settings::default();
settings.server.keepalive.interval_secs = 0;
assert!(settings.validate().is_err());
}
#[test]
fn zero_server_concurrency_counts_are_rejected() {
let mut settings = Settings::default();
settings.server.backpressure.max_inflight_per_conn = 0;
assert!(settings.validate().is_err());
let mut settings = Settings::default();
settings.server.subscriptions.max_concurrent = 0;
assert!(settings.validate().is_err());
let mut settings = Settings::default();
settings.server.backpressure.frame_queue_depth = 0;
assert!(settings.validate().is_err());
}
fn with_tls(mut settings: Settings) -> Settings {
settings.tls.cert = Some("server.crt".to_string());
settings.tls.key = Some("server.key".to_string());
settings
}
fn token(value: &str) -> TokenSettings {
TokenSettings {
token: value.to_string(),
}
}
#[test]
fn auth_tokens_require_tls_unless_allow_insecure() {
let mut settings = Settings::default();
settings.auth.tokens = vec![token("secret")];
assert!(settings.validate().is_err());
let mut settings = with_tls(Settings::default());
settings.auth.tokens = vec![token("secret")];
assert!(settings.validate().is_ok());
let mut settings = Settings::default();
settings.auth.tokens = vec![token("secret")];
settings.auth.allow_insecure = true;
assert!(settings.validate().is_ok());
}
#[test]
fn empty_auth_token_is_rejected() {
let mut settings = with_tls(Settings::default());
settings.auth.tokens = vec![token("")];
assert!(settings.validate().is_err());
}
#[test]
fn no_auth_tokens_is_open_and_valid() {
let settings = Settings::default();
assert!(settings.validate().is_ok());
assert!(settings.auth_tokens().is_none());
}
#[test]
fn auth_tokens_collects_configured_tokens() {
let mut settings = with_tls(Settings::default());
settings.auth.tokens = vec![token("alpha"), token("beta")];
assert_eq!(
settings.auth_tokens(),
Some(vec!["alpha".to_string(), "beta".to_string()])
);
}
#[test]
fn tls_cert_and_key_must_be_set_together() {
let mut settings = Settings::default();
settings.tls.cert = Some("server.crt".to_string());
assert!(settings.validate().is_err());
let mut settings = Settings::default();
settings.tls.key = Some("server.key".to_string());
assert!(settings.validate().is_err());
let mut settings = Settings::default();
settings.tls.cert = Some("server.crt".to_string());
settings.tls.key = Some("server.key".to_string());
assert!(settings.validate().is_ok());
}
}