Skip to main content

ante_protocol_shape/
id.rs

1use serde::{Deserialize, Serialize};
2use ulid::Ulid;
3
4#[derive(Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Default, Copy)]
5pub struct Id {
6    prefix: [u8; Self::PREFIX_SIZE],
7    id: Ulid,
8}
9
10impl Id {
11    pub const PREFIX_SIZE: usize = 4;
12
13    pub fn new<S: AsRef<str>>(id: S) -> Self {
14        Self { prefix: Self::clamp_prefix(id), id: Ulid::generate() }
15    }
16
17    fn clamp_prefix<S: AsRef<str>>(id: S) -> [u8; Self::PREFIX_SIZE] {
18        let mut prefix = [0u8; Self::PREFIX_SIZE];
19        let bytes = id.as_ref().as_bytes();
20        let len = std::cmp::min(bytes.len(), Self::PREFIX_SIZE);
21        prefix[..len].copy_from_slice(&bytes[..len]);
22        prefix
23    }
24
25    pub fn op() -> Self {
26        Self::new("op")
27    }
28
29    pub fn evt() -> Self {
30        Self::new("evt")
31    }
32
33    pub fn ses() -> Self {
34        Self::new("ses")
35    }
36
37    pub fn step() -> Self {
38        Self::new("step")
39    }
40}
41
42impl std::fmt::Display for Id {
43    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
44        for &b in self.prefix.iter().take_while(|&&b| b != 0) {
45            write!(f, "{}", b as char)?;
46        }
47        write!(f, "_{}", self.id)
48    }
49}
50
51impl std::fmt::Debug for Id {
52    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
53        write!(f, "{self}")
54    }
55}
56
57impl Serialize for Id {
58    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
59    where
60        S: serde::Serializer,
61    {
62        serializer.serialize_str(&self.to_string())
63    }
64}
65
66impl<'de> Deserialize<'de> for Id {
67    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
68    where
69        D: serde::Deserializer<'de>,
70    {
71        let value = String::deserialize(deserializer)?;
72        value.parse().map_err(serde::de::Error::custom)
73    }
74}
75
76#[derive(Debug, Clone, PartialEq, Eq)]
77pub enum ParseIdError {
78    MissingSeparator,
79    InvalidUlid,
80}
81
82impl std::fmt::Display for ParseIdError {
83    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
84        match self {
85            Self::MissingSeparator => write!(f, "missing separator"),
86            Self::InvalidUlid => write!(f, "invalid ULID"),
87        }
88    }
89}
90
91impl std::error::Error for ParseIdError {}
92
93impl std::str::FromStr for Id {
94    type Err = ParseIdError;
95
96    fn from_str(s: &str) -> Result<Self, Self::Err> {
97        let (prefix_str, ulid_str) = s.split_once('_').ok_or(ParseIdError::MissingSeparator)?;
98        let id = Ulid::from_string(ulid_str).map_err(|_| ParseIdError::InvalidUlid)?;
99        Ok(Self { prefix: Self::clamp_prefix(prefix_str), id })
100    }
101}
102
103#[cfg(test)]
104mod tests {
105    use super::*;
106    use std::str::FromStr;
107
108    #[test]
109    fn test_id_serde_roundtrip() {
110        let id = Id::new("evt");
111        let serialized = serde_json::to_string(&id).expect("serialize id");
112        assert_eq!(serialized, format!("\"{id}\""));
113
114        let deserialized: Id = serde_json::from_str(&serialized).expect("deserialize id");
115        assert_eq!(deserialized, id);
116    }
117
118    #[test]
119    fn test_id_from_str() {
120        let id = Id::new("ses");
121        let parsed = Id::from_str(&id.to_string()).expect("parse id");
122        assert_eq!(parsed, id);
123    }
124}