Skip to main content

wdl_modules/
version_requirement.rs

1//! Version-requirement type wrapping a constrained subset of
2//! [`semver::VersionReq`].
3
4use std::fmt;
5use std::ops::Deref;
6use std::str::FromStr;
7
8use semver::VersionReq;
9use serde_with::DeserializeFromStr;
10use serde_with::SerializeDisplay;
11use thiserror::Error;
12
13/// An error parsing a [`VersionRequirement`].
14#[derive(Debug, Error)]
15#[error("`{0}` is not a valid semantic version requirement")]
16pub struct VersionRequirementError(String);
17
18/// A version requirement, parsed by [`semver::VersionReq`].
19#[derive(Clone, Debug, PartialEq, Eq, SerializeDisplay, DeserializeFromStr)]
20pub struct VersionRequirement(VersionReq);
21
22impl VersionRequirement {
23    /// Consumes the [`VersionRequirement`] and returns the inner
24    /// [`VersionReq`].
25    pub fn into_inner(self) -> VersionReq {
26        self.0
27    }
28}
29
30impl Deref for VersionRequirement {
31    type Target = VersionReq;
32
33    fn deref(&self) -> &VersionReq {
34        &self.0
35    }
36}
37
38impl fmt::Display for VersionRequirement {
39    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
40        self.0.fmt(f)
41    }
42}
43
44impl FromStr for VersionRequirement {
45    type Err = VersionRequirementError;
46
47    fn from_str(s: &str) -> Result<Self, Self::Err> {
48        match VersionReq::parse(s.trim()) {
49            Ok(v) => Ok(Self(v)),
50            Err(_) => Err(VersionRequirementError(s.to_string())),
51        }
52    }
53}
54
55impl From<VersionRequirement> for String {
56    fn from(req: VersionRequirement) -> Self {
57        req.0.to_string()
58    }
59}
60
61#[cfg(test)]
62mod tests {
63    use semver::Version;
64
65    use super::*;
66
67    #[test]
68    fn accepts_spec_operators() {
69        for s in [
70            "^1.2.0",
71            "~1.2.0",
72            "=1.2.0",
73            ">=1.0.0, <2.0.0",
74            ">1.0.0",
75            "<2.0.0",
76            ">=1.0.0",
77            "<=2.0.0",
78            "*",
79            "1.2.0",
80        ] {
81            assert!(s.parse::<VersionRequirement>().is_ok(), "rejected `{s}`");
82        }
83    }
84
85    #[test]
86    fn rejects_invalid_format() {
87        for bad in ["", "   ", "not-a-req"] {
88            assert!(
89                bad.parse::<VersionRequirement>().is_err(),
90                "accepted `{bad}`"
91            );
92        }
93    }
94
95    #[test]
96    fn error_message_includes_input() {
97        let err = "^1.foo".parse::<VersionRequirement>().unwrap_err();
98        assert_eq!(
99            err.to_string(),
100            "`^1.foo` is not a valid semantic version requirement"
101        );
102    }
103
104    #[test]
105    fn matches_versions_correctly() {
106        let req: VersionRequirement = "^1.2.0".parse().unwrap();
107        assert!(req.matches(&Version::parse("1.2.0").unwrap()));
108        assert!(req.matches(&Version::parse("1.9.99").unwrap()));
109        assert!(!req.matches(&Version::parse("2.0.0").unwrap()));
110        assert!(!req.matches(&Version::parse("1.1.0").unwrap()));
111    }
112
113    #[test]
114    fn round_trips_via_serde() {
115        let req: VersionRequirement = "^1.2.0".parse().unwrap();
116        let json = serde_json::to_string(&req).unwrap();
117        assert_eq!(json, r#""^1.2.0""#);
118        let parsed: VersionRequirement = serde_json::from_str(&json).unwrap();
119        assert_eq!(parsed, req);
120    }
121}