use serde::Deserialize;
use std::net::SocketAddr;
use std::path::{Path, PathBuf};
use crate::error::{Error, Result};
use crate::lua::Builtins;
use crate::runtime::RuntimeOpts;
const DEFAULT_MEMORY_LIMIT: usize = 8 * 1024 * 1024;
const DEFAULT_EXEC_TIMEOUT_MS: u64 = 30_000;
#[derive(Debug, Clone, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct Config {
pub listen: SocketAddr,
pub handler_script: PathBuf,
pub config_script: Option<PathBuf>,
pub templates_dir: Option<PathBuf>,
pub database: Option<PathBuf>,
pub workers: usize,
pub dev_mode: bool,
pub builtins: Option<Vec<String>>,
pub lua: LuaConfig,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct LuaConfig {
pub stdlib: Vec<String>,
pub memory_limit: usize,
pub exec_timeout_ms: u64,
}
impl Default for Config {
fn default() -> Self {
Self {
listen: SocketAddr::from(([127, 0, 0, 1], 3000)),
handler_script: PathBuf::from("scripts/handler.lua"),
config_script: None,
templates_dir: None,
database: None,
workers: std::thread::available_parallelism().map_or(1, |n| n.get()),
dev_mode: false,
builtins: None,
lua: LuaConfig::default(),
}
}
}
impl Default for LuaConfig {
fn default() -> Self {
Self {
stdlib: ["math", "table", "string", "utf8", "coroutine", "package"]
.map(String::from)
.to_vec(),
memory_limit: DEFAULT_MEMORY_LIMIT,
exec_timeout_ms: DEFAULT_EXEC_TIMEOUT_MS,
}
}
}
impl Config {
pub fn from_file(path: &Path) -> Result<Self> {
let data = std::fs::read_to_string(path).map_err(|err| {
Error::Config(format!(
"failed to read the config file {}: {err}",
path.display()
))
})?;
toml::from_str(&data).map_err(|err| {
Error::Config(format!(
"failed to parse the config file {}: {err}",
path.display()
))
})
}
pub fn apply_env(&mut self) -> Result {
if let Some(v) = env_var("NITR_LISTEN") {
self.listen = parse_env("NITR_LISTEN", &v)?;
}
if let Some(v) = env_var("NITR_HANDLER_SCRIPT") {
self.handler_script = PathBuf::from(v);
}
if let Some(v) = env_var("NITR_CONFIG_SCRIPT") {
self.config_script = Some(PathBuf::from(v));
}
if let Some(v) = env_var("NITR_TEMPLATES_DIR") {
self.templates_dir = Some(PathBuf::from(v));
}
if let Some(v) = env_var("NITR_DATABASE") {
self.database = Some(PathBuf::from(v));
}
if let Some(v) = env_var("NITR_WORKERS") {
self.workers = parse_env("NITR_WORKERS", &v)?;
}
if let Some(v) = env_var("NITR_DEV_MODE") {
self.dev_mode = parse_env("NITR_DEV_MODE", &v)?;
}
if let Some(v) = env_var("NITR_LUA_MEMORY_LIMIT") {
self.lua.memory_limit = parse_env("NITR_LUA_MEMORY_LIMIT", &v)?;
}
if let Some(v) = env_var("NITR_LUA_EXEC_TIMEOUT_MS") {
self.lua.exec_timeout_ms = parse_env("NITR_LUA_EXEC_TIMEOUT_MS", &v)?;
}
Ok(())
}
pub fn builtins(&self) -> Result<Builtins> {
let Some(names) = &self.builtins else {
return Ok(Builtins::all());
};
let mut builtins = Builtins::empty();
for name in names {
let builtin = Builtins::from_config_name(name)
.ok_or_else(|| Error::Config(format!("unknown builtin `{name}`")))?;
if builtin == Builtins::TEMPLATE && self.templates_dir.is_none() {
return Err(Error::Config(
"builtin `template` is enabled but `templates_dir` is not set".into(),
));
}
if builtin == Builtins::DATABASE && self.database.is_none() {
return Err(Error::Config(
"builtin `db` is enabled but `database` is not set".into(),
));
}
builtins |= builtin;
}
Ok(builtins)
}
pub fn runtime_opts(&self) -> Result<RuntimeOpts> {
let package_dir = self
.handler_script
.parent()
.filter(|p| !p.as_os_str().is_empty())
.unwrap_or_else(|| Path::new("."))
.to_path_buf();
Ok(RuntimeOpts {
libs: self.lua.parse_stdlib()?,
memory_limit: self.lua.memory_limit,
dev_mode: self.dev_mode,
exec_timeout: match self.lua.exec_timeout_ms {
0 => None,
ms => Some(std::time::Duration::from_millis(ms)),
},
package_dir: Some(package_dir),
})
}
}
impl LuaConfig {
pub fn parse_stdlib(&self) -> Result<mlua::StdLib> {
use mlua::StdLib;
let mut libs = StdLib::NONE;
for name in &self.stdlib {
libs |= match name.as_str() {
"coroutine" => StdLib::COROUTINE,
"table" => StdLib::TABLE,
"io" => StdLib::IO,
"os" => StdLib::OS,
"string" => StdLib::STRING,
"utf8" => StdLib::UTF8,
"math" => StdLib::MATH,
"package" => StdLib::PACKAGE,
"debug" => StdLib::DEBUG,
_ => {
return Err(Error::Config(format!(
"unknown Lua standard library `{name}`"
)))
}
};
}
Ok(libs)
}
}
fn env_var(name: &str) -> Option<String> {
std::env::var(name).ok().filter(|v| !v.is_empty())
}
fn parse_env<T: std::str::FromStr>(name: &str, value: &str) -> Result<T>
where
T::Err: std::fmt::Display,
{
value
.parse()
.map_err(|err| Error::Config(format!("invalid value for {name}: {err}")))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::lua::Builtins;
fn write_temp_config(name: &str, content: &str) -> PathBuf {
let path = std::env::temp_dir().join(format!("nitr-test-{}-{name}", std::process::id()));
std::fs::write(&path, content).expect("write temp config");
path
}
#[test]
fn defaults_are_sane() {
let cfg = Config::default();
assert_eq!(cfg.listen, SocketAddr::from(([127, 0, 0, 1], 3000)));
assert_eq!(cfg.handler_script, PathBuf::from("scripts/handler.lua"));
assert!(cfg.workers >= 1);
assert!(!cfg.dev_mode);
assert_eq!(cfg.builtins().expect("builtins"), Builtins::all());
assert!(!cfg.lua.stdlib.iter().any(|s| s == "io" || s == "os"));
}
#[test]
fn parses_a_full_config_file() {
let path = write_temp_config(
"full.toml",
r#"
listen = "127.0.0.1:8080"
handler_script = "app/handler.lua"
database = "app.db"
workers = 2
dev_mode = true
builtins = ["dbg", "json", "db"]
[lua]
stdlib = ["math", "string", "package"]
memory_limit = 1048576
exec_timeout_ms = 500
"#,
);
let cfg = Config::from_file(&path).expect("parse config");
std::fs::remove_file(&path).ok();
assert_eq!(cfg.listen, SocketAddr::from(([127, 0, 0, 1], 8080)));
assert!(cfg.dev_mode);
assert_eq!(
cfg.builtins().expect("builtins"),
Builtins::DEBUG | Builtins::JSON | Builtins::DATABASE
);
let opts = cfg.runtime_opts().expect("runtime opts");
assert_eq!(opts.memory_limit, 1048576);
assert_eq!(
opts.exec_timeout,
Some(std::time::Duration::from_millis(500))
);
assert!(opts.dev_mode);
assert_eq!(opts.package_dir.as_deref(), Some(Path::new("app")));
}
#[test]
fn rejects_unknown_fields() {
let path = write_temp_config("typo.toml", "memroy_limit = 1\n");
let err = Config::from_file(&path).expect_err("typo must fail");
std::fs::remove_file(&path).ok();
assert!(err.to_string().contains("memroy_limit"));
}
#[test]
fn strict_builtins_require_their_settings() {
let mut cfg = Config {
builtins: Some(vec!["db".into()]),
..Config::default()
};
assert!(cfg.builtins().is_err());
cfg.database = Some(PathBuf::from("x.db"));
assert_eq!(cfg.builtins().expect("builtins"), Builtins::DATABASE);
cfg.builtins = Some(vec!["nope".into()]);
assert!(cfg.builtins().is_err());
}
#[test]
fn exec_timeout_zero_disables_the_budget() {
let mut cfg = Config::default();
cfg.lua.exec_timeout_ms = 0;
assert_eq!(cfg.runtime_opts().expect("opts").exec_timeout, None);
}
#[test]
fn unknown_stdlib_name_fails() {
let mut cfg = Config::default();
cfg.lua.stdlib.push("ffi".into());
assert!(cfg.runtime_opts().is_err());
}
}