use crate::{config::DaphneWorkerConfig, dap::dap_response_to_worker};
use daphne::{
constants,
messages::{Duration, Id, Time},
roles::{DapAggregator, DapHelper, DapLeader},
DapAbort, DapCollectJob, DapError, DapResponse,
};
use prio::codec::{Decode, Encode};
use serde::{Deserialize, Serialize};
use worker::*;
#[derive(Debug, Deserialize, Serialize)]
pub struct DaphneWorkerReportSelector {
pub max_agg_jobs: u64,
pub max_reports: u64,
}
macro_rules! parse_id {
(
$option_str:expr
) => {
match $option_str {
Some(ref id_base64url) => {
match base64::decode_config(id_base64url, base64::URL_SAFE_NO_PAD) {
Ok(ref id_raw) => match Id::get_decoded(id_raw) {
Ok(id) => id,
Err(_) => return Response::error("Bad Request", 400),
},
Err(_) => return Response::error("Bad Request", 400),
}
}
None => return Response::error("Bad Request", 400),
}
};
}
#[derive(Default)]
pub struct DaphneWorkerRouter {
pub enable_internal_test: bool,
}
impl DaphneWorkerRouter {
pub async fn handle_request(&self, req: Request, env: Env) -> Result<Response> {
let router = Router::new().get_async("/:version/hpke_config", |req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
let req = config.worker_request_to_dap(req).await?;
match config.http_get_hpke_config(&req).await {
Ok(req) => dap_response_to_worker(req),
Err(e) => abort(e),
}
});
let router = match env.var("DAP_AGGREGATOR_ROLE")?.to_string().as_ref() {
"leader" => {
router
.post_async("/:version/upload", |req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
let req = config.worker_request_to_dap(req).await?;
match config.http_post_upload(&req).await {
Ok(()) => Response::empty(),
Err(e) => abort(e),
}
})
.post_async("/:version/collect", |req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
let req = config.worker_request_to_dap(req).await?;
match config.http_post_collect(&req).await {
Ok(collect_uri) => {
let mut headers = Headers::new();
headers.set("Location", collect_uri.as_str())?;
Ok(Response::empty()
.unwrap()
.with_status(303)
.with_headers(headers))
}
Err(e) => abort(e),
}
})
.get_async(
"/:version/collect/task/:task_id/req/:collect_id",
|_req, ctx| async move {
let task_id = parse_id!(ctx.param("task_id"));
let collect_id = parse_id!(ctx.param("collect_id"));
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
match config.poll_collect_job(&task_id, &collect_id).await {
Ok(DapCollectJob::Done(collect_resp)) => {
dap_response_to_worker(DapResponse {
media_type: Some(constants::MEDIA_TYPE_COLLECT_RESP),
payload: collect_resp.get_encoded(),
})
}
Ok(DapCollectJob::Pending) => {
Ok(Response::empty().unwrap().with_status(202))
}
Ok(DapCollectJob::Unknown) => {
abort(DapAbort::BadRequest("unknown collect id".into()))
}
Err(e) => abort(e.into()),
}
},
)
.post_async("/internal/process", |mut req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
let report_sel: DaphneWorkerReportSelector = req.json().await?;
match config.process(&report_sel).await {
Ok(telem) => {
console_debug!("{:?}", telem);
Response::from_json(&telem)
}
Err(e) => abort(e),
}
})
.get_async(
"/internal/current_batch/task/:task_id",
|_req, ctx| async move {
let task_id = parse_id!(ctx.param("task_id"));
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
match config.internal_current_batch(&task_id).await {
Ok(batch_id) => Response::from_bytes(
batch_id.to_base64url().as_bytes().to_owned(),
),
Err(e) => abort(e.into()),
}
},
)
}
"helper" => router
.post_async("/:version/aggregate", |req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
let req = config.worker_request_to_dap(req).await?;
match config.http_post_aggregate(&req).await {
Ok(resp) => dap_response_to_worker(resp),
Err(e) => abort(e),
}
})
.post_async("/:version/aggregate_share", |req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
let req = config.worker_request_to_dap(req).await?;
match config.http_post_aggregate_share(&req).await {
Ok(resp) => dap_response_to_worker(resp),
Err(e) => abort(e),
}
}),
_ => return abort(DapError::fatal("unexpected role").into()),
};
let router = if self.enable_internal_test {
router
.post_async("/internal/delete_all", |_req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
match config.internal_delete_all().await {
Ok(()) => Response::empty(),
Err(e) => abort(e.into()),
}
})
.post_async("/internal/test/ready", |_req, _ctx| async move {
Response::from_json(&())
})
.post_async(
"/internal/test/endpoint_for_task",
|mut req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
let cmd: InternalTestEndpointForTask = req.json().await?;
config.internal_endpoint_for_task(cmd).await
},
)
.post_async("/internal/test/add_task", |mut req, ctx| async move {
let config = DaphneWorkerConfig::from_worker_context(ctx)?;
let cmd: InternalTestAddTask = req.json().await?;
config.internal_add_task(cmd).await?;
Response::from_json(&serde_json::json!({
"status": "success",
}))
})
} else {
router
};
let start = Date::now().as_millis();
let resp = router.run(req, env).await?;
let end = Date::now().as_millis();
console_log!("request completed in {}ms", end - start);
Ok(resp)
}
}
pub(crate) fn now() -> u64 {
Date::now().as_millis() / 1000
}
pub(crate) fn int_err<S: ToString>(s: S) -> Error {
console_error!("internal error: {}", s.to_string());
Error::RustError("internalError".to_string())
}
pub(crate) fn dap_err(e: Error) -> DapError {
DapError::Fatal(format!("worker: {}", e))
}
fn abort(e: DapAbort) -> Result<Response> {
match &e {
DapAbort::Internal(..) => {
console_error!("internal error: {}", e.to_string());
Err(Error::RustError("internalError".to_string()))
}
_ => {
let mut headers = Headers::new();
headers.set("Content-Type", "application/problem+json")?;
Ok(Response::from_json(&e.to_problem_details())?
.with_status(400)
.with_headers(headers))
}
}
}
#[derive(Clone, Copy, Debug, Deserialize)]
#[serde(rename_all = "snake_case")]
pub(crate) enum InternalTestRole {
Leader,
Helper,
}
#[derive(Deserialize)]
#[serde(rename_all = "snake_case")]
pub(crate) struct InternalTestEndpointForTask {
role: InternalTestRole,
}
#[derive(Deserialize)]
pub(crate) struct InternalTestVdaf {
#[serde(rename = "type")]
typ: String,
#[serde(skip_serializing_if = "Option::is_none")]
bits: Option<String>,
}
#[derive(Deserialize)]
#[serde(rename_all = "snake_case")]
pub(crate) struct InternalTestAddTask {
task_id: String, leader: Url,
helper: Url,
vdaf: InternalTestVdaf,
leader_authentication_token: String,
#[serde(skip_serializing_if = "Option::is_none")]
collector_authentication_token: Option<String>,
role: InternalTestRole,
verify_key: String, query_type: u8,
min_batch_size: u64,
#[serde(skip_serializing_if = "Option::is_none")]
max_batch_size: Option<u64>,
time_precision: Duration,
collector_hpke_config: String, task_expiration: Time,
}
mod config;
mod dap;
mod durable;