1use axum::Extension;
25use axum::Json;
26use axum::Router;
27use axum::http::{StatusCode, header};
28use axum::response::{IntoResponse, Response};
29use axum::routing::get;
30use serde::Serialize;
31use serde_json::json;
32use sqlx::{Pool, Sqlite};
33
34use crate::config::Config;
35
36pub const LABELER_KEY_FRAGMENT: &str = "#atproto_label";
39
40const LABELER_SERVICE_FRAGMENT: &str = "#atproto_labeler";
46const LABELER_SERVICE_TYPE: &str = "AtprotoLabeler";
47
48#[derive(Debug, Clone)]
52pub struct SigningKeyRow {
53 pub id: i64,
56 pub public_key_multibase: String,
59}
60
61#[derive(Debug, Serialize)]
65pub struct DidDocumentWire {
66 #[serde(rename = "@context")]
69 pub context: Vec<&'static str>,
70 pub id: String,
73 #[serde(rename = "verificationMethod")]
77 pub verification_method: Vec<VerificationMethodWire>,
78 pub service: Vec<ServiceWire>,
81}
82
83#[derive(Debug, Serialize)]
87pub struct VerificationMethodWire {
88 pub id: String,
91 #[serde(rename = "type")]
93 pub r#type: &'static str,
94 pub controller: String,
97 #[serde(rename = "publicKeyMultibase")]
99 pub public_key_multibase: String,
100}
101
102#[derive(Debug, Serialize)]
106pub struct ServiceWire {
107 pub id: &'static str,
109 #[serde(rename = "type")]
112 pub r#type: &'static str,
113 #[serde(rename = "serviceEndpoint")]
116 pub service_endpoint: String,
117}
118
119pub fn build_did_document(
128 service_did: &str,
129 service_endpoint: &str,
130 keys: &[SigningKeyRow],
131) -> DidDocumentWire {
132 let single = keys.len() == 1;
133 let verification_method = keys
134 .iter()
135 .map(|k| {
136 let fragment = if single {
137 LABELER_KEY_FRAGMENT.to_string()
138 } else {
139 format!("{LABELER_KEY_FRAGMENT}_{}", k.id)
140 };
141 VerificationMethodWire {
142 id: format!("{service_did}{fragment}"),
143 r#type: "Multikey",
144 controller: service_did.to_string(),
145 public_key_multibase: k.public_key_multibase.clone(),
146 }
147 })
148 .collect();
149 DidDocumentWire {
150 context: vec![
151 "https://www.w3.org/ns/did/v1",
152 "https://w3id.org/security/multikey/v1",
153 ],
154 id: service_did.to_string(),
155 verification_method,
156 service: vec![ServiceWire {
157 id: LABELER_SERVICE_FRAGMENT,
158 r#type: LABELER_SERVICE_TYPE,
159 service_endpoint: service_endpoint.to_string(),
160 }],
161 }
162}
163
164pub fn did_document_router(pool: Pool<Sqlite>, config: Config) -> Router {
167 Router::new()
168 .route("/.well-known/did.json", get(serve_did_document))
169 .layer(Extension(DidDocumentState { pool, config }))
170}
171
172#[derive(Clone)]
173struct DidDocumentState {
174 pool: Pool<Sqlite>,
175 config: Config,
176}
177
178async fn serve_did_document(Extension(state): Extension<DidDocumentState>) -> Response {
179 let rows = match sqlx::query_as!(
180 SigningKeyRow,
181 "SELECT id, public_key_multibase FROM signing_keys WHERE valid_to IS NULL ORDER BY id"
182 )
183 .fetch_all(&state.pool)
184 .await
185 {
186 Ok(r) => r,
187 Err(_) => return internal_error(),
188 };
189 if rows.is_empty() {
190 return service_unavailable();
191 }
192 let doc = build_did_document(
193 &state.config.service_did,
194 &state.config.service_endpoint,
195 &rows,
196 );
197 let body = serde_json::to_vec(&doc).expect("DidDocumentWire always serializes");
198 (
199 StatusCode::OK,
200 [(header::CONTENT_TYPE, "application/json")],
201 body,
202 )
203 .into_response()
204}
205
206fn service_unavailable() -> Response {
207 (
208 StatusCode::SERVICE_UNAVAILABLE,
209 Json(json!({
210 "error": "ServiceUnavailable",
211 "message": "signing key not bootstrapped",
212 })),
213 )
214 .into_response()
215}
216
217fn internal_error() -> Response {
218 (
219 StatusCode::INTERNAL_SERVER_ERROR,
220 Json(json!({
221 "error": "InternalServerError",
222 "message": "service temporarily unavailable",
223 })),
224 )
225 .into_response()
226}
227
228#[cfg(test)]
229mod tests {
230 use super::*;
231 use serde_json::Value;
232
233 fn row(id: i64, mb: &str) -> SigningKeyRow {
234 SigningKeyRow {
235 id,
236 public_key_multibase: mb.into(),
237 }
238 }
239
240 #[test]
241 fn single_key_uses_unsuffixed_fragment() {
242 let doc = build_did_document(
243 "did:web:labeler.example",
244 "https://labeler.example",
245 &[row(1, "zXYZ")],
246 );
247 assert_eq!(doc.verification_method.len(), 1);
248 assert_eq!(
249 doc.verification_method[0].id,
250 "did:web:labeler.example#atproto_label"
251 );
252 assert_eq!(doc.verification_method[0].r#type, "Multikey");
253 assert_eq!(
254 doc.verification_method[0].controller,
255 "did:web:labeler.example"
256 );
257 assert_eq!(doc.verification_method[0].public_key_multibase, "zXYZ");
258 }
259
260 #[test]
261 fn multi_key_uses_id_suffixed_fragments() {
262 let doc = build_did_document(
264 "did:web:labeler.example",
265 "https://labeler.example",
266 &[row(1, "zOLD"), row(2, "zNEW")],
267 );
268 assert_eq!(doc.verification_method.len(), 2);
269 assert_eq!(
270 doc.verification_method[0].id,
271 "did:web:labeler.example#atproto_label_1"
272 );
273 assert_eq!(
274 doc.verification_method[1].id,
275 "did:web:labeler.example#atproto_label_2"
276 );
277 }
278
279 #[test]
280 fn context_contains_required_entries() {
281 let doc = build_did_document(
282 "did:web:labeler.example",
283 "https://labeler.example",
284 &[row(1, "z")],
285 );
286 assert!(doc.context.contains(&"https://www.w3.org/ns/did/v1"));
287 assert!(
288 doc.context
289 .contains(&"https://w3id.org/security/multikey/v1")
290 );
291 }
292
293 #[test]
294 fn service_entry_is_atproto_labeler() {
295 let doc = build_did_document(
296 "did:web:labeler.example",
297 "https://labeler.example",
298 &[row(1, "z")],
299 );
300 assert_eq!(doc.service.len(), 1);
301 assert_eq!(doc.service[0].id, "#atproto_labeler");
302 assert_eq!(doc.service[0].r#type, "AtprotoLabeler");
303 assert_eq!(doc.service[0].service_endpoint, "https://labeler.example");
304 }
305
306 #[test]
307 fn serialized_json_has_camelcase_field_names() {
308 let doc = build_did_document(
309 "did:web:labeler.example",
310 "https://labeler.example",
311 &[row(1, "z")],
312 );
313 let v: Value = serde_json::to_value(&doc).unwrap();
314 assert!(v.get("@context").is_some());
316 assert!(
317 v["verificationMethod"][0]
318 .get("publicKeyMultibase")
319 .is_some()
320 );
321 assert!(v["service"][0].get("serviceEndpoint").is_some());
322 }
323}