use axum::Extension;
use axum::Json;
use axum::Router;
use axum::http::{StatusCode, header};
use axum::response::{IntoResponse, Response};
use axum::routing::get;
use serde::Serialize;
use serde_json::json;
use sqlx::{Pool, Sqlite};
use crate::config::Config;
pub const LABELER_KEY_FRAGMENT: &str = "#atproto_label";
const LABELER_SERVICE_FRAGMENT: &str = "#atproto_labeler";
const LABELER_SERVICE_TYPE: &str = "AtprotoLabeler";
#[derive(Debug, Clone)]
pub struct SigningKeyRow {
pub id: i64,
pub public_key_multibase: String,
}
#[derive(Debug, Serialize)]
pub struct DidDocumentWire {
#[serde(rename = "@context")]
pub context: Vec<&'static str>,
pub id: String,
#[serde(rename = "verificationMethod")]
pub verification_method: Vec<VerificationMethodWire>,
pub service: Vec<ServiceWire>,
}
#[derive(Debug, Serialize)]
pub struct VerificationMethodWire {
pub id: String,
#[serde(rename = "type")]
pub r#type: &'static str,
pub controller: String,
#[serde(rename = "publicKeyMultibase")]
pub public_key_multibase: String,
}
#[derive(Debug, Serialize)]
pub struct ServiceWire {
pub id: &'static str,
#[serde(rename = "type")]
pub r#type: &'static str,
#[serde(rename = "serviceEndpoint")]
pub service_endpoint: String,
}
pub fn build_did_document(
service_did: &str,
service_endpoint: &str,
keys: &[SigningKeyRow],
) -> DidDocumentWire {
let single = keys.len() == 1;
let verification_method = keys
.iter()
.map(|k| {
let fragment = if single {
LABELER_KEY_FRAGMENT.to_string()
} else {
format!("{LABELER_KEY_FRAGMENT}_{}", k.id)
};
VerificationMethodWire {
id: format!("{service_did}{fragment}"),
r#type: "Multikey",
controller: service_did.to_string(),
public_key_multibase: k.public_key_multibase.clone(),
}
})
.collect();
DidDocumentWire {
context: vec![
"https://www.w3.org/ns/did/v1",
"https://w3id.org/security/multikey/v1",
],
id: service_did.to_string(),
verification_method,
service: vec![ServiceWire {
id: LABELER_SERVICE_FRAGMENT,
r#type: LABELER_SERVICE_TYPE,
service_endpoint: service_endpoint.to_string(),
}],
}
}
pub fn did_document_router(pool: Pool<Sqlite>, config: Config) -> Router {
Router::new()
.route("/.well-known/did.json", get(serve_did_document))
.layer(Extension(DidDocumentState { pool, config }))
}
#[derive(Clone)]
struct DidDocumentState {
pool: Pool<Sqlite>,
config: Config,
}
async fn serve_did_document(Extension(state): Extension<DidDocumentState>) -> Response {
let rows = match sqlx::query_as!(
SigningKeyRow,
"SELECT id, public_key_multibase FROM signing_keys WHERE valid_to IS NULL ORDER BY id"
)
.fetch_all(&state.pool)
.await
{
Ok(r) => r,
Err(_) => return internal_error(),
};
if rows.is_empty() {
return service_unavailable();
}
let doc = build_did_document(
&state.config.service_did,
&state.config.service_endpoint,
&rows,
);
let body = serde_json::to_vec(&doc).expect("DidDocumentWire always serializes");
(
StatusCode::OK,
[(header::CONTENT_TYPE, "application/json")],
body,
)
.into_response()
}
fn service_unavailable() -> Response {
(
StatusCode::SERVICE_UNAVAILABLE,
Json(json!({
"error": "ServiceUnavailable",
"message": "signing key not bootstrapped",
})),
)
.into_response()
}
fn internal_error() -> Response {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({
"error": "InternalServerError",
"message": "service temporarily unavailable",
})),
)
.into_response()
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::Value;
fn row(id: i64, mb: &str) -> SigningKeyRow {
SigningKeyRow {
id,
public_key_multibase: mb.into(),
}
}
#[test]
fn single_key_uses_unsuffixed_fragment() {
let doc = build_did_document(
"did:web:labeler.example",
"https://labeler.example",
&[row(1, "zXYZ")],
);
assert_eq!(doc.verification_method.len(), 1);
assert_eq!(
doc.verification_method[0].id,
"did:web:labeler.example#atproto_label"
);
assert_eq!(doc.verification_method[0].r#type, "Multikey");
assert_eq!(
doc.verification_method[0].controller,
"did:web:labeler.example"
);
assert_eq!(doc.verification_method[0].public_key_multibase, "zXYZ");
}
#[test]
fn multi_key_uses_id_suffixed_fragments() {
let doc = build_did_document(
"did:web:labeler.example",
"https://labeler.example",
&[row(1, "zOLD"), row(2, "zNEW")],
);
assert_eq!(doc.verification_method.len(), 2);
assert_eq!(
doc.verification_method[0].id,
"did:web:labeler.example#atproto_label_1"
);
assert_eq!(
doc.verification_method[1].id,
"did:web:labeler.example#atproto_label_2"
);
}
#[test]
fn context_contains_required_entries() {
let doc = build_did_document(
"did:web:labeler.example",
"https://labeler.example",
&[row(1, "z")],
);
assert!(doc.context.contains(&"https://www.w3.org/ns/did/v1"));
assert!(
doc.context
.contains(&"https://w3id.org/security/multikey/v1")
);
}
#[test]
fn service_entry_is_atproto_labeler() {
let doc = build_did_document(
"did:web:labeler.example",
"https://labeler.example",
&[row(1, "z")],
);
assert_eq!(doc.service.len(), 1);
assert_eq!(doc.service[0].id, "#atproto_labeler");
assert_eq!(doc.service[0].r#type, "AtprotoLabeler");
assert_eq!(doc.service[0].service_endpoint, "https://labeler.example");
}
#[test]
fn serialized_json_has_camelcase_field_names() {
let doc = build_did_document(
"did:web:labeler.example",
"https://labeler.example",
&[row(1, "z")],
);
let v: Value = serde_json::to_value(&doc).unwrap();
assert!(v.get("@context").is_some());
assert!(
v["verificationMethod"][0]
.get("publicKeyMultibase")
.is_some()
);
assert!(v["service"][0].get("serviceEndpoint").is_some());
}
}