use axum::Json;
use axum::extract::rejection::JsonRejection;
use axum::extract::{Path, State};
use futures::StreamExt;
use serde::Serialize;
use serde_json::Value;
use uuid::Uuid;
use crw_core::error::CrwError;
use crw_core::types::{CrawlState, ScrapeRequest};
use crate::error::AppError;
use crate::state::AppState;
pub const MAX_BATCH_BODY_BYTES: usize = 8 * 1024 * 1024;
const VALIDATE_CONCURRENCY: usize = 256;
#[derive(Debug, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct BatchStartResponse {
pub success: bool,
pub id: String,
pub invalid_urls: Vec<String>,
}
pub async fn start_batch(
State(state): State<AppState>,
body: Result<Json<Value>, JsonRejection>,
) -> Result<Json<BatchStartResponse>, AppError> {
let Json(mut raw) = body.map_err(AppError::from)?;
let obj = raw
.as_object_mut()
.ok_or_else(|| CrwError::InvalidRequest("request body must be a JSON object".into()))?;
let urls_val = obj
.remove("urls")
.ok_or_else(|| CrwError::InvalidRequest("`urls` is required".into()))?;
let urls: Vec<String> = serde_json::from_value(urls_val)
.map_err(|e| CrwError::InvalidRequest(format!("invalid `urls`: {e}")))?;
if urls.is_empty() {
return Err(AppError::from(CrwError::InvalidRequest(
"`urls` must contain at least one URL".into(),
)));
}
let max_urls = state.config.crawler.max_batch_urls;
if urls.len() > max_urls {
return Err(AppError::from(CrwError::InvalidRequest(format!(
"`urls` exceeds the maximum of {max_urls} URLs per batch (got {})",
urls.len()
))));
}
let ignore_invalid = obj
.remove("ignoreInvalidUrls")
.as_ref()
.and_then(Value::as_bool)
.unwrap_or(true);
let max_concurrency_override = obj
.remove("maxConcurrency")
.as_ref()
.and_then(Value::as_u64)
.map(|n| n as usize);
obj.insert(
"url".to_string(),
Value::String("https://placeholder.invalid/".into()),
);
let mut template: ScrapeRequest = serde_json::from_value(Value::Object(obj.clone()))
.map_err(|e| CrwError::InvalidRequest(format!("invalid batch scrape options: {e}")))?;
crate::state::validate_renderer_pin(template.renderer, template.render_js, &state)?;
template.url = String::new();
let checks = futures::stream::iter(urls.into_iter().map(|u| async move {
let ok = match url::Url::parse(&u) {
Ok(parsed) => crw_core::url_safety::validate_safe_url_resolved(&parsed)
.await
.is_ok(),
Err(_) => false,
};
(u, ok)
}))
.buffered(VALIDATE_CONCURRENCY)
.collect::<Vec<_>>()
.await;
let mut valid = Vec::new();
let mut invalid = Vec::new();
for (u, ok) in checks {
if ok { valid.push(u) } else { invalid.push(u) }
}
if valid.is_empty() {
return Err(AppError::from(CrwError::InvalidRequest(
"no valid URLs to scrape".into(),
)));
}
if !invalid.is_empty() && !ignore_invalid {
return Err(AppError::from(CrwError::InvalidRequest(format!(
"invalid URLs: {}",
invalid.join(", ")
))));
}
let id = state
.start_batch_job(valid, template, max_concurrency_override)
.await;
Ok(Json(BatchStartResponse {
success: true,
id: id.to_string(),
invalid_urls: invalid,
}))
}
pub async fn get_batch(
State(state): State<AppState>,
Path(id): Path<Uuid>,
) -> Result<Json<CrawlState>, AppError> {
let jobs = state.crawl_jobs.read().await;
let job = jobs
.get(&id)
.ok_or_else(|| CrwError::NotFound(format!("Batch job {id} not found")))?;
Ok(Json(job.rx.borrow().clone()))
}
pub async fn cancel_batch(state: State<AppState>, id: Path<Uuid>) -> Result<Json<Value>, AppError> {
super::crawl::cancel_crawl(state, id).await
}
#[cfg(test)]
mod tests {
use super::BatchStartResponse;
#[test]
fn batch_start_response_is_camel_case() {
let v = serde_json::to_value(BatchStartResponse {
success: true,
id: "job-1".into(),
invalid_urls: vec!["bad".into()],
})
.unwrap();
assert!(
v.get("invalidUrls").is_some(),
"expected camelCase invalidUrls"
);
assert!(v.get("invalid_urls").is_none(), "snake_case key leaked");
assert!(
v.get("invalidURLs").is_none(),
"v2 capital-URL key leaked into v1"
);
}
}