Skip to main content

wowlab_common/
node_public_key.rs

1use std::fmt;
2
3use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD as BASE64};
4use serde::{Deserialize, Serialize};
5
6const PUBLIC_KEY_LENGTH: usize = 32;
7
8/// A node's Ed25519 public key (32 bytes, base64-encoded for display/storage).
9#[derive(Clone, Deserialize, Eq, Hash, PartialEq)]
10#[serde(try_from = "String")]
11pub struct NodePublicKey([u8; PUBLIC_KEY_LENGTH]);
12
13/// Error returned when a base64 string cannot be decoded into a valid 32-byte public key.
14#[derive(Debug, thiserror::Error)]
15#[error("invalid node public key: expected 32 bytes base64")]
16pub struct InvalidNodePublicKey {
17    #[source]
18    kind: InvalidNodePublicKeyKind,
19}
20
21#[derive(Debug, thiserror::Error)]
22enum InvalidNodePublicKeyKind {
23    #[error("base64 decoding failed")]
24    Base64(#[source] base64::DecodeError),
25    #[error("decoded public key contained {actual} bytes")]
26    Length { actual: usize },
27}
28
29impl InvalidNodePublicKey {
30    /// Return the decoded byte length when the base64 text decoded successfully but had the wrong length.
31    #[must_use]
32    pub const fn actual_length(&self) -> Option<usize> {
33        match self.kind {
34            InvalidNodePublicKeyKind::Length { actual } => Some(actual),
35            InvalidNodePublicKeyKind::Base64(_) => None,
36        }
37    }
38
39    /// Return whether the input failed base64 decoding.
40    #[must_use]
41    pub const fn is_invalid_base64(&self) -> bool {
42        matches!(self.kind, InvalidNodePublicKeyKind::Base64(_))
43    }
44}
45
46impl From<Vec<u8>> for InvalidNodePublicKey {
47    fn from(bytes: Vec<u8>) -> Self {
48        Self {
49            kind: InvalidNodePublicKeyKind::Length {
50                actual: bytes.len(),
51            },
52        }
53    }
54}
55
56impl From<base64::DecodeError> for InvalidNodePublicKey {
57    fn from(source: base64::DecodeError) -> Self {
58        Self {
59            kind: InvalidNodePublicKeyKind::Base64(source),
60        }
61    }
62}
63
64impl NodePublicKey {
65    /// Create from raw 32-byte key.
66    #[must_use]
67    pub fn from_bytes(bytes: [u8; PUBLIC_KEY_LENGTH]) -> Self {
68        Self(bytes)
69    }
70
71    /// Parse from base64 string.
72    ///
73    /// # Errors
74    ///
75    /// Returns [`InvalidNodePublicKey`] when the value is not base64 or is not exactly 32 bytes.
76    pub fn from_base64(s: &str) -> Result<Self, InvalidNodePublicKey> {
77        let bytes = BASE64.decode(s)?;
78        let arr: [u8; PUBLIC_KEY_LENGTH] = bytes.try_into()?;
79
80        Ok(Self(arr))
81    }
82
83    /// Get raw bytes.
84    #[must_use]
85    pub fn as_bytes(&self) -> &[u8; PUBLIC_KEY_LENGTH] {
86        &self.0
87    }
88
89    /// Get base64 representation.
90    #[must_use]
91    pub fn to_base64(&self) -> String {
92        BASE64.encode(self.0)
93    }
94}
95
96impl fmt::Display for NodePublicKey {
97    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
98        write!(f, "{}", self.to_base64())
99    }
100}
101
102impl fmt::Debug for NodePublicKey {
103    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
104        write!(f, "NodePublicKey({})", self.to_base64())
105    }
106}
107
108impl std::str::FromStr for NodePublicKey {
109    type Err = InvalidNodePublicKey;
110
111    fn from_str(s: &str) -> Result<Self, Self::Err> {
112        Self::from_base64(s)
113    }
114}
115
116impl TryFrom<String> for NodePublicKey {
117    type Error = InvalidNodePublicKey;
118
119    fn try_from(value: String) -> Result<Self, Self::Error> {
120        Self::from_base64(&value)
121    }
122}
123
124impl Serialize for NodePublicKey {
125    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
126    where
127        S: serde::Serializer,
128    {
129        serializer.serialize_str(&self.to_base64())
130    }
131}
132
133#[cfg(feature = "crypto")]
134impl From<ed25519_dalek::VerifyingKey> for NodePublicKey {
135    fn from(key: ed25519_dalek::VerifyingKey) -> Self {
136        Self(*key.as_bytes())
137    }
138}
139
140#[cfg(feature = "crypto")]
141impl From<&ed25519_dalek::VerifyingKey> for NodePublicKey {
142    fn from(key: &ed25519_dalek::VerifyingKey) -> Self {
143        Self(*key.as_bytes())
144    }
145}
146
147#[cfg(feature = "crypto")]
148impl TryFrom<&NodePublicKey> for ed25519_dalek::VerifyingKey {
149    type Error = ed25519_dalek::SignatureError;
150
151    fn try_from(key: &NodePublicKey) -> Result<Self, Self::Error> {
152        ed25519_dalek::VerifyingKey::from_bytes(&key.0)
153    }
154}
155
156#[cfg(feature = "sqlx")]
157impl sqlx::Encode<'_, sqlx::Postgres> for NodePublicKey {
158    fn encode_by_ref(
159        &self,
160        buf: &mut sqlx::postgres::PgArgumentBuffer,
161    ) -> Result<sqlx::encode::IsNull, sqlx::error::BoxDynError> {
162        <String as sqlx::Encode<sqlx::Postgres>>::encode_by_ref(&self.to_base64(), buf)
163    }
164}
165
166#[cfg(feature = "sqlx")]
167impl sqlx::Decode<'_, sqlx::Postgres> for NodePublicKey {
168    fn decode(value: sqlx::postgres::PgValueRef<'_>) -> Result<Self, sqlx::error::BoxDynError> {
169        let s = <&str as sqlx::Decode<sqlx::Postgres>>::decode(value)?;
170
171        Self::from_base64(s).map_err(|e| Box::new(e) as sqlx::error::BoxDynError)
172    }
173}
174
175#[cfg(feature = "sqlx")]
176impl sqlx::Type<sqlx::Postgres> for NodePublicKey {
177    fn type_info() -> sqlx::postgres::PgTypeInfo {
178        <String as sqlx::Type<sqlx::Postgres>>::type_info()
179    }
180
181    fn compatible(ty: &sqlx::postgres::PgTypeInfo) -> bool {
182        <String as sqlx::Type<sqlx::Postgres>>::compatible(ty)
183    }
184}
185
186#[cfg(feature = "sqlx")]
187impl sqlx::postgres::PgHasArrayType for NodePublicKey {
188    fn array_type_info() -> sqlx::postgres::PgTypeInfo {
189        <String as sqlx::postgres::PgHasArrayType>::array_type_info()
190    }
191
192    fn array_compatible(ty: &sqlx::postgres::PgTypeInfo) -> bool {
193        <String as sqlx::postgres::PgHasArrayType>::array_compatible(ty)
194    }
195}
196
197#[cfg(test)]
198mod tests {
199    use std::{error::Error as _, str::FromStr};
200
201    use googletest::prelude::*;
202    use rstest::rstest;
203
204    use super::*;
205
206    fn zero_key_b64() -> String {
207        NodePublicKey::from_bytes([0u8; PUBLIC_KEY_LENGTH]).to_base64()
208    }
209
210    #[gtest]
211    fn from_base64_valid_32_bytes() -> Result<()> {
212        let key = NodePublicKey::from_base64(&zero_key_b64())?;
213
214        verify_that!(key.as_bytes(), eq(&[0u8; PUBLIC_KEY_LENGTH]))
215    }
216
217    #[gtest]
218    #[rstest]
219    #[case::invalid_chars("!!!not base64!!!")]
220    fn from_base64_invalid_literals(#[case] input: &str) -> Result<()> {
221        let error = NodePublicKey::from_base64(input).err().or_fail()?;
222
223        verify_that!(error.is_invalid_base64(), eq(true))?;
224
225        verify_that!(error.source(), some(anything()))
226    }
227
228    #[gtest]
229    fn from_base64_wrong_length_short() -> Result<()> {
230        let short = BASE64.encode([0u8; 16]);
231        let error = NodePublicKey::from_base64(&short).err().or_fail()?;
232
233        verify_that!(error.actual_length(), some(eq(16)))
234    }
235
236    #[gtest]
237    fn from_base64_empty_has_zero_decoded_length() -> Result<()> {
238        let error = NodePublicKey::from_base64("").err().or_fail()?;
239
240        verify_that!(error.actual_length(), some(eq(0)))
241    }
242
243    #[gtest]
244    fn from_base64_wrong_length_long() -> Result<()> {
245        let long = BASE64.encode([0u8; 33]);
246        let error = NodePublicKey::from_base64(&long).err().or_fail()?;
247
248        verify_that!(error.actual_length(), some(eq(33)))
249    }
250
251    #[gtest]
252    fn round_trip_and_accessors() -> Result<()> {
253        let key = NodePublicKey::from_bytes([7u8; PUBLIC_KEY_LENGTH]);
254
255        verify_that!(key.as_bytes(), eq(&[7u8; PUBLIC_KEY_LENGTH]))?;
256        let decoded = NodePublicKey::from_base64(&key.to_base64())?;
257
258        verify_that!(decoded.as_bytes(), eq(&[7u8; PUBLIC_KEY_LENGTH]))
259    }
260
261    #[gtest]
262    fn from_str_matches_from_base64() -> Result<()> {
263        let valid = zero_key_b64();
264
265        verify_that!(
266            NodePublicKey::from_str(&valid).ok(),
267            some(eq(&NodePublicKey::from_bytes([0u8; PUBLIC_KEY_LENGTH])))
268        )?;
269
270        verify_that!(NodePublicKey::from_str("!!!"), err(anything()))
271    }
272
273    #[gtest]
274    fn display_is_base64() -> Result<()> {
275        let key = NodePublicKey::from_bytes([0u8; PUBLIC_KEY_LENGTH]);
276        let expected = "A".repeat(43);
277
278        verify_that!(format!("{key}"), eq(expected.as_str()))
279    }
280
281    #[gtest]
282    fn debug_wraps_base64() -> Result<()> {
283        let key = NodePublicKey::from_bytes([0u8; PUBLIC_KEY_LENGTH]);
284        let debug = format!("{key:?}");
285
286        verify_that!(debug, contains_substring("NodePublicKey("))?;
287
288        verify_that!(debug, contains_substring("A".repeat(43).as_str()))
289    }
290
291    #[gtest]
292    fn serde_round_trips() -> Result<()> {
293        let key = NodePublicKey::from_bytes([7u8; PUBLIC_KEY_LENGTH]);
294        let json = serde_json::to_string(&key).or_fail()?;
295
296        verify_that!(json, eq(format!("\"{}\"", key.to_base64()).as_str()))?;
297        let back: NodePublicKey = serde_json::from_str(&json).or_fail()?;
298
299        verify_that!(back, eq(&key))
300    }
301
302    #[gtest]
303    fn serde_rejects_invalid() -> Result<()> {
304        verify_that!(
305            serde_json::from_str::<NodePublicKey>("\"!!!\""),
306            err(anything())
307        )
308    }
309
310    #[gtest]
311    fn invalid_error_display() -> Result<()> {
312        let error = NodePublicKey::from_base64("!!!").err().or_fail()?;
313
314        verify_that!(
315            error.to_string(),
316            eq("invalid node public key: expected 32 bytes base64")
317        )
318    }
319}