1use std::path::PathBuf;
37use std::sync::Arc;
38use std::time::Duration;
39
40use axum::Extension;
41use axum::Json;
42use axum::Router;
43use axum::body::Bytes;
44use axum::http::{HeaderMap, HeaderValue, StatusCode};
45use axum::response::{IntoResponse, Response};
46use axum::routing::post;
47use serde::{Deserialize, Serialize};
48use sqlx::{Pool, Sqlite};
49use time::OffsetDateTime;
50use time::format_description::FormatItem;
51use time::macros::format_description;
52
53use crate::auth::AuthContext;
54
55#[derive(Debug, Clone)]
58pub struct CreateReportConfig {
59 pub per_did_limit: u32,
61 pub per_did_window: Duration,
63 pub global_pending_cap: u32,
65 pub disk_size_limit_bytes: u64,
68 pub db_path: PathBuf,
71 pub cors_allowed_origins: Vec<String>,
75 pub max_body_bytes: usize,
79}
80
81impl Default for CreateReportConfig {
82 fn default() -> Self {
83 Self {
84 per_did_limit: 10,
85 per_did_window: Duration::from_secs(3600),
86 global_pending_cap: 10_000,
87 disk_size_limit_bytes: 5 * 1024 * 1024 * 1024,
88 db_path: PathBuf::new(),
89 cors_allowed_origins: Vec::new(),
90 max_body_bytes: 32 * 1024,
91 }
92 }
93}
94
95const REASON_MAX_BYTES: usize = 2048;
100
101const ACCEPTED_REASON_TYPES: &[&str] = &[
108 "com.atproto.moderation.defs#reasonSpam",
109 "com.atproto.moderation.defs#reasonViolation",
110 "com.atproto.moderation.defs#reasonMisleading",
111 "com.atproto.moderation.defs#reasonSexual",
112 "com.atproto.moderation.defs#reasonRude",
113 "com.atproto.moderation.defs#reasonOther",
114];
115
116const CTS_FORMAT: &[FormatItem<'_>] =
120 format_description!("[year]-[month]-[day]T[hour]:[minute]:[second].[subsecond digits:3]");
121
122#[derive(Debug, Deserialize)]
125pub(super) struct CreateReportInput {
126 #[serde(rename = "reasonType")]
127 reason_type: String,
128 subject: Subject,
129 reason: Option<String>,
130 #[serde(rename = "modTool", default)]
135 #[allow(dead_code)]
136 mod_tool: Option<serde_json::Value>,
137}
138
139#[derive(Debug, Deserialize)]
140#[serde(tag = "$type")]
141pub(super) enum Subject {
142 #[serde(rename = "com.atproto.admin.defs#repoRef")]
143 Repo { did: String },
144 #[serde(rename = "com.atproto.repo.strongRef")]
145 Strong { uri: String, cid: String },
146}
147
148#[derive(Debug, Serialize)]
149pub(super) struct CreateReportOutput {
150 id: i64,
151 #[serde(rename = "createdAt")]
152 created_at: String,
153 #[serde(rename = "reasonType")]
154 reason_type: String,
155 subject: serde_json::Value,
156 #[serde(rename = "reportedBy")]
157 reported_by: String,
158 #[serde(skip_serializing_if = "Option::is_none")]
159 reason: Option<String>,
160}
161
162#[derive(Debug, Serialize)]
163struct XrpcErrorBody {
164 error: &'static str,
165 message: &'static str,
166}
167
168const RATE_LIMIT_BODY: XrpcErrorBody = XrpcErrorBody {
171 error: "RateLimitExceeded",
172 message: "rate limit exceeded",
173};
174
175const REASON_TOO_LONG_BODY: XrpcErrorBody = XrpcErrorBody {
180 error: "InvalidRequest",
181 message: "reason exceeds maximum length",
182};
183
184#[derive(Clone)]
190pub(super) struct AppState {
191 pub pool: Pool<Sqlite>,
192 pub auth: Arc<AuthContext>,
193 pub config: Arc<CreateReportConfig>,
194}
195
196pub fn create_report_router(
199 pool: Pool<Sqlite>,
200 auth: Arc<AuthContext>,
201 config: CreateReportConfig,
202) -> Router {
203 let state = AppState {
204 pool,
205 auth,
206 config: Arc::new(config),
207 };
208 Router::new()
209 .route(
210 "/xrpc/com.atproto.moderation.createReport",
211 post(post_handler).options(preflight_handler),
212 )
213 .layer(Extension(state))
214}
215
216async fn post_handler(
219 Extension(state): Extension<AppState>,
220 headers: HeaderMap,
221 body: Bytes,
222) -> Response {
223 if let Some(reject) = check_cors(&state.config, &headers) {
225 return reject;
226 }
227
228 if body.len() > state.config.max_body_bytes {
230 return xrpc_error(
231 StatusCode::PAYLOAD_TOO_LARGE,
232 "InvalidRequest",
233 "request body too large",
234 );
235 }
236
237 let token = match headers
239 .get("authorization")
240 .and_then(|h| h.to_str().ok())
241 .and_then(|s| s.strip_prefix("Bearer "))
242 {
243 Some(t) => t,
244 None => return auth_required(),
245 };
246 let caller = match state
247 .auth
248 .verify_service_auth(token, "com.atproto.moderation.createReport")
249 .await
250 {
251 Ok(v) => v,
252 Err(_) => return auth_required(),
253 };
254
255 let trusted = crate::xrpc_gateway::is_trusted_pds(&state.pool, &caller.iss)
269 .await
270 .unwrap_or(false);
271 if trusted {
272 return crate::xrpc_gateway::handlers::create_report::dispatch_pds_forwarded_report(
273 &state.pool,
274 &body,
275 )
276 .await;
277 }
278
279 let input: CreateReportInput = match serde_json::from_slice(&body) {
281 Ok(v) => v,
282 Err(_) => {
283 return xrpc_error(
284 StatusCode::BAD_REQUEST,
285 "InvalidRequest",
286 "malformed request body",
287 );
288 }
289 };
290
291 if !ACCEPTED_REASON_TYPES.contains(&input.reason_type.as_str()) {
293 return xrpc_error(
294 StatusCode::BAD_REQUEST,
295 "InvalidRequest",
296 "unsupported reasonType",
297 );
298 }
299 if let Some(reason) = &input.reason
300 && reason.len() > REASON_MAX_BYTES
301 {
302 return (StatusCode::BAD_REQUEST, Json(REASON_TOO_LONG_BODY)).into_response();
305 }
306 let (subject_type, subject_did, subject_uri, subject_cid) = match &input.subject {
307 Subject::Repo { did } => {
308 if !did.starts_with("did:") {
309 return xrpc_error(
310 StatusCode::BAD_REQUEST,
311 "InvalidRequest",
312 "subject.did malformed",
313 );
314 }
315 ("account", did.clone(), None, None)
316 }
317 Subject::Strong { uri, cid } => {
318 if !uri.starts_with("at://") || cid.is_empty() {
319 return xrpc_error(
320 StatusCode::BAD_REQUEST,
321 "InvalidRequest",
322 "subject malformed",
323 );
324 }
325 let did = match extract_did_from_at_uri(uri) {
326 Some(d) => d.to_owned(),
327 None => {
328 return xrpc_error(
329 StatusCode::BAD_REQUEST,
330 "InvalidRequest",
331 "subject.uri missing DID authority",
332 );
333 }
334 };
335 ("record", did, Some(uri.clone()), Some(cid.clone()))
336 }
337 };
338
339 if let Ok(meta) = std::fs::metadata(&state.config.db_path)
341 && meta.len() >= state.config.disk_size_limit_bytes
342 {
343 return xrpc_error(
344 StatusCode::INTERNAL_SERVER_ERROR,
345 "InternalServerError",
346 "service temporarily unavailable",
347 );
348 }
349
350 let suppressed: i64 = sqlx::query_scalar!(
352 "SELECT COUNT(*) FROM suppressed_reporters WHERE did = ?1",
353 caller.iss
354 )
355 .fetch_one(&state.pool)
356 .await
357 .unwrap_or(0);
358 if suppressed > 0 {
359 return rate_limited();
360 }
361
362 let window_start = match rfc3339_minus_secs(state.config.per_did_window.as_secs() as i64) {
365 Ok(s) => s,
366 Err(()) => {
367 return xrpc_error(
368 StatusCode::INTERNAL_SERVER_ERROR,
369 "InternalServerError",
370 "service temporarily unavailable",
371 );
372 }
373 };
374 let recent: i64 = sqlx::query_scalar!(
375 "SELECT COUNT(*) FROM reports WHERE reported_by = ?1 AND created_at >= ?2",
376 caller.iss,
377 window_start
378 )
379 .fetch_one(&state.pool)
380 .await
381 .unwrap_or(0);
382 if recent >= state.config.per_did_limit as i64 {
383 return rate_limited();
384 }
385
386 let pending: i64 = sqlx::query_scalar!("SELECT COUNT(*) FROM reports WHERE status = 'pending'")
388 .fetch_one(&state.pool)
389 .await
390 .unwrap_or(0);
391 if pending >= state.config.global_pending_cap as i64 {
392 return rate_limited();
393 }
394
395 let created_at = match rfc3339_now() {
397 Ok(s) => s,
398 Err(()) => {
399 return xrpc_error(
400 StatusCode::INTERNAL_SERVER_ERROR,
401 "InternalServerError",
402 "service temporarily unavailable",
403 );
404 }
405 };
406 let insert_result = sqlx::query_scalar!(
407 "INSERT INTO reports (
408 created_at, reported_by, reason_type, reason,
409 subject_type, subject_did, subject_uri, subject_cid, status
410 )
411 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, 'pending')
412 RETURNING id",
413 created_at,
414 caller.iss,
415 input.reason_type,
416 input.reason,
417 subject_type,
418 subject_did,
419 subject_uri,
420 subject_cid,
421 )
422 .fetch_one(&state.pool)
423 .await;
424 let id = match insert_result {
425 Ok(id) => id,
426 Err(_) => {
427 return xrpc_error(
428 StatusCode::INTERNAL_SERVER_ERROR,
429 "InternalServerError",
430 "service temporarily unavailable",
431 );
432 }
433 };
434
435 let subject_json = subject_to_json(&input.subject);
436 let out = CreateReportOutput {
437 id,
438 created_at,
439 reason_type: input.reason_type,
440 subject: subject_json,
441 reported_by: caller.iss,
442 reason: input.reason,
443 };
444
445 let mut resp = (StatusCode::OK, Json(out)).into_response();
446 attach_cors_headers(&mut resp, &state.config, &headers);
447 resp
448}
449
450async fn preflight_handler(Extension(state): Extension<AppState>, headers: HeaderMap) -> Response {
454 let Some(origin) = headers.get("origin").and_then(|v| v.to_str().ok()) else {
455 return StatusCode::METHOD_NOT_ALLOWED.into_response();
457 };
458 if !state
459 .config
460 .cors_allowed_origins
461 .iter()
462 .any(|a| a == origin)
463 {
464 return StatusCode::FORBIDDEN.into_response();
465 }
466
467 let mut resp = StatusCode::NO_CONTENT.into_response();
468 let h = resp.headers_mut();
469 h.insert(
470 "access-control-allow-origin",
471 HeaderValue::from_str(origin).unwrap_or(HeaderValue::from_static("")),
472 );
473 h.insert(
474 "access-control-allow-methods",
475 HeaderValue::from_static("POST"),
476 );
477 h.insert(
478 "access-control-allow-headers",
479 HeaderValue::from_static("content-type, authorization"),
480 );
481 h.insert("vary", HeaderValue::from_static("Origin"));
482 resp
483}
484
485fn check_cors(config: &CreateReportConfig, headers: &HeaderMap) -> Option<Response> {
488 let origin = headers.get("origin").and_then(|v| v.to_str().ok())?;
489 if config.cors_allowed_origins.iter().any(|a| a == origin) {
490 None
491 } else {
492 Some(StatusCode::FORBIDDEN.into_response())
493 }
494}
495
496fn attach_cors_headers(resp: &mut Response, config: &CreateReportConfig, req_headers: &HeaderMap) {
497 let Some(origin) = req_headers.get("origin").and_then(|v| v.to_str().ok()) else {
498 return;
499 };
500 if !config.cors_allowed_origins.iter().any(|a| a == origin) {
501 return;
502 }
503 let h = resp.headers_mut();
504 if let Ok(v) = HeaderValue::from_str(origin) {
505 h.insert("access-control-allow-origin", v);
506 }
507 h.insert("vary", HeaderValue::from_static("Origin"));
508}
509
510fn auth_required() -> Response {
511 xrpc_error(
512 StatusCode::UNAUTHORIZED,
513 "AuthenticationRequired",
514 "authentication required",
515 )
516}
517
518fn rate_limited() -> Response {
519 (StatusCode::TOO_MANY_REQUESTS, Json(RATE_LIMIT_BODY)).into_response()
520}
521
522fn xrpc_error(status: StatusCode, error: &'static str, message: &'static str) -> Response {
523 (status, Json(XrpcErrorBody { error, message })).into_response()
524}
525
526fn extract_did_from_at_uri(uri: &str) -> Option<&str> {
527 uri.strip_prefix("at://")
528 .and_then(|rest| rest.split('/').next())
529 .filter(|did| did.starts_with("did:"))
530}
531
532fn rfc3339_now() -> Result<String, ()> {
533 let dt = OffsetDateTime::now_utc();
534 let formatted = dt.format(&CTS_FORMAT).map_err(|_| ())?;
535 Ok(format!("{formatted}Z"))
536}
537
538fn rfc3339_minus_secs(secs: i64) -> Result<String, ()> {
539 let dt = OffsetDateTime::now_utc() - Duration::from_secs(secs.max(0) as u64);
540 let formatted = dt.format(&CTS_FORMAT).map_err(|_| ())?;
541 Ok(format!("{formatted}Z"))
542}
543
544fn subject_to_json(s: &Subject) -> serde_json::Value {
545 match s {
546 Subject::Repo { did } => serde_json::json!({
547 "$type": "com.atproto.admin.defs#repoRef",
548 "did": did,
549 }),
550 Subject::Strong { uri, cid } => serde_json::json!({
551 "$type": "com.atproto.repo.strongRef",
552 "uri": uri,
553 "cid": cid,
554 }),
555 }
556}
557
558#[cfg(test)]
559mod tests {
560 use super::*;
561
562 #[test]
563 fn extract_did_accepts_valid_at_uri() {
564 let did = extract_did_from_at_uri("at://did:plc:abc/col/rkey").unwrap();
565 assert_eq!(did, "did:plc:abc");
566 }
567
568 #[test]
569 fn extract_did_rejects_non_at_uri() {
570 assert!(extract_did_from_at_uri("https://example.com").is_none());
571 }
572
573 #[test]
574 fn extract_did_rejects_non_did_authority() {
575 assert!(extract_did_from_at_uri("at://example.com/foo").is_none());
578 }
579
580 #[test]
581 fn reason_types_allowlist_pins_f11_literal_set() {
582 assert_eq!(ACCEPTED_REASON_TYPES.len(), 6);
586 }
587}