use std::borrow::Cow;
use std::fmt;
use schemars::{json_schema, JsonSchema, Schema, SchemaGenerator};
use serde::{de, Deserialize, Deserializer, Serialize, Serializer};
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
pub struct Version(Vec<u64>);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VersionReq(Vec<u64>);
fn parse_components(s: &str) -> Result<Vec<u64>, String> {
let s = s.trim();
if s.is_empty() {
return Err("version is empty".to_string());
}
let parts: Vec<&str> = s.split('.').collect();
if parts.len() > 3 {
return Err(format!(
"version {s:?} has too many components (expected 1 to 3, like 1, 1.2, or 1.2.3)"
));
}
let mut out = Vec::with_capacity(parts.len());
for p in parts {
let n: u64 = p
.parse()
.map_err(|_| format!("version {s:?} component {p:?} is not a non-negative integer"))?;
out.push(n);
}
Ok(out)
}
fn join(components: &[u64]) -> String {
components
.iter()
.map(u64::to_string)
.collect::<Vec<_>>()
.join(".")
}
impl Version {
pub fn parse(s: &str) -> Result<Self, String> {
Ok(Version(parse_components(s)?))
}
pub fn components(&self) -> &[u64] {
&self.0
}
}
impl fmt::Display for Version {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&join(&self.0))
}
}
impl VersionReq {
pub fn parse(s: &str) -> Result<Self, String> {
Ok(VersionReq(parse_components(s)?))
}
pub fn matches(&self, version: &Version) -> bool {
version.0.len() >= self.0.len() && version.0[..self.0.len()] == self.0[..]
}
}
impl fmt::Display for VersionReq {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&join(&self.0))
}
}
impl Serialize for Version {
fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
s.serialize_str(&self.to_string())
}
}
impl<'de> Deserialize<'de> for Version {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
Version::parse(&scalar_text(d, "version")?).map_err(de::Error::custom)
}
}
impl Serialize for VersionReq {
fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
s.serialize_str(&self.to_string())
}
}
impl<'de> Deserialize<'de> for VersionReq {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
VersionReq::parse(&scalar_text(d, "version pin")?).map_err(de::Error::custom)
}
}
fn scalar_text<'de, D: Deserializer<'de>>(d: D, what: &str) -> Result<String, D::Error> {
let value = serde_yaml_ng::Value::deserialize(d)?;
match &value {
serde_yaml_ng::Value::Number(n) => Ok(n.to_string()),
serde_yaml_ng::Value::String(s) => Ok(s.clone()),
other => Err(de::Error::custom(format!(
"{what} must be a number or string like 1, 1.2, or \"1.2.3\"; got {other:?}"
))),
}
}
impl JsonSchema for Version {
fn schema_name() -> Cow<'static, str> {
"Version".into()
}
fn json_schema(_generator: &mut SchemaGenerator) -> Schema {
json_schema!({
"type": ["integer", "number", "string"],
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_one_to_three_components() {
assert_eq!(Version::parse("1").unwrap().components(), &[1]);
assert_eq!(Version::parse("1.2").unwrap().components(), &[1, 2]);
assert_eq!(Version::parse(" 1.2.3 ").unwrap().components(), &[1, 2, 3]);
}
#[test]
fn rejects_malformed_versions() {
assert!(Version::parse("").is_err());
assert!(Version::parse("1.2.3.4").is_err());
assert!(Version::parse("1.x").is_err());
assert!(Version::parse("-1").is_err());
}
#[test]
fn display_roundtrips() {
assert_eq!(Version::parse("1.2.3").unwrap().to_string(), "1.2.3");
assert_eq!(VersionReq::parse("1.2").unwrap().to_string(), "1.2");
}
#[test]
fn pin_matches_by_prefix() {
let req = VersionReq::parse("1").unwrap();
assert!(req.matches(&Version::parse("1").unwrap()));
assert!(req.matches(&Version::parse("1.4").unwrap()));
assert!(req.matches(&Version::parse("1.4.2").unwrap()));
assert!(!req.matches(&Version::parse("2.0").unwrap()));
let req = VersionReq::parse("1.4.2").unwrap();
assert!(req.matches(&Version::parse("1.4.2").unwrap()));
assert!(!req.matches(&Version::parse("1.4").unwrap()));
assert!(!req.matches(&Version::parse("1.4.3").unwrap()));
}
#[test]
fn versions_order_by_components() {
let v = |s: &str| Version::parse(s).unwrap();
assert!(v("1.4") > v("1.2"));
assert!(v("1.10") > v("1.9"));
assert!(v("2") > v("1.99.99"));
assert!(v("1") < v("1.0"));
let mut all = [v("1.2"), v("2.0"), v("1.10")];
all.sort();
assert_eq!(all.iter().max().unwrap(), &v("2.0"));
}
#[test]
fn deserializes_from_int_float_and_string() {
let v: Version = serde_yaml_ng::from_str("1").unwrap();
assert_eq!(v.to_string(), "1");
let v: Version = serde_yaml_ng::from_str("1.2").unwrap();
assert_eq!(v.to_string(), "1.2");
let v: Version = serde_yaml_ng::from_str("\"1.2.3\"").unwrap();
assert_eq!(v.to_string(), "1.2.3");
}
#[test]
fn rejects_non_scalar_version() {
assert!(serde_yaml_ng::from_str::<Version>("[1, 2]").is_err());
}
#[test]
fn serializes_as_string() {
let v = Version::parse("1.2").unwrap();
assert_eq!(serde_json::to_string(&v).unwrap(), "\"1.2\"");
let req = VersionReq::parse("1").unwrap();
assert_eq!(serde_json::to_string(&req).unwrap(), "\"1\"");
}
#[test]
fn pins_round_trip_through_a_persisted_record() {
let req = VersionReq::parse("1.2").unwrap();
let json = serde_json::to_string(&req).unwrap();
assert_eq!(serde_json::from_str::<VersionReq>(&json).unwrap(), req);
assert_eq!(
serde_json::from_str::<VersionReq>("1").unwrap(),
VersionReq::parse("1").unwrap()
);
assert!(serde_json::from_str::<VersionReq>("\"1.x\"").is_err());
assert!(serde_json::from_str::<VersionReq>("[1]").is_err());
assert!(serde_json::from_str::<Version>("[1]").is_err());
}
}