use axum::Json;
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use sqlx::{Pool, Sqlite};
use time::OffsetDateTime;
use time::format_description::FormatItem;
use time::macros::format_description;
const ACCEPTED_REASON_TYPES: &[&str] = &[
"com.atproto.moderation.defs#reasonSpam",
"com.atproto.moderation.defs#reasonViolation",
"com.atproto.moderation.defs#reasonMisleading",
"com.atproto.moderation.defs#reasonSexual",
"com.atproto.moderation.defs#reasonRude",
"com.atproto.moderation.defs#reasonOther",
];
const REASON_MAX_BYTES: usize = 2048;
const CTS_FORMAT: &[FormatItem<'_>] =
format_description!("[year]-[month]-[day]T[hour]:[minute]:[second].[subsecond digits:3]");
#[derive(Debug, Deserialize)]
pub struct CreateReportRequest {
#[serde(rename = "reasonType")]
pub reason_type: String,
#[serde(default)]
pub reason: Option<String>,
pub subject: ReportSubject,
#[serde(rename = "reportedBy")]
pub reported_by: String,
}
#[derive(Debug, Deserialize)]
#[serde(tag = "$type")]
pub enum ReportSubject {
#[serde(rename = "com.atproto.admin.defs#repoRef")]
RepoRef {
did: String,
},
#[serde(rename = "com.atproto.repo.strongRef")]
StrongRef {
uri: String,
cid: String,
},
}
#[derive(Debug, Serialize)]
pub struct ReportView {
pub id: i64,
#[serde(rename = "reasonType")]
pub reason_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub reason: Option<String>,
pub subject: Value,
#[serde(rename = "reportedBy")]
pub reported_by: String,
#[serde(rename = "createdAt")]
pub created_at: String,
}
pub async fn dispatch_pds_forwarded_report(pool: &Pool<Sqlite>, body: &[u8]) -> Response {
let req: CreateReportRequest = match serde_json::from_slice(body) {
Ok(r) => r,
Err(e) => return invalid_request(format!("malformed request body: {e}")),
};
if !ACCEPTED_REASON_TYPES.contains(&req.reason_type.as_str()) {
return invalid_request(format!("unsupported reasonType: {}", req.reason_type));
}
if let Some(reason) = &req.reason
&& reason.len() > REASON_MAX_BYTES
{
return invalid_request("reason exceeds maximum length".to_string());
}
if !req.reported_by.starts_with("did:") {
return invalid_request(format!("reportedBy {:?} is not a DID", req.reported_by));
}
let translated = match translate_subject(&req.subject) {
Ok(t) => t,
Err(msg) => return invalid_request(msg),
};
let created_at = match rfc3339_now() {
Ok(s) => s,
Err(()) => {
tracing::error!("xrpc_gateway createReport: rfc3339 formatting failed");
return internal_server_error();
}
};
let id = match insert_report(pool, &req, &translated, &created_at).await {
Ok(id) => id,
Err(e) => {
tracing::error!(
error = %e,
"xrpc_gateway createReport: reports INSERT failed"
);
return internal_server_error();
}
};
let view = ReportView {
id,
reason_type: req.reason_type,
reason: req.reason,
subject: subject_to_json(&req.subject),
reported_by: req.reported_by,
created_at,
};
(StatusCode::OK, Json(view)).into_response()
}
#[derive(Debug)]
struct TranslatedSubject {
subject_type: &'static str,
subject_did: String,
subject_uri: Option<String>,
subject_cid: Option<String>,
}
fn translate_subject(subject: &ReportSubject) -> Result<TranslatedSubject, String> {
match subject {
ReportSubject::RepoRef { did } => {
if !did.starts_with("did:") {
return Err(format!("subject.did {did:?} is not a DID"));
}
Ok(TranslatedSubject {
subject_type: "account",
subject_did: did.clone(),
subject_uri: None,
subject_cid: None,
})
}
ReportSubject::StrongRef { uri, cid } => {
if !uri.starts_with("at://") {
return Err(format!("subject.uri {uri:?} is not an AT-URI"));
}
if cid.is_empty() {
return Err("subject.cid must be non-empty".to_string());
}
let did = extract_did_from_at_uri(uri)
.ok_or_else(|| format!("subject.uri {uri:?} missing DID authority"))?
.to_string();
Ok(TranslatedSubject {
subject_type: "record",
subject_did: did,
subject_uri: Some(uri.clone()),
subject_cid: Some(cid.clone()),
})
}
}
}
async fn insert_report(
pool: &Pool<Sqlite>,
req: &CreateReportRequest,
translated: &TranslatedSubject,
created_at: &str,
) -> sqlx::Result<i64> {
let subject_type = translated.subject_type;
sqlx::query_scalar!(
r#"INSERT INTO reports (
created_at, reported_by, reason_type, reason,
subject_type, subject_did, subject_uri, subject_cid, status
)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, 'pending')
RETURNING id as "id!: i64""#,
created_at,
req.reported_by,
req.reason_type,
req.reason,
subject_type,
translated.subject_did,
translated.subject_uri,
translated.subject_cid,
)
.fetch_one(pool)
.await
}
fn extract_did_from_at_uri(uri: &str) -> Option<&str> {
uri.strip_prefix("at://")
.and_then(|rest| rest.split('/').next())
.filter(|did| did.starts_with("did:"))
}
fn rfc3339_now() -> Result<String, ()> {
let dt = OffsetDateTime::now_utc();
let formatted = dt.format(&CTS_FORMAT).map_err(|_| ())?;
Ok(format!("{formatted}Z"))
}
fn subject_to_json(s: &ReportSubject) -> Value {
match s {
ReportSubject::RepoRef { did } => serde_json::json!({
"$type": "com.atproto.admin.defs#repoRef",
"did": did,
}),
ReportSubject::StrongRef { uri, cid } => serde_json::json!({
"$type": "com.atproto.repo.strongRef",
"uri": uri,
"cid": cid,
}),
}
}
fn invalid_request(message: String) -> Response {
let body = ErrorEnvelope {
error: "InvalidRequest",
message,
};
(StatusCode::BAD_REQUEST, Json(body)).into_response()
}
fn internal_server_error() -> Response {
let body = ErrorEnvelope {
error: "InternalServerError",
message: "service temporarily unavailable".to_string(),
};
(StatusCode::INTERNAL_SERVER_ERROR, Json(body)).into_response()
}
#[derive(Serialize)]
struct ErrorEnvelope {
error: &'static str,
message: String,
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn accepted_reason_types_match_user_direct_intake() {
assert_eq!(ACCEPTED_REASON_TYPES.len(), 6);
for t in [
"com.atproto.moderation.defs#reasonSpam",
"com.atproto.moderation.defs#reasonViolation",
"com.atproto.moderation.defs#reasonMisleading",
"com.atproto.moderation.defs#reasonSexual",
"com.atproto.moderation.defs#reasonRude",
"com.atproto.moderation.defs#reasonOther",
] {
assert!(ACCEPTED_REASON_TYPES.contains(&t), "missing {t}");
}
}
#[test]
fn translate_repo_ref_subject() {
let s = ReportSubject::RepoRef {
did: "did:plc:account".into(),
};
let t = translate_subject(&s).unwrap();
assert_eq!(t.subject_type, "account");
assert_eq!(t.subject_did, "did:plc:account");
assert!(t.subject_uri.is_none());
assert!(t.subject_cid.is_none());
}
#[test]
fn translate_strong_ref_extracts_did_authority() {
let s = ReportSubject::StrongRef {
uri: "at://did:plc:author/app.bsky.feed.post/abc".into(),
cid: "bafy123".into(),
};
let t = translate_subject(&s).unwrap();
assert_eq!(t.subject_type, "record");
assert_eq!(t.subject_did, "did:plc:author");
assert_eq!(
t.subject_uri.as_deref(),
Some("at://did:plc:author/app.bsky.feed.post/abc")
);
assert_eq!(t.subject_cid.as_deref(), Some("bafy123"));
}
#[test]
fn translate_repo_ref_rejects_non_did_authority() {
let s = ReportSubject::RepoRef {
did: "not-a-did".into(),
};
let err = translate_subject(&s).unwrap_err();
assert!(err.contains("not a DID"), "{err}");
}
#[test]
fn translate_strong_ref_rejects_https_uri() {
let s = ReportSubject::StrongRef {
uri: "https://example.com/post".into(),
cid: "bafy".into(),
};
let err = translate_subject(&s).unwrap_err();
assert!(err.contains("AT-URI"), "{err}");
}
#[test]
fn translate_strong_ref_rejects_empty_cid() {
let s = ReportSubject::StrongRef {
uri: "at://did:plc:a/c/r".into(),
cid: "".into(),
};
let err = translate_subject(&s).unwrap_err();
assert!(err.contains("cid"), "{err}");
}
#[test]
fn translate_strong_ref_rejects_uri_without_did_authority() {
let s = ReportSubject::StrongRef {
uri: "at://example.com/c/r".into(),
cid: "bafy".into(),
};
let err = translate_subject(&s).unwrap_err();
assert!(err.contains("DID authority"), "{err}");
}
#[test]
fn deserialize_repo_ref_subject() {
let v = json!({
"$type": "com.atproto.admin.defs#repoRef",
"did": "did:plc:abc"
});
let s: ReportSubject = serde_json::from_value(v).unwrap();
assert!(matches!(s, ReportSubject::RepoRef { did } if did == "did:plc:abc"));
}
#[test]
fn deserialize_strong_ref_subject() {
let v = json!({
"$type": "com.atproto.repo.strongRef",
"uri": "at://did:plc:abc/c/r",
"cid": "bafy"
});
let s: ReportSubject = serde_json::from_value(v).unwrap();
match s {
ReportSubject::StrongRef { uri, cid } => {
assert_eq!(uri, "at://did:plc:abc/c/r");
assert_eq!(cid, "bafy");
}
other => panic!("expected StrongRef, got {other:?}"),
}
}
#[test]
fn deserialize_full_request_round_trip() {
let body = json!({
"reasonType": "com.atproto.moderation.defs#reasonSpam",
"reason": "looks like spam",
"subject": {
"$type": "com.atproto.admin.defs#repoRef",
"did": "did:plc:target",
},
"reportedBy": "did:plc:reporter",
});
let req: CreateReportRequest = serde_json::from_value(body).unwrap();
assert_eq!(req.reason_type, "com.atproto.moderation.defs#reasonSpam");
assert_eq!(req.reason.as_deref(), Some("looks like spam"));
assert_eq!(req.reported_by, "did:plc:reporter");
}
#[test]
fn extract_did_from_at_uri_basic() {
assert_eq!(
extract_did_from_at_uri("at://did:plc:abc/c/r"),
Some("did:plc:abc")
);
assert_eq!(extract_did_from_at_uri("https://x"), None);
assert_eq!(extract_did_from_at_uri("at://example.com/c/r"), None);
}
}