Skip to main content

a3s_runtime/
provider.rs

1use crate::{RuntimeError, RuntimeResult};
2use serde::{Deserialize, Deserializer, Serialize, Serializer};
3use std::fmt;
4
5#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
6pub struct ProviderId(String);
7
8impl ProviderId {
9    pub fn parse(value: impl Into<String>) -> RuntimeResult<Self> {
10        let value = value.into();
11        let valid = !value.is_empty()
12            && value.len() <= 64
13            && value.bytes().enumerate().all(|(index, byte)| {
14                byte.is_ascii_lowercase()
15                    || byte.is_ascii_digit()
16                    || (byte == b'-' && index > 0 && index + 1 < value.len())
17            });
18        if !valid {
19            return Err(RuntimeError::InvalidRequest(format!(
20                "Runtime provider ID {value:?} must use lowercase ASCII letters, digits, and interior hyphens"
21            )));
22        }
23        Ok(Self(value))
24    }
25
26    pub fn as_str(&self) -> &str {
27        &self.0
28    }
29}
30
31impl fmt::Display for ProviderId {
32    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
33        self.0.fmt(formatter)
34    }
35}
36
37impl Serialize for ProviderId {
38    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
39    where
40        S: Serializer,
41    {
42        serializer.serialize_str(&self.0)
43    }
44}
45
46impl<'de> Deserialize<'de> for ProviderId {
47    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
48    where
49        D: Deserializer<'de>,
50    {
51        let value = String::deserialize(deserializer)?;
52        Self::parse(value).map_err(serde::de::Error::custom)
53    }
54}
55
56#[cfg(test)]
57mod tests {
58    use super::*;
59
60    #[test]
61    fn provider_ids_are_portable_and_have_no_default_policy() {
62        for invalid in ["", "Docker", "a3s_box", "-docker", "docker-", "a/b"] {
63            assert!(ProviderId::parse(invalid).is_err(), "accepted {invalid:?}");
64        }
65        assert_eq!(
66            ProviderId::parse("vendor-runtime-2").unwrap().as_str(),
67            "vendor-runtime-2"
68        );
69    }
70}