wdl_modules/
version_requirement.rs1use 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#[derive(Debug, Error)]
15#[error("`{0}` is not a valid semantic version requirement")]
16pub struct VersionRequirementError(String);
17
18#[derive(Clone, Debug, PartialEq, Eq, SerializeDisplay, DeserializeFromStr)]
20pub struct VersionRequirement(VersionReq);
21
22impl VersionRequirement {
23 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}