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 input: CreateReportInput = match serde_json::from_slice(&body) {
257 Ok(v) => v,
258 Err(_) => {
259 return xrpc_error(
260 StatusCode::BAD_REQUEST,
261 "InvalidRequest",
262 "malformed request body",
263 );
264 }
265 };
266
267 if !ACCEPTED_REASON_TYPES.contains(&input.reason_type.as_str()) {
269 return xrpc_error(
270 StatusCode::BAD_REQUEST,
271 "InvalidRequest",
272 "unsupported reasonType",
273 );
274 }
275 if let Some(reason) = &input.reason
276 && reason.len() > REASON_MAX_BYTES
277 {
278 return (StatusCode::BAD_REQUEST, Json(REASON_TOO_LONG_BODY)).into_response();
281 }
282 let (subject_type, subject_did, subject_uri, subject_cid) = match &input.subject {
283 Subject::Repo { did } => {
284 if !did.starts_with("did:") {
285 return xrpc_error(
286 StatusCode::BAD_REQUEST,
287 "InvalidRequest",
288 "subject.did malformed",
289 );
290 }
291 ("account", did.clone(), None, None)
292 }
293 Subject::Strong { uri, cid } => {
294 if !uri.starts_with("at://") || cid.is_empty() {
295 return xrpc_error(
296 StatusCode::BAD_REQUEST,
297 "InvalidRequest",
298 "subject malformed",
299 );
300 }
301 let did = match extract_did_from_at_uri(uri) {
302 Some(d) => d.to_owned(),
303 None => {
304 return xrpc_error(
305 StatusCode::BAD_REQUEST,
306 "InvalidRequest",
307 "subject.uri missing DID authority",
308 );
309 }
310 };
311 ("record", did, Some(uri.clone()), Some(cid.clone()))
312 }
313 };
314
315 if let Ok(meta) = std::fs::metadata(&state.config.db_path)
317 && meta.len() >= state.config.disk_size_limit_bytes
318 {
319 return xrpc_error(
320 StatusCode::INTERNAL_SERVER_ERROR,
321 "InternalServerError",
322 "service temporarily unavailable",
323 );
324 }
325
326 let suppressed: i64 = sqlx::query_scalar!(
328 "SELECT COUNT(*) FROM suppressed_reporters WHERE did = ?1",
329 caller.iss
330 )
331 .fetch_one(&state.pool)
332 .await
333 .unwrap_or(0);
334 if suppressed > 0 {
335 return rate_limited();
336 }
337
338 let window_start = match rfc3339_minus_secs(state.config.per_did_window.as_secs() as i64) {
341 Ok(s) => s,
342 Err(()) => {
343 return xrpc_error(
344 StatusCode::INTERNAL_SERVER_ERROR,
345 "InternalServerError",
346 "service temporarily unavailable",
347 );
348 }
349 };
350 let recent: i64 = sqlx::query_scalar!(
351 "SELECT COUNT(*) FROM reports WHERE reported_by = ?1 AND created_at >= ?2",
352 caller.iss,
353 window_start
354 )
355 .fetch_one(&state.pool)
356 .await
357 .unwrap_or(0);
358 if recent >= state.config.per_did_limit as i64 {
359 return rate_limited();
360 }
361
362 let pending: i64 = sqlx::query_scalar!("SELECT COUNT(*) FROM reports WHERE status = 'pending'")
364 .fetch_one(&state.pool)
365 .await
366 .unwrap_or(0);
367 if pending >= state.config.global_pending_cap as i64 {
368 return rate_limited();
369 }
370
371 let created_at = match rfc3339_now() {
373 Ok(s) => s,
374 Err(()) => {
375 return xrpc_error(
376 StatusCode::INTERNAL_SERVER_ERROR,
377 "InternalServerError",
378 "service temporarily unavailable",
379 );
380 }
381 };
382 let insert_result = sqlx::query_scalar!(
383 "INSERT INTO reports (
384 created_at, reported_by, reason_type, reason,
385 subject_type, subject_did, subject_uri, subject_cid, status
386 )
387 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, 'pending')
388 RETURNING id",
389 created_at,
390 caller.iss,
391 input.reason_type,
392 input.reason,
393 subject_type,
394 subject_did,
395 subject_uri,
396 subject_cid,
397 )
398 .fetch_one(&state.pool)
399 .await;
400 let id = match insert_result {
401 Ok(id) => id,
402 Err(_) => {
403 return xrpc_error(
404 StatusCode::INTERNAL_SERVER_ERROR,
405 "InternalServerError",
406 "service temporarily unavailable",
407 );
408 }
409 };
410
411 let subject_json = subject_to_json(&input.subject);
412 let out = CreateReportOutput {
413 id,
414 created_at,
415 reason_type: input.reason_type,
416 subject: subject_json,
417 reported_by: caller.iss,
418 reason: input.reason,
419 };
420
421 let mut resp = (StatusCode::OK, Json(out)).into_response();
422 attach_cors_headers(&mut resp, &state.config, &headers);
423 resp
424}
425
426async fn preflight_handler(Extension(state): Extension<AppState>, headers: HeaderMap) -> Response {
430 let Some(origin) = headers.get("origin").and_then(|v| v.to_str().ok()) else {
431 return StatusCode::METHOD_NOT_ALLOWED.into_response();
433 };
434 if !state
435 .config
436 .cors_allowed_origins
437 .iter()
438 .any(|a| a == origin)
439 {
440 return StatusCode::FORBIDDEN.into_response();
441 }
442
443 let mut resp = StatusCode::NO_CONTENT.into_response();
444 let h = resp.headers_mut();
445 h.insert(
446 "access-control-allow-origin",
447 HeaderValue::from_str(origin).unwrap_or(HeaderValue::from_static("")),
448 );
449 h.insert(
450 "access-control-allow-methods",
451 HeaderValue::from_static("POST"),
452 );
453 h.insert(
454 "access-control-allow-headers",
455 HeaderValue::from_static("content-type, authorization"),
456 );
457 h.insert("vary", HeaderValue::from_static("Origin"));
458 resp
459}
460
461fn check_cors(config: &CreateReportConfig, headers: &HeaderMap) -> Option<Response> {
464 let origin = headers.get("origin").and_then(|v| v.to_str().ok())?;
465 if config.cors_allowed_origins.iter().any(|a| a == origin) {
466 None
467 } else {
468 Some(StatusCode::FORBIDDEN.into_response())
469 }
470}
471
472fn attach_cors_headers(resp: &mut Response, config: &CreateReportConfig, req_headers: &HeaderMap) {
473 let Some(origin) = req_headers.get("origin").and_then(|v| v.to_str().ok()) else {
474 return;
475 };
476 if !config.cors_allowed_origins.iter().any(|a| a == origin) {
477 return;
478 }
479 let h = resp.headers_mut();
480 if let Ok(v) = HeaderValue::from_str(origin) {
481 h.insert("access-control-allow-origin", v);
482 }
483 h.insert("vary", HeaderValue::from_static("Origin"));
484}
485
486fn auth_required() -> Response {
487 xrpc_error(
488 StatusCode::UNAUTHORIZED,
489 "AuthenticationRequired",
490 "authentication required",
491 )
492}
493
494fn rate_limited() -> Response {
495 (StatusCode::TOO_MANY_REQUESTS, Json(RATE_LIMIT_BODY)).into_response()
496}
497
498fn xrpc_error(status: StatusCode, error: &'static str, message: &'static str) -> Response {
499 (status, Json(XrpcErrorBody { error, message })).into_response()
500}
501
502fn extract_did_from_at_uri(uri: &str) -> Option<&str> {
503 uri.strip_prefix("at://")
504 .and_then(|rest| rest.split('/').next())
505 .filter(|did| did.starts_with("did:"))
506}
507
508fn rfc3339_now() -> Result<String, ()> {
509 let dt = OffsetDateTime::now_utc();
510 let formatted = dt.format(&CTS_FORMAT).map_err(|_| ())?;
511 Ok(format!("{formatted}Z"))
512}
513
514fn rfc3339_minus_secs(secs: i64) -> Result<String, ()> {
515 let dt = OffsetDateTime::now_utc() - Duration::from_secs(secs.max(0) as u64);
516 let formatted = dt.format(&CTS_FORMAT).map_err(|_| ())?;
517 Ok(format!("{formatted}Z"))
518}
519
520fn subject_to_json(s: &Subject) -> serde_json::Value {
521 match s {
522 Subject::Repo { did } => serde_json::json!({
523 "$type": "com.atproto.admin.defs#repoRef",
524 "did": did,
525 }),
526 Subject::Strong { uri, cid } => serde_json::json!({
527 "$type": "com.atproto.repo.strongRef",
528 "uri": uri,
529 "cid": cid,
530 }),
531 }
532}
533
534#[cfg(test)]
535mod tests {
536 use super::*;
537
538 #[test]
539 fn extract_did_accepts_valid_at_uri() {
540 let did = extract_did_from_at_uri("at://did:plc:abc/col/rkey").unwrap();
541 assert_eq!(did, "did:plc:abc");
542 }
543
544 #[test]
545 fn extract_did_rejects_non_at_uri() {
546 assert!(extract_did_from_at_uri("https://example.com").is_none());
547 }
548
549 #[test]
550 fn extract_did_rejects_non_did_authority() {
551 assert!(extract_did_from_at_uri("at://example.com/foo").is_none());
554 }
555
556 #[test]
557 fn reason_types_allowlist_pins_f11_literal_set() {
558 assert_eq!(ACCEPTED_REASON_TYPES.len(), 6);
562 }
563}