diff --git a/bytescale/Cargo.toml b/bytescale/Cargo.toml index ce0cb07..6a7a50e 100644 --- a/bytescale/Cargo.toml +++ b/bytescale/Cargo.toml @@ -15,6 +15,7 @@ default = ["std", "derive"] std = ["humanbyte/std"] derive = [] arbitrary = ["dep:arbitrary", "std"] +schemars = ["humanbyte/schemars"] serde = ["humanbyte/serde"] [dependencies] @@ -22,6 +23,7 @@ arbitrary = { version = "1", features = ["derive"], optional = true } humanbyte = { version = "0.2.1-alpha.0", path = "../humanbyte", features = ["derive"] } [dev-dependencies] +schemars = "1" serde = { version = "1.0", features = ["derive"] } serde_json = { version = "1.0", features = ["std"] } toml = "0.8" diff --git a/bytescale/src/lib.rs b/bytescale/src/lib.rs index d631148..15d5b40 100644 --- a/bytescale/src/lib.rs +++ b/bytescale/src/lib.rs @@ -145,10 +145,14 @@ mod tests { #[test] fn when_err() { // shortcut for writing test cases - fn parse(s: &str) -> Result { + fn parse(s: &str) -> Result { s.parse::() } + // the error type chains with `?` in std contexts + fn assert_error(_: &E) {} + assert_error(&parse("oops").unwrap_err()); + assert!(parse("").is_err()); assert!(parse("a124GB").is_err()); assert!(parse("1.3 42.0 B").is_err()); @@ -158,6 +162,37 @@ mod tests { assert!(parse("1 000 B").is_err()); } + #[test] + fn test_div() { + let file_size = ByteScale::gib(4); + let chunk_size = ByteScale::mib(64); + assert_eq!(file_size / chunk_size, 64u64); + assert_eq!(file_size / 4u64, ByteScale::gib(1)); + assert_eq!(file_size % chunk_size, ByteScale::b(0)); + + let mut x = ByteScale::mb(10); + x /= 2u64; + assert_eq!(x, ByteScale::mb(5)); + } + + #[test] + fn test_to_string_with_precision() { + use humanbyte::Format; + let x = ByteScale::tib(1); + assert_eq!(x.to_string_with_precision(Format::IEC, 2), "1.00 TiB"); + assert_eq!(x.to_string_with_precision(Format::IEC, 0), "1 TiB"); + assert_eq!( + ByteScale::kib(1) + 512u64, + "1.500 KiB".parse::().unwrap() + ); + } + + #[test] + fn test_standalone_parse() { + // no newtype required + assert_eq!(humanbyte::parse("1.5 KiB"), Ok(1536)); + } + #[test] #[should_panic(expected = "byte size overflows u64")] fn test_constructor_overflow() { @@ -200,4 +235,62 @@ mod tests { let s: S = toml::from_str(r#"x = "9223372036854775807""#).unwrap(); assert_eq!(s.x, "9223372036854775807".parse::().unwrap()); } + + #[cfg(feature = "serde")] + #[test] + fn test_serde_with_plain_integers() { + use std::collections::BTreeMap; + + use serde::{Deserialize, Serialize}; + + #[derive(Serialize, Deserialize, PartialEq, Debug)] + struct Config { + #[serde(with = "humanbyte::serde")] + buffer_size: usize, + #[serde(with = "humanbyte::serde")] + max_size: u64, + #[serde(with = "humanbyte::serde::map_keys")] + pools: BTreeMap, + } + + let config: Config = serde_json::from_str( + r#"{ + "buffer_size": "1.5 KiB", + "max_size": 1048576, + "pools": { "4 KiB": "small", "2 MiB": "large" } + }"#, + ) + .unwrap(); + assert_eq!(config.buffer_size, 1536); + assert_eq!(config.max_size, 1_048_576); + assert_eq!(config.pools[&4096], "small"); + assert_eq!(config.pools[&2_097_152], "large"); + + // roundtrip: serializes human-readable, parses back to the same values + let json = serde_json::to_string(&config).unwrap(); + assert!(json.contains(r#""buffer_size":"1.5 KiB""#)); + assert!(json.contains(r#""4.0 KiB":"small""#)); + let back: Config = serde_json::from_str(&json).unwrap(); + assert_eq!(back, config); + + // toml too + let config: Config = toml::from_str( + "buffer_size = \"2 KiB\"\nmax_size = \"1 MiB\"\n[pools]\n\"64 KiB\" = \"medium\"", + ) + .unwrap(); + assert_eq!(config.buffer_size, 2048); + assert_eq!(config.pools[&65536], "medium"); + } + + #[cfg(feature = "schemars")] + #[test] + fn test_json_schema() { + let schema = schemars::schema_for!(ByteScale); + let json = serde_json::to_value(&schema).unwrap(); + assert_eq!( + json["type"], + serde_json::json!(["string", "integer"]), + "schema should accept both forms: {json}" + ); + } } diff --git a/humanbyte-derive/Cargo.toml b/humanbyte-derive/Cargo.toml index d769c53..ec7c49c 100644 --- a/humanbyte-derive/Cargo.toml +++ b/humanbyte-derive/Cargo.toml @@ -12,6 +12,7 @@ proc-macro = true [features] default = [] +schemars = [] serde = [] [dependencies] diff --git a/humanbyte-derive/src/lib.rs b/humanbyte-derive/src/lib.rs index 54b1263..55bce6b 100644 --- a/humanbyte-derive/src/lib.rs +++ b/humanbyte-derive/src/lib.rs @@ -16,6 +16,9 @@ pub fn humanbyte(input: TokenStream) -> TokenStream { if cfg!(feature = "serde") { combined.extend(serde_tokens(name)); } + if cfg!(feature = "schemars") { + combined.extend(schemars_tokens(name)); + } TokenStream::from(combined) } @@ -185,6 +188,47 @@ fn ops_tokens(name: &syn::Ident) -> proc_macro2::TokenStream { } } + impl core::ops::Div<#name> for #name { + /// Dividing two byte sizes yields a dimensionless count, + /// e.g. `file_size / chunk_size` chunks. + type Output = u64; + + #[inline(always)] + fn div(self, rhs: #name) -> u64 { + self.0 / rhs.0 + } + } + + impl core::ops::Div for #name + where + T: Into, + { + type Output = #name; + #[inline(always)] + fn div(self, rhs: T) -> #name { + #name(self.0 / rhs.into()) + } + } + + impl core::ops::DivAssign for #name + where + T: Into, + { + #[inline(always)] + fn div_assign(&mut self, rhs: T) { + self.0 /= rhs.into(); + } + } + + impl core::ops::Rem<#name> for #name { + type Output = #name; + + #[inline(always)] + fn rem(self, rhs: #name) -> #name { + #name(self.0 % rhs.0) + } + } + impl core::ops::Add<#name> for u64 { type Output = #name; #[inline(always)] @@ -327,29 +371,10 @@ pub fn humanbyte_fromstr(input: TokenStream) -> TokenStream { fn fromstr_tokens(name: &syn::Ident) -> proc_macro2::TokenStream { quote! { impl core::str::FromStr for #name { - type Err = ::humanbyte::String; + type Err = ::humanbyte::ParseError; fn from_str(value: &str) -> core::result::Result { - if let Ok(v) = value.parse::() { - return Ok(Self(v)); - } - let number = ::humanbyte::take_while(value, |c| c.is_ascii_digit() || c == '.'); - match number.parse::() { - Ok(v) => { - let suffix = ::humanbyte::skip_while(&value[number.len()..], char::is_whitespace); - match suffix.parse::<::humanbyte::Unit>() { - Ok(u) => Ok(Self((v * u64::from(u) as f64) as u64)), - Err(error) => Err(::humanbyte::format!( - "couldn't parse {:?} into a known SI unit, {}", - suffix, error - )), - } - } - Err(error) => Err(::humanbyte::format!( - "couldn't parse {:?} into a ByteSize, {}", - value, error - )), - } + ::humanbyte::parse(value).map(Self) } } } @@ -370,6 +395,16 @@ fn parse_tokens(name: &syn::Ident) -> proc_macro2::TokenStream { ::humanbyte::to_string(self.0, format) } + /// Returns the size as a string with the given number of decimals. + #[inline(always)] + pub fn to_string_with_precision( + &self, + format: ::humanbyte::Format, + precision: usize, + ) -> ::humanbyte::String { + ::humanbyte::to_string_with_precision(self.0, format, precision) + } + /// Returns the inner u64 value. #[inline(always)] pub const fn as_u64(&self) -> u64 { @@ -394,41 +429,41 @@ pub fn humanbyte_serde(input: TokenStream) -> TokenStream { fn serde_tokens(name: &syn::Ident) -> proc_macro2::TokenStream { quote! { - impl<'de> ::humanbyte::serde::Deserialize<'de> for #name { + impl<'de> ::humanbyte::serde_crate::Deserialize<'de> for #name { fn deserialize(deserializer: D) -> core::result::Result where - D: ::humanbyte::serde::Deserializer<'de>, + D: ::humanbyte::serde_crate::Deserializer<'de>, { struct ByteSizeVisitor; - impl<'de> ::humanbyte::serde::de::Visitor<'de> for ByteSizeVisitor { + impl<'de> ::humanbyte::serde_crate::de::Visitor<'de> for ByteSizeVisitor { type Value = #name; fn expecting(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { formatter.write_str("an integer or string") } - fn visit_i64(self, value: i64) -> core::result::Result { + fn visit_i64(self, value: i64) -> core::result::Result { if let Ok(val) = u64::try_from(value) { Ok(#name(val)) } else { Err(E::invalid_value( - ::humanbyte::serde::de::Unexpected::Signed(value), + ::humanbyte::serde_crate::de::Unexpected::Signed(value), &"integer overflow", )) } } - fn visit_u64(self, value: u64) -> core::result::Result { + fn visit_u64(self, value: u64) -> core::result::Result { Ok(#name(value)) } - fn visit_str(self, value: &str) -> core::result::Result { + fn visit_str(self, value: &str) -> core::result::Result { if let Ok(val) = value.parse() { Ok(val) } else { Err(E::invalid_value( - ::humanbyte::serde::de::Unexpected::Str(value), + ::humanbyte::serde_crate::de::Unexpected::Str(value), &"parsable string", )) } @@ -442,10 +477,10 @@ fn serde_tokens(name: &syn::Ident) -> proc_macro2::TokenStream { } } } - impl ::humanbyte::serde::Serialize for #name { + impl ::humanbyte::serde_crate::Serialize for #name { fn serialize(&self, serializer: S) -> core::result::Result where - S: ::humanbyte::serde::Serializer, + S: ::humanbyte::serde_crate::Serializer, { if serializer.is_human_readable() { ::serialize(self.to_string().as_str(), serializer) @@ -456,3 +491,32 @@ fn serde_tokens(name: &syn::Ident) -> proc_macro2::TokenStream { } } } + +#[proc_macro_derive(HumanByteSchema)] +pub fn humanbyte_schema(input: TokenStream) -> TokenStream { + let input = parse_macro_input!(input as DeriveInput); + TokenStream::from(schemars_tokens(&input.ident)) +} + +fn schemars_tokens(name: &syn::Ident) -> proc_macro2::TokenStream { + quote! { + impl ::humanbyte::schemars_crate::JsonSchema for #name { + fn schema_name() -> ::humanbyte::Cow<'static, str> { + ::humanbyte::Cow::Borrowed(stringify!(#name)) + } + + fn schema_id() -> ::humanbyte::Cow<'static, str> { + ::humanbyte::Cow::Borrowed(concat!(module_path!(), "::", stringify!(#name))) + } + + fn json_schema( + _generator: &mut ::humanbyte::schemars_crate::SchemaGenerator, + ) -> ::humanbyte::schemars_crate::Schema { + ::humanbyte::schemars_crate::json_schema!({ + "type": ["string", "integer"], + "description": "A byte size, as either a human-readable string (e.g. \"1.5 KiB\") or a number of bytes", + }) + } + } + } +} diff --git a/humanbyte/Cargo.toml b/humanbyte/Cargo.toml index 40f39d9..a0373c7 100644 --- a/humanbyte/Cargo.toml +++ b/humanbyte/Cargo.toml @@ -12,7 +12,9 @@ default = ["std"] std = [] derive = ["dep:humanbyte-derive"] serde = ["dep:serde", "std", "humanbyte-derive/serde"] +schemars = ["dep:schemars", "std", "humanbyte-derive/schemars"] [dependencies] humanbyte-derive = { version = "0.2.1-alpha.0", path = "../humanbyte-derive", optional = true } +schemars = { version = "1", optional = true } serde = { version = "1.0", features = ["derive"], optional = true } diff --git a/humanbyte/README.md b/humanbyte/README.md index bafa370..e373a0f 100644 --- a/humanbyte/README.md +++ b/humanbyte/README.md @@ -58,6 +58,38 @@ a la carte fashion: * HumanByteOps * HumanByteFromStr * HumanByteSerde (requires the `serde` feature) +* HumanByteSchema (requires the `schemars` feature) + +## Without a newtype + +Plain `u64`/`usize` fields can use human-readable serde directly — no newtype required: + +```rust,ignore +#[derive(Serialize, Deserialize)] +struct Config { + #[serde(with = "humanbyte::serde")] + buffer_size: usize, + #[serde(with = "humanbyte::serde::map_keys")] + pools: BTreeMap, +} +``` + +And free functions mirror the derived methods: + +```rust +assert_eq!(humanbyte::parse("1.5 KiB"), Ok(1536)); +assert_eq!(humanbyte::to_string(1536, humanbyte::Format::IEC), "1.5 KiB"); +assert_eq!( + humanbyte::to_string_with_precision(1536, humanbyte::Format::IEC, 2), + "1.50 KiB" +); +``` + +## JSON schema + +With the `schemars` feature, derived types implement `schemars::JsonSchema` (accepting a +string like `"1.5 KiB"` or a raw byte count), so config types need no manual +`#[schemars(with = "String")]` annotations. [bytescale]: https://docs.rs/bytescale/latest/bytescale [bytesize]: https://docs.rs/bytesize/latest/bytesize diff --git a/humanbyte/src/lib.rs b/humanbyte/src/lib.rs index 256a0b9..132caec 100644 --- a/humanbyte/src/lib.rs +++ b/humanbyte/src/lib.rs @@ -6,10 +6,18 @@ extern crate alloc; #[cfg(feature = "std")] extern crate std; +/// Re-export of the `serde` crate for use by derive-generated code. #[cfg(feature = "serde")] -pub use serde; +#[doc(hidden)] +pub use ::serde as serde_crate; + +/// Re-export of the `schemars` crate for use by derive-generated code. +#[cfg(feature = "schemars")] +#[doc(hidden)] +pub use ::schemars as schemars_crate; // Re-export necessary types to avoid users needing explicit extern crate declarations +pub use alloc::borrow::Cow; pub use alloc::{ format, string::{String, ToString}, @@ -51,14 +59,22 @@ const UNITS_IEC: &str = "KMGTPE"; /// /// See . const UNITS_SI: &str = "kMGTPE"; -#[derive(Debug, Clone, Default)] + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] pub enum Format { #[default] IEC, SI, } +/// Formats `bytes` as a human-readable string with one decimal, e.g. `"1.5 KiB"`. pub fn to_string(bytes: u64, format: Format) -> String { + to_string_with_precision(bytes, format, 1) +} + +/// Formats `bytes` as a human-readable string with `precision` decimals, +/// e.g. `to_string_with_precision(1 << 40, Format::IEC, 2)` is `"1.00 TiB"`. +pub fn to_string_with_precision(bytes: u64, format: Format, precision: usize) -> String { let unit = match format { Format::IEC => KIB, Format::SI => KB, @@ -83,7 +99,8 @@ pub fn to_string(bytes: u64, format: Format) -> String { exp += 1; } format!( - "{:.1} {}{}", + "{:.*} {}{}", + precision, (bytes as f64 / unit.pow(exp) as f64), unit_prefix[(exp - 1) as usize] as char, unit_suffix @@ -91,7 +108,34 @@ pub fn to_string(bytes: u64, format: Format) -> String { } } -#[derive(Debug)] +/// Parses a human-readable byte size string into a byte count, e.g. `"1.5 KiB"` to `1536`. +/// +/// Accepts a plain integer, or a number followed by an optional SI/IEC unit +/// (case-insensitive): `"1024"`, `"1.5 KB"`, `"2MiB"`, `"3 g"`. +pub fn parse(value: &str) -> Result { + if let Ok(v) = value.parse::() { + return Ok(v); + } + let number = take_while(value, |c| c.is_ascii_digit() || c == '.'); + match number.parse::() { + Ok(v) => { + let suffix = skip_while(&value[number.len()..], char::is_whitespace); + match suffix.parse::() { + Ok(u) => Ok((v * u64::from(u) as f64) as u64), + Err(error) => Err(ParseError(format!( + "couldn't parse {:?} into a known SI unit, {}", + suffix, error + ))), + } + } + Err(error) => Err(ParseError(format!( + "couldn't parse {:?} into a byte size, {}", + value, error + ))), + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] pub struct ParseError(pub String); impl core::fmt::Display for ParseError { @@ -164,7 +208,7 @@ impl From for u64 { } impl FromStr for Unit { - type Err = String; + type Err = ParseError; fn from_str(unit: &str) -> Result { match unit.to_lowercase().as_str() { @@ -181,7 +225,7 @@ impl FromStr for Unit { "gi" | "gib" => Ok(Self::GibiByte), "ti" | "tib" => Ok(Self::TebiByte), "pi" | "pib" => Ok(Self::PebiByte), - _ => Err(format!("couldn't parse unit of {:?}", unit)), + _ => Err(ParseError(format!("couldn't parse unit of {:?}", unit))), } } } @@ -209,3 +253,184 @@ impl> core::ops::RangeBounds for HumanByteRange { core::ops::Bound::Included(&self.stop) } } + +/// Serde support for plain integer fields, without declaring a newtype. +/// +/// Serializes as a human-readable string (`"1.5 KiB"`) in human-readable +/// formats (JSON, TOML, ...) and as a raw integer in binary formats. +/// Deserializes from either form. +/// +/// ```ignore +/// #[derive(Serialize, Deserialize)] +/// struct Config { +/// #[serde(with = "humanbyte::serde")] +/// buffer_size: usize, +/// #[serde(with = "humanbyte::serde::map_keys")] +/// pools: BTreeMap, +/// } +/// ``` +#[cfg(feature = "serde")] +pub mod serde { + use ::serde::{Deserializer, Serializer, de}; + + use crate::{Format, to_string}; + + pub fn serialize(value: &T, serializer: S) -> Result + where + T: Copy + TryInto, + S: Serializer, + { + let value: u64 = (*value) + .try_into() + .map_err(|_| ::serde::ser::Error::custom("byte size doesn't fit in u64"))?; + if serializer.is_human_readable() { + serializer.serialize_str(&to_string(value, Format::IEC)) + } else { + serializer.serialize_u64(value) + } + } + + pub fn deserialize<'de, T, D>(deserializer: D) -> Result + where + T: TryFrom, + D: Deserializer<'de>, + { + let value = if deserializer.is_human_readable() { + deserializer.deserialize_any(ByteVisitor)? + } else { + deserializer.deserialize_u64(ByteVisitor)? + }; + T::try_from(value).map_err(|_| de::Error::custom("byte size overflows target type")) + } + + struct ByteVisitor; + + impl de::Visitor<'_> for ByteVisitor { + type Value = u64; + + fn expecting(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { + formatter.write_str("an integer or a byte size string") + } + + fn visit_i64(self, value: i64) -> Result { + u64::try_from(value).map_err(|_| { + E::invalid_value(de::Unexpected::Signed(value), &"a non-negative integer") + }) + } + + fn visit_u64(self, value: u64) -> Result { + Ok(value) + } + + fn visit_str(self, value: &str) -> Result { + crate::parse(value) + .map_err(|_| E::invalid_value(de::Unexpected::Str(value), &"a byte size string")) + } + } + + /// Like the parent module, but for `BTreeMap`s keyed by byte sizes. + pub mod map_keys { + use alloc::collections::BTreeMap; + + use ::serde::ser::SerializeMap; + use ::serde::{Deserialize, Deserializer, Serialize, Serializer, de}; + + use super::{ByteVisitor, Format, to_string}; + + pub fn serialize(map: &BTreeMap, serializer: S) -> Result + where + K: Copy + TryInto, + V: Serialize, + S: Serializer, + { + let human_readable = serializer.is_human_readable(); + let mut ser = serializer.serialize_map(Some(map.len()))?; + for (key, value) in map { + let key: u64 = (*key) + .try_into() + .map_err(|_| ::serde::ser::Error::custom("byte size doesn't fit in u64"))?; + if human_readable { + ser.serialize_entry(&to_string(key, Format::IEC), value)?; + } else { + ser.serialize_entry(&key, value)?; + } + } + ser.end() + } + + pub fn deserialize<'de, K, V, D>(deserializer: D) -> Result, D::Error> + where + K: TryFrom + Ord, + V: Deserialize<'de>, + D: Deserializer<'de>, + { + struct Key(K); + + impl<'de, K: TryFrom> Deserialize<'de> for Key { + fn deserialize>(deserializer: D) -> Result { + let value = if deserializer.is_human_readable() { + deserializer.deserialize_any(ByteVisitor)? + } else { + deserializer.deserialize_u64(ByteVisitor)? + }; + K::try_from(value) + .map(Key) + .map_err(|_| de::Error::custom("byte size overflows target type")) + } + } + + struct MapVisitor(core::marker::PhantomData<(K, V)>); + + impl<'de, K, V> de::Visitor<'de> for MapVisitor + where + K: TryFrom + Ord, + V: Deserialize<'de>, + { + type Value = BTreeMap; + + fn expecting(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { + formatter.write_str("a map keyed by byte sizes") + } + + fn visit_map>( + self, + mut access: A, + ) -> Result { + let mut map = BTreeMap::new(); + while let Some((Key(key), value)) = access.next_entry::, V>()? { + map.insert(key, value); + } + Ok(map) + } + } + + deserializer.deserialize_map(MapVisitor(core::marker::PhantomData)) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse() { + assert_eq!(parse("1024"), Ok(1024)); + assert_eq!(parse("1 KiB"), Ok(1024)); + assert_eq!(parse("1.5KiB"), Ok(1536)); + assert_eq!(parse("2 mb"), Ok(2_000_000)); + assert_eq!(parse("3 g"), Ok(3_000_000_000)); + assert!(parse("").is_err()); + assert!(parse("1.5 XB").is_err()); + } + + #[test] + fn test_precision() { + assert_eq!(to_string_with_precision(TIB, Format::IEC, 2), "1.00 TiB"); + assert_eq!(to_string_with_precision(1536, Format::IEC, 0), "2 KiB"); + assert_eq!(to_string_with_precision(1536, Format::IEC, 3), "1.500 KiB"); + // sub-unit values are plain byte counts regardless of precision + assert_eq!(to_string_with_precision(215, Format::IEC, 2), "215 B"); + assert_eq!(to_string(TIB, Format::IEC), "1.0 TiB"); + } +}