extern crate std;
use std::format;
use std::prelude::v1::*;
use std::fs;
use std::io;
use std::path::Path;
use super::codegen;
use super::curve::DefinitionsFile;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ValueType {
U8,
U16,
}
impl ValueType {
pub fn parse(name: &str) -> Result<Self, Error> {
match name {
"u8" => Ok(Self::U8),
"u16" => Ok(Self::U16),
other => Err(Error::Validation(format!(
"unsupported value type `{other}` (expected `u8` or `u16`)"
))),
}
}
pub const fn as_str(self) -> &'static str {
match self {
Self::U8 => "u8",
Self::U16 => "u16",
}
}
pub const fn required_lut_size(self) -> usize {
match self {
Self::U8 => 256,
Self::U16 => 65_536,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GenerateOptions {
pub value_type: ValueType,
pub lut_size: usize,
}
impl Default for GenerateOptions {
fn default() -> Self {
Self {
value_type: ValueType::U8,
lut_size: ValueType::U8.required_lut_size(),
}
}
}
#[derive(Debug)]
pub enum Error {
Io(io::Error),
Toml(toml::de::Error),
Validation(String),
}
impl core::fmt::Display for Error {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Io(error) => write!(f, "{error}"),
Self::Toml(error) => write!(f, "invalid TOML: {error}"),
Self::Validation(message) => write!(f, "{message}"),
}
}
}
impl std::error::Error for Error {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Io(error) => Some(error),
Self::Toml(error) => Some(error),
Self::Validation(_) => None,
}
}
}
impl From<io::Error> for Error {
fn from(error: io::Error) -> Self {
Self::Io(error)
}
}
impl From<toml::de::Error> for Error {
fn from(error: toml::de::Error) -> Self {
Self::Toml(error)
}
}
pub fn generate_from_toml(path: impl AsRef<Path>, opts: &GenerateOptions) -> Result<String, Error> {
let toml_str = fs::read_to_string(path)?;
generate_from_str(&toml_str, opts)
}
pub fn generate_from_str(toml: &str, opts: &GenerateOptions) -> Result<String, Error> {
let defs: DefinitionsFile = toml::from_str(toml)?;
generate(&defs, opts)
}
pub fn generate_to_path(
input: impl AsRef<Path>,
output: impl AsRef<Path>,
opts: &GenerateOptions,
) -> Result<(), Error> {
let source = generate_from_toml(input, opts)?;
fs::write(output, source)?;
Ok(())
}
pub fn generate(defs: &DefinitionsFile, opts: &GenerateOptions) -> Result<String, Error> {
validate_options(opts)?;
codegen::generate(defs, opts.value_type.as_str(), opts.lut_size).map_err(Error::Validation)
}
fn validate_options(opts: &GenerateOptions) -> Result<(), Error> {
let required = opts.value_type.required_lut_size();
if opts.lut_size != required {
return Err(Error::Validation(format!(
"--lut-size must be {required} for --value-type {} \
(the LUT must cover the full {} domain)",
opts.value_type.as_str(),
opts.value_type.as_str(),
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn lut_size_must_cover_full_u8_domain() {
let ok = GenerateOptions {
value_type: ValueType::U8,
lut_size: 256,
};
assert!(validate_options(&ok).is_ok());
let bad = GenerateOptions {
value_type: ValueType::U8,
lut_size: 255,
};
assert_eq!(
validate_options(&bad).unwrap_err().to_string(),
"--lut-size must be 256 for --value-type u8 \
(the LUT must cover the full u8 domain)"
);
}
#[test]
fn lut_size_must_cover_full_u16_domain() {
let ok = GenerateOptions {
value_type: ValueType::U16,
lut_size: 65_536,
};
assert!(validate_options(&ok).is_ok());
let bad = GenerateOptions {
value_type: ValueType::U16,
lut_size: 256,
};
assert_eq!(
validate_options(&bad).unwrap_err().to_string(),
"--lut-size must be 65536 for --value-type u16 \
(the LUT must cover the full u16 domain)"
);
}
#[test]
fn value_type_parse_rejects_unknown() {
let error = ValueType::parse("u32").unwrap_err().to_string();
assert!(error.contains("unsupported value type"));
}
#[test]
fn generate_from_str_linear_curve() {
let toml = r#"
[curves.linear]
builtin = "linear"
"#;
let out = generate_from_str(toml, &GenerateOptions::default()).unwrap();
assert!(out.contains("LINEAR_FWD"));
assert!(out.contains("LINEAR_INV"));
assert!(out.contains("// Auto-generated by ph-curves-gen."));
}
#[test]
fn semantic_validation_returns_error_instead_of_panicking() {
let error = generate_from_str("[curves.bad]\n", &GenerateOptions::default()).unwrap_err();
assert!(matches!(error, Error::Validation(_)));
assert!(error.to_string().contains("exactly one"));
}
#[test]
fn invalid_formula_returns_validation_error() {
let error = generate_from_str(
"[curves.bad]\nformula = \"t @ 2\"\n",
&GenerateOptions::default(),
)
.unwrap_err();
assert!(matches!(error, Error::Validation(_)));
assert!(error.to_string().contains("unexpected character"));
}
}