use std::fs;
use std::path::{Component, Path, PathBuf};
use serde::Deserialize;
use thiserror::Error;
use url::{Host, Url};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LocalModelEndpoint {
base_url: Url,
model: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
pub struct ToolPaths {
pub pdftotext: PathBuf,
pub pdftoppm: PathBuf,
pub tesseract: PathBuf,
pub pandoc: PathBuf,
pub zk: PathBuf,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Settings {
pub vault_root: PathBuf,
pub inbox_directory: PathBuf,
pub research_directory: PathBuf,
pub archive_directory: PathBuf,
pub database_path: PathBuf,
pub tools: ToolPaths,
pub classifier: LocalModelEndpoint,
pub distiller: LocalModelEndpoint,
pub minimum_text_characters: usize,
pub process_timeout_seconds: u64,
pub maximum_output_bytes: usize,
pub maximum_source_bytes: u64,
pub watch_interval_seconds: u64,
}
#[derive(Debug, Deserialize)]
struct RawSettings {
vault_root: PathBuf,
inbox_directory: PathBuf,
research_directory: PathBuf,
archive_directory: PathBuf,
database_path: PathBuf,
tools: ToolPaths,
classifier: RawEndpoint,
distiller: RawEndpoint,
minimum_text_characters: usize,
process_timeout_seconds: u64,
maximum_output_bytes: usize,
maximum_source_bytes: u64,
watch_interval_seconds: u64,
}
#[derive(Debug, Deserialize)]
struct RawEndpoint {
base_url: String,
model: String,
}
impl Settings {
pub fn load(path: &Path) -> Result<Self, SettingsError> {
let text = fs::read_to_string(path).map_err(SettingsError::Read)?;
let raw: RawSettings = toml::from_str(&text)
.map_err(|error| SettingsError::InvalidConfiguration(error.to_string()))?;
for relative in [
&raw.inbox_directory,
&raw.research_directory,
&raw.archive_directory,
] {
validate_relative_path(relative)?;
}
for tool in [
&raw.tools.pdftotext,
&raw.tools.pdftoppm,
&raw.tools.tesseract,
&raw.tools.pandoc,
&raw.tools.zk,
] {
if !tool.is_absolute() {
return Err(SettingsError::UnsafeToolPath(tool.display().to_string()));
}
}
if raw.minimum_text_characters == 0
|| raw.process_timeout_seconds == 0
|| raw.maximum_output_bytes == 0
|| raw.maximum_source_bytes == 0
|| raw.watch_interval_seconds == 0
{
return Err(SettingsError::InvalidLimit);
}
Ok(Self {
vault_root: raw.vault_root,
inbox_directory: raw.inbox_directory,
research_directory: raw.research_directory,
archive_directory: raw.archive_directory,
database_path: raw.database_path,
tools: raw.tools,
classifier: LocalModelEndpoint::new(&raw.classifier.base_url, raw.classifier.model)?,
distiller: LocalModelEndpoint::new(&raw.distiller.base_url, raw.distiller.model)?,
minimum_text_characters: raw.minimum_text_characters,
process_timeout_seconds: raw.process_timeout_seconds,
maximum_output_bytes: raw.maximum_output_bytes,
maximum_source_bytes: raw.maximum_source_bytes,
watch_interval_seconds: raw.watch_interval_seconds,
})
}
#[must_use]
pub fn required_tools(&self) -> Vec<crate::doctor::RequiredTool> {
[
("pdftotext", &self.tools.pdftotext),
("pdftoppm", &self.tools.pdftoppm),
("tesseract", &self.tools.tesseract),
("pandoc", &self.tools.pandoc),
("zk", &self.tools.zk),
]
.into_iter()
.map(|(name, path)| crate::doctor::RequiredTool {
name: name.to_owned(),
path: path.clone(),
})
.collect()
}
}
fn validate_relative_path(path: &Path) -> Result<(), SettingsError> {
let valid = !path.as_os_str().is_empty()
&& path.is_relative()
&& path
.components()
.all(|component| matches!(component, Component::Normal(_)));
if !valid {
return Err(SettingsError::UnsafeRelativePath(
path.display().to_string(),
));
}
Ok(())
}
impl LocalModelEndpoint {
pub fn new(base_url: &str, model: impl Into<String>) -> Result<Self, SettingsError> {
let base_url = Url::parse(base_url)
.map_err(|error| SettingsError::InvalidModelEndpoint(error.to_string()))?;
if !matches!(base_url.scheme(), "http" | "https") {
return Err(SettingsError::UnsupportedModelEndpointScheme(
base_url.scheme().to_owned(),
));
}
if !base_url.username().is_empty() || base_url.password().is_some() {
return Err(SettingsError::ModelEndpointContainsCredentials);
}
let is_loopback = match base_url.host() {
Some(Host::Ipv4(address)) => address.is_loopback(),
Some(Host::Ipv6(address)) => address.is_loopback(),
Some(Host::Domain(_)) | None => false,
};
if !is_loopback {
return Err(SettingsError::ModelEndpointNotLoopback(
base_url.to_string(),
));
}
if base_url.query().is_some() || base_url.fragment().is_some() {
return Err(SettingsError::ModelEndpointContainsQueryOrFragment);
}
let model = model.into();
if model.trim().is_empty() || model.chars().any(char::is_control) {
return Err(SettingsError::EmptyModelName);
}
Ok(Self { base_url, model })
}
#[must_use]
pub fn base_url(&self) -> &Url {
&self.base_url
}
#[must_use]
pub fn model(&self) -> &str {
&self.model
}
}
#[derive(Debug, Error)]
pub enum SettingsError {
#[error("failed to read configuration: {0}")]
Read(#[source] std::io::Error),
#[error("invalid configuration: {0}")]
InvalidConfiguration(String),
#[error("unsafe vault-relative path: {0}")]
UnsafeRelativePath(String),
#[error("configured limits and intervals must be greater than zero")]
InvalidLimit,
#[error("tool path must be absolute: {0}")]
UnsafeToolPath(String),
#[error("invalid model endpoint: {0}")]
InvalidModelEndpoint(String),
#[error("unsupported model endpoint scheme: {0}")]
UnsupportedModelEndpointScheme(String),
#[error("model endpoint must use an explicit loopback IP address: {0}")]
ModelEndpointNotLoopback(String),
#[error("model endpoint URL must not contain credentials")]
ModelEndpointContainsCredentials,
#[error("model endpoint URL must not contain a query or fragment")]
ModelEndpointContainsQueryOrFragment,
#[error("model name must not be empty")]
EmptyModelName,
#[error("live integration tests require a temporary vault: {0}")]
UnsafeLiveTestVault(String),
}
pub fn validate_live_test_vault(
vault_root: &Path,
temporary_root: &Path,
) -> Result<(), SettingsError> {
let vault_root = vault_root.canonicalize().map_err(SettingsError::Read)?;
let temporary_root = temporary_root.canonicalize().map_err(SettingsError::Read)?;
if !vault_root.starts_with(temporary_root) {
return Err(SettingsError::UnsafeLiveTestVault(
vault_root.display().to_string(),
));
}
Ok(())
}