use std::{
cmp::Ordering,
fmt::{Display, Formatter},
str::FromStr,
};
use alpm_parsers::iter_str_context;
use serde::{Deserialize, Serialize};
use strum::VariantNames;
use winnow::{
ModalResult,
Parser,
combinator::{alt, eof, fail, seq},
error::{StrContext, StrContextValue},
token::take_while,
};
use crate::{Error, Version};
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct VersionRequirement {
pub comparison: VersionComparison,
pub version: Version,
}
impl VersionRequirement {
pub fn new(comparison: VersionComparison, version: Version) -> Self {
VersionRequirement {
comparison,
version,
}
}
pub fn is_satisfied_by(&self, ver: &Version) -> bool {
let other_version = if self.version.pkgrel.is_none() {
&Version {
pkgrel: None,
..ver.clone()
}
} else {
ver
};
self.comparison
.is_compatible_with(other_version.cmp(&self.version))
}
pub fn parser(input: &mut &str) -> ModalResult<Self> {
seq!(Self {
comparison: take_while(1.., ('<', '>', '='))
.context(StrContext::Expected(StrContextValue::Description(
"version comparison operator"
)))
.and_then(VersionComparison::parser),
version: Version::parser,
})
.parse_next(input)
}
pub fn is_intersection(&self, other: &VersionRequirement) -> bool {
let version_comparison = self.version.cmp(&other.version);
match self.comparison {
VersionComparison::Less => {
match version_comparison {
Ordering::Less => matches!(
other.comparison,
VersionComparison::Less | VersionComparison::LessOrEqual
),
Ordering::Equal => matches!(
other.comparison,
VersionComparison::Less | VersionComparison::LessOrEqual
),
Ordering::Greater => true,
}
}
VersionComparison::LessOrEqual => {
match version_comparison {
Ordering::Less => matches!(
other.comparison,
VersionComparison::Less | VersionComparison::LessOrEqual
),
Ordering::Equal => matches!(
other.comparison,
VersionComparison::Less
| VersionComparison::LessOrEqual
| VersionComparison::Equal
| VersionComparison::GreaterOrEqual
),
Ordering::Greater => true,
}
}
VersionComparison::Equal => match version_comparison {
Ordering::Equal => matches!(
other.comparison,
VersionComparison::LessOrEqual
| VersionComparison::Equal
| VersionComparison::GreaterOrEqual
),
Ordering::Less | Ordering::Greater => false,
},
VersionComparison::GreaterOrEqual => match version_comparison {
Ordering::Less => true,
Ordering::Equal => matches!(
other.comparison,
VersionComparison::LessOrEqual
| VersionComparison::Equal
| VersionComparison::GreaterOrEqual
| VersionComparison::Greater
),
Ordering::Greater => matches!(
other.comparison,
VersionComparison::GreaterOrEqual | VersionComparison::Greater
),
},
VersionComparison::Greater => {
match version_comparison {
Ordering::Less => true,
Ordering::Equal => matches!(
other.comparison,
VersionComparison::GreaterOrEqual | VersionComparison::Greater
),
Ordering::Greater => matches!(
other.comparison,
VersionComparison::GreaterOrEqual | VersionComparison::Greater
),
}
}
}
}
}
impl Display for VersionRequirement {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "{}{}", self.comparison, self.version)
}
}
impl FromStr for VersionRequirement {
type Err = Error;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(Self::parser.parse(s)?)
}
}
#[derive(
strum::AsRefStr,
Clone,
Copy,
Debug,
strum::Display,
strum::EnumIter,
PartialEq,
Eq,
strum::VariantNames,
Serialize,
Deserialize,
)]
pub enum VersionComparison {
#[strum(to_string = "<=")]
LessOrEqual,
#[strum(to_string = ">=")]
GreaterOrEqual,
#[strum(to_string = "=")]
Equal,
#[strum(to_string = "<")]
Less,
#[strum(to_string = ">")]
Greater,
}
impl VersionComparison {
fn is_compatible_with(self, ord: Ordering) -> bool {
match (self, ord) {
(VersionComparison::Less, Ordering::Less)
| (VersionComparison::LessOrEqual, Ordering::Less | Ordering::Equal)
| (VersionComparison::Equal, Ordering::Equal)
| (VersionComparison::GreaterOrEqual, Ordering::Greater | Ordering::Equal)
| (VersionComparison::Greater, Ordering::Greater) => true,
(VersionComparison::Less, Ordering::Equal | Ordering::Greater)
| (VersionComparison::LessOrEqual, Ordering::Greater)
| (VersionComparison::Equal, Ordering::Less | Ordering::Greater)
| (VersionComparison::GreaterOrEqual, Ordering::Less)
| (VersionComparison::Greater, Ordering::Less | Ordering::Equal) => false,
}
}
pub fn parser(input: &mut &str) -> ModalResult<Self> {
alt((
("<=", eof).value(Self::LessOrEqual),
(">=", eof).value(Self::GreaterOrEqual),
("=", eof).value(Self::Equal),
("<", eof).value(Self::Less),
(">", eof).value(Self::Greater),
fail.context(StrContext::Label("comparison operator"))
.context_with(iter_str_context!([VersionComparison::VARIANTS])),
))
.parse_next(input)
}
}
impl FromStr for VersionComparison {
type Err = Error;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(Self::parser.parse(s)?)
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use testresult::TestResult;
use super::*;
#[rstest]
#[case("<", VersionComparison::Less)]
#[case("<=", VersionComparison::LessOrEqual)]
#[case("=", VersionComparison::Equal)]
#[case(">=", VersionComparison::GreaterOrEqual)]
#[case(">", VersionComparison::Greater)]
fn valid_version_comparison(#[case] comparison: &str, #[case] expected: VersionComparison) {
assert_eq!(comparison.parse(), Ok(expected));
}
#[rstest]
#[case("", "invalid comparison operator")]
#[case("<<", "invalid comparison operator")]
#[case("==", "invalid comparison operator")]
#[case("!=", "invalid comparison operator")]
#[case(" =", "invalid comparison operator")]
#[case("= ", "invalid comparison operator")]
#[case("<1", "invalid comparison operator")]
fn invalid_version_comparison(#[case] comparison: &str, #[case] err_snippet: &str) {
let Err(Error::ParseError(err_msg)) = VersionComparison::from_str(comparison) else {
panic!("'{comparison}' did not fail as expected")
};
assert!(
err_msg.contains(err_snippet),
"Error:\n=====\n{err_msg}\n=====\nshould contain snippet:\n\n{err_snippet}"
);
}
#[rstest]
#[case("=1", VersionRequirement {
comparison: VersionComparison::Equal,
version: Version::from_str("1").unwrap(),
})]
#[case("<=42:abcd-2.4", VersionRequirement {
comparison: VersionComparison::LessOrEqual,
version: Version::from_str("42:abcd-2.4").unwrap(),
})]
#[case(">3.1", VersionRequirement {
comparison: VersionComparison::Greater,
version: Version::from_str("3.1").unwrap(),
})]
fn valid_version_requirement(#[case] requirement: &str, #[case] expected: VersionRequirement) {
assert_eq!(
requirement.parse(),
Ok(expected),
"Expected successful parse for version requirement '{requirement}'"
);
}
#[rstest]
#[case::bad_operator("<>3.1", "invalid comparison operator")]
#[case::no_operator("3.1", "expected version comparison operator")]
#[case::arrow_operator("=>3.1", "invalid comparison operator")]
#[case::no_version("<=", "expected pkgver string")]
fn invalid_version_requirement(#[case] requirement: &str, #[case] err_snippet: &str) {
let Err(Error::ParseError(err_msg)) = VersionRequirement::from_str(requirement) else {
panic!("'{requirement}' erroneously parsed as VersionRequirement")
};
assert!(
err_msg.contains(err_snippet),
"Error:\n=====\n{err_msg}\n=====\nshould contain snippet:\n\n{err_snippet}"
);
}
#[rstest]
#[case("<3.1>3.2", "invalid pkgver character")]
fn invalid_version_requirement_pkgver_parse(
#[case] requirement: &str,
#[case] err_snippet: &str,
) {
let Err(Error::ParseError(err_msg)) = VersionRequirement::from_str(requirement) else {
panic!("'{requirement}' erroneously parsed as VersionRequirement")
};
assert!(
err_msg.contains(err_snippet),
"Error:\n=====\n{err_msg}\n=====\nshould contain snippet:\n\n{err_snippet}"
);
}
#[rstest]
#[case("=1", "1", true)]
#[case("=1", "1.0", false)]
#[case("=1", "1-1", true)]
#[case("=1", "1:1", false)]
#[case("=1", "0.9", false)]
#[case("<42", "41", true)]
#[case("<42", "42", false)]
#[case("<42", "43", false)]
#[case("<=42", "41", true)]
#[case("<=42", "42", true)]
#[case("<=42", "43", false)]
#[case(">42", "41", false)]
#[case(">42", "42", false)]
#[case(">42", "43", true)]
#[case(">=42", "41", false)]
#[case(">=42", "42", true)]
#[case(">=42", "43", true)]
fn version_requirement_satisfied(
#[case] requirement: &str,
#[case] version: &str,
#[case] result: bool,
) {
let requirement = VersionRequirement::from_str(requirement).unwrap();
let version = Version::from_str(version).unwrap();
assert_eq!(requirement.is_satisfied_by(&version), result);
}
#[rstest]
#[case::self_less_matching_other_less("<1", "<1")]
#[case::self_less_matching_other_less_or_equal("<1", "<=1")]
#[case::self_less_bigger_other_less("<1", "<2")]
#[case::self_less_bigger_other_less_or_equal("<1", "<=2")]
#[case::self_less_smaller_other_less("<1", "<0.1")]
#[case::self_less_smaller_other_less_or_equal("<1", "<=0.1")]
#[case::self_less_smaller_other_equal("<1", "=0.1")]
#[case::self_less_smaller_other_greater_or_equal("<1", ">=0.1")]
#[case::self_less_smaller_other_greater("<1", ">0.1")]
#[case::self_less_smaller_other_equal("<1", "=0.1")]
#[case::self_less_or_equal_matching_other_less("<=1", "<1")]
#[case::self_less_or_equal_matching_other_less_or_equal("<=1", "<=1")]
#[case::self_less_or_equal_matching_other_equal("<=1", "=1")]
#[case::self_less_or_equal_matching_other_greater_or_equal("<=1", ">=1")]
#[case::self_less_or_equal_bigger_other_less("<=1", "<2")]
#[case::self_less_or_equal_bigger_other_less_or_equal("<=1", "<=2")]
#[case::self_less_or_equal_smaller_other_greater_or_equal("<=1", ">=0.1")]
#[case::self_less_or_equal_smaller_other_greater("<=1", ">0.1")]
#[case::self_equal_matching_other_less_or_equal("=1", "<=1")]
#[case::self_equal_matching_other_equal("=1", "=1")]
#[case::self_equal_matching_other_greater_or_equal("=1", ">=1")]
#[case::self_greater_or_equal_matching_other_less_or_equal(">=1", "<=1")]
#[case::self_greater_or_equal_matching_other_equal(">=1", "=1")]
#[case::self_greater_or_equal_matching_other_greater_or_equal(">=1", ">=1")]
#[case::self_greater_or_equal_matching_other_greater(">=1", ">1")]
#[case::self_greater_or_equal_bigger_other_less(">=1", "<2")]
#[case::self_greater_or_equal_bigger_other_less_or_equal(">=1", "<=2")]
#[case::self_greater_or_equal_bigger_other_equal(">=1", "=2")]
#[case::self_greater_or_equal_bigger_other_greater_or_equal(">=1", ">=2")]
#[case::self_greater_or_equal_bigger_other_greater(">=1", ">2")]
#[case::self_greater_or_equal_smaller_other_greater_or_equal(">=1", ">=0.1")]
#[case::self_greater_or_equal_smaller_other_greater(">=1", ">0.1")]
#[case::self_greater_matching_other_greater_or_equal(">1", ">=1")]
#[case::self_greater_matching_other_greater(">1", ">1")]
#[case::self_greater_bigger_other_less(">1", "<2")]
#[case::self_greater_bigger_other_less_or_equal(">1", "<=2")]
#[case::self_greater_bigger_other_equal(">1", "=2")]
#[case::self_greater_bigger_other_greater_or_equal(">1", ">=2")]
#[case::self_greater_bigger_other_greater(">1", ">2")]
#[case::self_greater_smaller_other_greater_or_equal(">1", ">=0.1")]
#[case::self_greater_smaller_other_greater(">1", ">0.1")]
fn version_requirements_form_intersection(
#[case] self_requirement: &str,
#[case] other_requirement: &str,
) -> TestResult {
let self_requirement: VersionRequirement = self_requirement.parse()?;
let other_requirement: VersionRequirement = other_requirement.parse()?;
assert!(self_requirement.is_intersection(&other_requirement));
Ok(())
}
#[rstest]
#[case::self_less_matching_other_equal("<1", "=1")]
#[case::self_less_matching_other_greater_or_equal("<1", ">=1")]
#[case::self_less_matching_other_greater("<1", ">1")]
#[case::self_less_or_equal_matching_other_greater("<=1", ">1")]
#[case::self_equal_matching_other_less("=1", "<1")]
#[case::self_equal_matching_other_greater("=1", ">1")]
#[case::self_equal_bigger_other_less("=1", "<2")]
#[case::self_equal_bigger_other_greater("=1", ">2")]
#[case::self_equal_smaller_other_less("=1", "<0.1")]
#[case::self_equal_smaller_other_greater("=1", ">0.1")]
#[case::self_greater_or_equal_matching_other_less(">=1", "<1")]
#[case::self_greater_matching_other_less(">1", "<1")]
#[case::self_greater_matching_other_less_or_equal(">1", "<=1")]
#[case::self_greater_matching_other_equal(">1", "=1")]
fn version_requirements_do_not_form_intersection(
#[case] self_requirement: &str,
#[case] other_requirement: &str,
) -> TestResult {
let self_requirement: VersionRequirement = self_requirement.parse()?;
let other_requirement: VersionRequirement = other_requirement.parse()?;
assert!(!self_requirement.is_intersection(&other_requirement));
Ok(())
}
}