use std::collections::HashSet;
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)]
pub struct ClassifierProfile {
pub name: String,
pub endpoint: LocalModelEndpoint,
pub maximum_input_characters: usize,
}
#[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: ClassifierProfile,
pub classifier_fallbacks: Vec<ClassifierProfile>,
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,
#[serde(default)]
classifier_fallbacks: Vec<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 {
#[serde(default)]
name: Option<String>,
base_url: String,
model: String,
#[serde(default)]
maximum_input_characters: Option<usize>,
}
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);
}
let classifier = classifier_profile(raw.classifier, "primary", false)?;
let classifier_fallbacks = raw
.classifier_fallbacks
.into_iter()
.enumerate()
.map(|(index, profile)| {
classifier_profile(profile, &format!("fallback-{}", index + 1), true)
})
.collect::<Result<Vec<_>, _>>()?;
let mut names = HashSet::new();
if !names.insert(classifier.name.clone())
|| classifier_fallbacks
.iter()
.any(|profile| !names.insert(profile.name.clone()))
{
return Err(SettingsError::DuplicateClassifierProfile);
}
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,
classifier_fallbacks,
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 classifier_profile(
raw: RawEndpoint,
default_name: &str,
require_limit: bool,
) -> Result<ClassifierProfile, SettingsError> {
let name = match raw.name {
Some(name) => name,
None if require_limit => return Err(SettingsError::MissingClassifierProfileName),
None => default_name.to_owned(),
};
if name.trim().is_empty() || name.chars().any(char::is_control) {
return Err(SettingsError::InvalidClassifierProfileName);
}
let maximum_input_characters = match raw.maximum_input_characters {
Some(0) => return Err(SettingsError::InvalidLimit),
Some(limit) => limit,
None if require_limit => return Err(SettingsError::MissingClassifierProfileLimit),
None => usize::MAX,
};
Ok(ClassifierProfile {
name,
endpoint: LocalModelEndpoint::new(&raw.base_url, raw.model)?,
maximum_input_characters,
})
}
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("fallback classifier profiles require maximum_input_characters")]
MissingClassifierProfileLimit,
#[error("fallback classifier profiles require an explicit name")]
MissingClassifierProfileName,
#[error("classifier profile name must be non-empty and contain no control characters")]
InvalidClassifierProfileName,
#[error("classifier profile names must be unique")]
DuplicateClassifierProfile,
#[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(())
}