add support for toml file parsing, tweak some configuration names

This commit is contained in:
Chris Beck
2025-12-06 22:53:08 -07:00
parent df1616daeb
commit c2ad134393
4 changed files with 41 additions and 21 deletions
+1 -1
View File
@@ -29,10 +29,10 @@ reqwest = { workspace = true }
serde = { workspace = true }
serde_json = { workspace = true }
syslog_rfc5424 = { workspace = true }
toml = { workspace = true }
tokio = { workspace = true, features = ["rt-multi-thread", "signal"] }
tokio-util = { workspace = true }
tracing = { workspace = true }
tracing-subscriber = { workspace = true }
[dev-dependencies]
toml = { workspace = true }
+27 -7
View File
@@ -6,7 +6,7 @@ use conf::{Conf, Subcommands};
use hyper::service::service_fn;
use hyper_util::{rt::TokioIo, server::conn::auto};
use signal_gateway::{Gateway, GatewayConfig};
use std::{net::SocketAddr, sync::Arc, time::Duration};
use std::{env, fs, net::SocketAddr, path::PathBuf, sync::Arc, time::Duration};
use tokio::net::TcpListener;
use tokio_util::sync::CancellationToken;
use tracing::{error, info, warn};
@@ -33,6 +33,11 @@ pub enum AdminHandlerCommand {
#[derive(Conf, Debug)]
#[conf(serde, test)]
pub struct Config {
/// Path to a TOML config file (optional).
/// This is parsed before other args, so config file values can be overridden by CLI args.
#[allow(dead_code)] // Parsed early via find_parameter, kept here for --help
#[conf(long)]
config_file: Option<PathBuf>,
/// If true, just validate config and don't start
#[conf(long)]
dry_run: bool,
@@ -79,15 +84,28 @@ fn init_logging() {
}
#[tokio::main]
async fn main() {
async fn main() -> Result<(), Box<dyn std::error::Error>> {
init_logging();
let config = Config::parse();
// Check for --config-file before the main parse, so we can load it and pass to conf
let config_file_path = conf::find_parameter("config-file", env::args_os());
let config = if let Some(config_path) = config_file_path {
let path_display = config_path.to_string_lossy();
let file_contents = fs::read_to_string(&config_path)
.map_err(|err| format!("Could not open config file '{path_display}': {err}"))?;
let doc: toml::Value = toml::from_str(&file_contents)
.map_err(|err| format!("Config file '{path_display}' is not valid TOML: {err}"))?;
info!("Loaded config file: {path_display}");
Config::conf_builder().doc(path_display, doc).parse()
} else {
Config::parse()
};
info!("Config = {config:#?}");
if config.dry_run {
return;
return Ok(());
}
let token = CancellationToken::new();
@@ -123,6 +141,8 @@ async fn main() {
// Run gateway task and block on it returning. Note that it exits if the token is canceled.
gateway.run().await;
Ok(())
}
fn start_http_task(listener: TcpListener, gateway: Arc<Gateway>) -> tokio::task::JoinHandle<()> {
@@ -175,7 +195,7 @@ signal_account = "+15551234567"
signal_cli_tcp_addr = "127.0.0.1:7583"
signal_cli_retry_delay = "10s"
[admin_signal_uuids]
[signal_admins]
"abc-123-uuid" = ["12345 67890 12345 67890 12345 67890"]
"def-456-uuid" = []
@@ -225,11 +245,11 @@ limits = [
config.gateway.signal_cli_retry_delay,
Duration::from_secs(10)
);
assert_eq!(config.gateway.admin_signal_uuids.len(), 2);
assert_eq!(config.gateway.signal_admins.len(), 2);
assert!(
config
.gateway
.admin_signal_uuids
.signal_admins
.get("abc-123-uuid")
.is_some()
);
+5 -5
View File
@@ -62,10 +62,10 @@ pub struct GatewayConfig {
/// Delay before retrying connection to signal-cli after an error.
#[conf(long, env, default_value = "5s", value_parser = conf_extra::parse_duration, serde(use_value_parser))]
pub signal_cli_retry_delay: Duration,
/// Admin UUIDs mapped to their safety numbers (can be empty).
/// Signal admin UUIDs mapped to their safety numbers (can be empty).
/// Accepts either a map `{"uuid1": ["12345..."], "uuid2": []}` or a list `["uuid1", "uuid2"]`.
#[conf(long, env, value_parser = serde_json::from_str)]
pub admin_signal_uuids: SignalTrustSet,
pub signal_admins: SignalTrustSet,
/// If set, alerts are sent to this group instead of individual admins.
#[conf(long, env)]
pub alert_group_id: Option<String>,
@@ -306,7 +306,7 @@ impl Gateway {
loop {
match self
.config
.admin_signal_uuids
.signal_admins
.update_trust(signal_cli, &self.config.signal_account)
.await
{
@@ -347,7 +347,7 @@ impl Gateway {
if let Some(group_id) = &self.config.alert_group_id {
MessageTarget::Group(group_id.clone())
} else {
MessageTarget::Recipients(self.config.admin_signal_uuids.uuids().cloned().collect())
MessageTarget::Recipients(self.config.signal_admins.uuids().cloned().collect())
}
}
};
@@ -384,7 +384,7 @@ impl Gateway {
let from_group = data_message.group_info.as_ref().map(|g| g.group_id.clone());
// Check if sender is an admin
if !self.config.admin_signal_uuids.is_trusted(&msg.envelope) {
if !self.config.signal_admins.is_trusted(&msg.envelope) {
warn!("Ignoring message from non-admin: {msg:?}");
continue;
}
+8 -8
View File
@@ -24,9 +24,9 @@ use std::{path::PathBuf, str::FromStr, time::Duration};
#[derive(Clone, Conf, Debug)]
#[conf(serde)]
pub struct PrometheusConfig {
/// Address of prometheus host. Should start with http and usually indicate port 9090
/// URL of the Prometheus server query API (e.g., `http://localhost:9090`).
#[conf(long, env)]
pub prometheus_host: String,
pub prometheus_url: String,
#[cfg(feature = "plot")]
/// Configuration options for generated plots
#[conf(flatten, prefix)]
@@ -58,7 +58,7 @@ impl Prometheus {
) -> Result<(ExtractLabels, Vec<Option<(f64, MetricVal)>>), BoxError> {
info!("Prom query: {query}");
let vector: Vec<MetricValue> = QueryRequest { query, time: None }
.send_with_client(&self.reqwest_client, &self.config.prometheus_host)
.send_with_client(&self.reqwest_client, &self.config.prometheus_url)
.await?
.into_vector()?;
@@ -76,7 +76,7 @@ impl Prometheus {
Ok(SeriesRequest {
matches: matches.into_iter().map(|s| s.as_ref().to_owned()).collect(),
}
.send_with_client(&self.reqwest_client, &self.config.prometheus_host)
.send_with_client(&self.reqwest_client, &self.config.prometheus_url)
.await?)
}
@@ -88,14 +88,14 @@ impl Prometheus {
Ok(LabelsRequest {
matches: matches.into_iter().map(|s| s.as_ref().to_owned()).collect(),
}
.send_with_client::<()>(&self.reqwest_client, &self.config.prometheus_host)
.send_with_client::<()>(&self.reqwest_client, &self.config.prometheus_url)
.await?)
}
/// Get the list of current alerts
pub async fn alerts(&self) -> Result<Vec<AlertInfo>, BoxError> {
Ok(AlertsRequest {}
.send_with_client(&self.reqwest_client, &self.config.prometheus_host)
.send_with_client(&self.reqwest_client, &self.config.prometheus_url)
.await?
.alerts)
}
@@ -151,7 +151,7 @@ impl Prometheus {
let matrix = QueryRangeRequest::builder(query.clone())
.range(now - range..now)
.build()
.send_with_client(&self.reqwest_client, &self.config.prometheus_host)
.send_with_client(&self.reqwest_client, &self.config.prometheus_url)
.await?
.into_matrix()?;
@@ -170,7 +170,7 @@ impl Prometheus {
let matrix = QueryRangeRequest::builder(query.to_owned())
.since(since)
.build()
.send_with_client(&self.reqwest_client, &self.config.prometheus_host)
.send_with_client(&self.reqwest_client, &self.config.prometheus_url)
.await?
.into_matrix()?;