wowlab_common/
node_public_key.rs1use 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#[derive(Clone, Deserialize, Eq, Hash, PartialEq)]
10#[serde(try_from = "String")]
11pub struct NodePublicKey([u8; PUBLIC_KEY_LENGTH]);
12
13#[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 #[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 #[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 #[must_use]
67 pub fn from_bytes(bytes: [u8; PUBLIC_KEY_LENGTH]) -> Self {
68 Self(bytes)
69 }
70
71 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 #[must_use]
85 pub fn as_bytes(&self) -> &[u8; PUBLIC_KEY_LENGTH] {
86 &self.0
87 }
88
89 #[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}