use crate::{
auth::{BearerToken, BearerTokenProvider},
hpke::{HpkeDecrypter, HpkeReceiverConfig},
messages::{
BatchSelector, CollectReq, CollectResp, HpkeCiphertext, HpkeConfig, Id,
PartialBatchSelector, Report, ReportId, ReportMetadata, Time, TransitionFailure,
},
roles::{DapAggregator, DapAuthorizedSender, DapHelper, DapLeader},
DapAbort, DapAggregateShare, DapBatchBucket, DapCollectJob, DapError, DapGlobalConfig,
DapHelperState, DapOutputShare, DapQueryConfig, DapRequest, DapResponse, DapTaskConfig,
};
use assert_matches::assert_matches;
use async_trait::async_trait;
use rand::{thread_rng, Rng};
use serde::{Deserialize, Serialize};
use std::{
borrow::{Borrow, Cow},
collections::{HashMap, HashSet, VecDeque},
hash::Hash,
ops::DerefMut,
sync::{Arc, Mutex},
time::SystemTime,
};
use url::Url;
#[derive(Eq, Hash, PartialEq)]
pub(crate) enum DapBatchBucketOwned {
FixedSize { batch_id: Id },
TimeInterval { batch_window: Time },
}
impl From<DapBatchBucketOwned> for PartialBatchSelector {
fn from(bucket: DapBatchBucketOwned) -> Self {
match bucket {
DapBatchBucketOwned::FixedSize { batch_id } => Self::FixedSize { batch_id },
DapBatchBucketOwned::TimeInterval { .. } => Self::TimeInterval,
}
}
}
impl<'a> DapBatchBucket<'a> {
pub(crate) fn to_owned_bucket(&self) -> DapBatchBucketOwned {
match self {
Self::FixedSize { batch_id } => DapBatchBucketOwned::FixedSize {
batch_id: (*batch_id).clone(),
},
Self::TimeInterval { batch_window } => DapBatchBucketOwned::TimeInterval {
batch_window: *batch_window,
},
}
}
}
pub(crate) struct MockAggregatorReportSelector(pub(crate) Id);
#[allow(dead_code)]
pub(crate) struct MockAggregator {
pub(crate) now: Time,
pub(crate) global_config: DapGlobalConfig,
pub(crate) tasks: HashMap<Id, DapTaskConfig>,
pub(crate) hpke_receiver_config_list: Vec<HpkeReceiverConfig>,
pub(crate) leader_token: BearerToken,
pub(crate) collector_token: Option<BearerToken>, pub(crate) report_store: Arc<Mutex<HashMap<Id, ReportStore>>>,
pub(crate) leader_state_store: Arc<Mutex<HashMap<Id, LeaderState>>>,
pub(crate) helper_state_store: Arc<Mutex<HashMap<HelperStateInfo, DapHelperState>>>,
pub(crate) agg_store: Arc<Mutex<HashMap<Id, HashMap<DapBatchBucketOwned, AggStore>>>>,
}
#[allow(dead_code)]
impl MockAggregator {
async fn check_report_early_fail(
&self,
task_id: &Id,
bucket: &DapBatchBucketOwned,
metadata: &ReportMetadata,
) -> Option<TransitionFailure> {
let mut guard = self.agg_store.lock().expect("agg_store: failed to lock");
let agg_store = guard.entry(task_id.clone()).or_default();
if matches!(agg_store.get(bucket), Some(inner_agg_store) if inner_agg_store.collected) {
return Some(TransitionFailure::BatchCollected);
}
let mut guard = self
.report_store
.lock()
.expect("report_store: failed to lock");
let report_store = guard.entry(task_id.clone()).or_default();
if report_store.processed.contains(&metadata.id) {
return Some(TransitionFailure::ReportReplayed);
}
None
}
fn get_hpke_receiver_config_for(&self, hpke_config_id: u8) -> Option<&HpkeReceiverConfig> {
self.hpke_receiver_config_list
.iter()
.find(|&hpke_receiver_config| hpke_config_id == hpke_receiver_config.config.id)
}
fn assign_report_to_bucket(&self, report: &Report) -> Option<DapBatchBucketOwned> {
let mut rng = thread_rng();
let task_config = self
.tasks
.get(&report.task_id)
.expect("tasks: unrecognized task");
match task_config.query {
DapQueryConfig::FixedSize { .. } => {
let mut guard = self
.leader_state_store
.lock()
.expect("leader_state_store: failed to lock");
let leader_state_store = guard.entry(report.task_id.clone()).or_default();
for (batch_id, report_count) in leader_state_store.batch_queue.iter_mut() {
if *report_count < task_config.min_batch_size {
*report_count += 1;
return Some(DapBatchBucketOwned::FixedSize {
batch_id: batch_id.clone(),
});
}
}
let batch_id = Id(rng.gen());
leader_state_store
.batch_queue
.push_back((batch_id.clone(), 1));
Some(DapBatchBucketOwned::FixedSize { batch_id })
}
DapQueryConfig::TimeInterval => Some(DapBatchBucketOwned::TimeInterval {
batch_window: task_config.truncate_time(report.metadata.time),
}),
}
}
pub(crate) fn current_batch(&self, task_id: &Id) -> Option<Id> {
let task_config = self.tasks.get(task_id).expect("tasks: unrecognized task");
assert_matches!(task_config.query, DapQueryConfig::FixedSize { .. });
let guard = self
.leader_state_store
.lock()
.expect("leader_state_store: failed to lock");
let leader_state_store = guard
.get(task_id)
.expect("leader_state_store: unrecognized task");
leader_state_store
.batch_queue
.front()
.cloned() .map(|(batch_id, _report_count)| batch_id)
}
}
#[async_trait(?Send)]
impl<'a> BearerTokenProvider<'a> for MockAggregator {
type WrappedBearerToken = &'a BearerToken;
async fn get_leader_bearer_token_for(
&'a self,
_task_id: &'a Id,
) -> Result<Option<&'a BearerToken>, DapError> {
Ok(Some(&self.leader_token))
}
async fn get_collector_bearer_token_for(
&'a self,
_task_id: &'a Id,
) -> Result<Option<&'a BearerToken>, DapError> {
if let Some(ref collector_token) = self.collector_token {
Ok(Some(collector_token))
} else {
Err(DapError::fatal(
"MockAggregator not configured with Collector bearer token",
))
}
}
}
#[async_trait(?Send)]
impl<'a> HpkeDecrypter<'a> for MockAggregator {
type WrappedHpkeConfig = &'a HpkeConfig;
async fn get_hpke_config_for(
&'a self,
task_id: Option<&Id>,
) -> Result<&'a HpkeConfig, DapError> {
if self.hpke_receiver_config_list.is_empty() {
return Err(DapError::fatal("emtpy HPKE receiver config list"));
}
if task_id.is_none() {
return Err(DapError::Abort(DapAbort::MissingTaskId));
}
Ok(&self.hpke_receiver_config_list[0].config)
}
async fn can_hpke_decrypt(&self, _task_id: &Id, config_id: u8) -> Result<bool, DapError> {
Ok(self.get_hpke_receiver_config_for(config_id).is_some())
}
async fn hpke_decrypt(
&self,
_task_id: &Id,
info: &[u8],
aad: &[u8],
ciphertext: &HpkeCiphertext,
) -> Result<Vec<u8>, DapError> {
if let Some(hpke_receiver_config) = self.get_hpke_receiver_config_for(ciphertext.config_id)
{
Ok(hpke_receiver_config.decrypt(info, aad, &ciphertext.enc, &ciphertext.payload)?)
} else {
Err(DapError::Transition(TransitionFailure::HpkeUnknownConfigId))
}
}
}
#[async_trait(?Send)]
impl DapAuthorizedSender<BearerToken> for MockAggregator {
async fn authorize(
&self,
task_id: &Id,
media_type: &'static str,
_payload: &[u8],
) -> Result<BearerToken, DapError> {
Ok(self
.authorize_with_bearer_token(task_id, media_type)
.await?
.clone())
}
}
#[async_trait(?Send)]
impl<'srv, 'req> DapAggregator<'srv, 'req, BearerToken> for MockAggregator
where
'srv: 'req,
{
type WrappedDapTaskConfig = &'req DapTaskConfig;
async fn authorized(&self, req: &DapRequest<BearerToken>) -> Result<bool, DapError> {
self.bearer_token_authorized(req).await
}
fn get_global_config(&self) -> &DapGlobalConfig {
&self.global_config
}
async fn get_task_config_for(
&'srv self,
task_id: Cow<'req, Id>,
) -> Result<Option<&'req DapTaskConfig>, DapError> {
Ok(self.tasks.get(task_id.as_ref()))
}
fn get_current_time(&self) -> Time {
SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap()
.as_secs()
}
async fn is_batch_overlapping(
&self,
task_id: &Id,
batch_sel: &BatchSelector,
) -> Result<bool, DapError> {
let guard = self.agg_store.lock().expect("agg_store: failed to lock");
let task_config = self.tasks.get(task_id).expect("tasks: unrecognized task");
let agg_store = if let Some(agg_store) = guard.get(task_id) {
agg_store
} else {
return Ok(false);
};
for bucket in task_config.batch_span_for_sel(batch_sel)? {
if let Some(inner_agg_store) = agg_store.get(&bucket.to_owned_bucket()) {
if inner_agg_store.collected {
return Ok(true);
}
}
}
Ok(false)
}
async fn batch_exists(&self, task_id: &Id, batch_id: &Id) -> Result<bool, DapError> {
let guard = self.agg_store.lock().expect("agg_store: failed to lock");
if let Some(agg_store) = guard.get(task_id) {
Ok(agg_store
.get(&DapBatchBucketOwned::FixedSize {
batch_id: batch_id.clone(),
})
.is_some())
} else {
Ok(false)
}
}
async fn put_out_shares(
&self,
task_id: &Id,
part_batch_sel: &PartialBatchSelector,
out_shares: Vec<DapOutputShare>,
) -> Result<(), DapError> {
let task_config = self
.get_task_config_for(Cow::Borrowed(task_id))
.await?
.ok_or_else(|| DapError::fatal("task not found"))?;
let mut guard = self.agg_store.lock().expect("agg_store: failed to lock");
let agg_store = guard.entry(task_id.clone()).or_default();
for (bucket, agg_share_delta) in task_config
.batch_span_for_out_shares(part_batch_sel, out_shares)?
.into_iter()
{
let inner_agg_store = agg_store.entry(bucket.to_owned_bucket()).or_default();
inner_agg_store.agg_share.merge(agg_share_delta)?;
}
Ok(())
}
async fn get_agg_share(
&self,
task_id: &Id,
batch_sel: &BatchSelector,
) -> Result<DapAggregateShare, DapError> {
let mut guard = self.agg_store.lock().expect("agg_store: failed to lock");
let agg_store = guard.entry(task_id.clone()).or_default();
let task_config = self.tasks.get(task_id).expect("tasks: unrecognized task");
let mut agg_share = DapAggregateShare::default();
for bucket in task_config.batch_span_for_sel(batch_sel)? {
if let Some(inner_agg_store) = agg_store.get(&bucket.to_owned_bucket()) {
if inner_agg_store.collected {
return Err(DapError::Abort(DapAbort::BatchOverlap));
} else {
agg_share.merge(inner_agg_store.agg_share.clone())?;
}
}
}
Ok(agg_share)
}
async fn check_early_reject<'b>(
&self,
task_id: &Id,
part_batch_sel: &'b PartialBatchSelector,
report_meta: impl Iterator<Item = &'b ReportMetadata>,
) -> Result<HashMap<ReportId, TransitionFailure>, DapError> {
let task_config = self.tasks.get(task_id).expect("tasks: unrecognized task");
let span = task_config.batch_span_for_meta(part_batch_sel, report_meta)?;
let mut early_fails = HashMap::new();
for (bucket, report_meta) in span.iter() {
for metadata in report_meta.iter() {
if let Some(transition_failure) = self
.check_report_early_fail(task_id, &bucket.to_owned_bucket(), metadata)
.await
{
early_fails.insert(metadata.id.clone(), transition_failure);
};
let mut guard = self
.report_store
.lock()
.expect("report_store: failed to lock");
let report_store = guard.entry(task_id.clone()).or_default();
report_store.processed.insert(metadata.id.clone());
}
}
Ok(early_fails)
}
async fn mark_collected(
&self,
task_id: &Id,
batch_sel: &BatchSelector,
) -> Result<(), DapError> {
let mut guard = self.agg_store.lock().expect("agg_store: failed to lock");
let agg_store = guard.entry(task_id.clone()).or_default();
let task_config = self.tasks.get(task_id).expect("tasks: unrecognized task");
for bucket in task_config.batch_span_for_sel(batch_sel)? {
if let Some(inner_agg_store) = agg_store.get_mut(&bucket.to_owned_bucket()) {
inner_agg_store.collected = true;
}
}
Ok(())
}
}
#[async_trait(?Send)]
impl<'srv, 'req> DapHelper<'srv, 'req, BearerToken> for MockAggregator
where
'srv: 'req,
{
async fn put_helper_state(
&self,
task_id: &Id,
agg_job_id: &Id,
helper_state: &DapHelperState,
) -> Result<(), DapError> {
let helper_state_info = HelperStateInfo {
task_id: task_id.clone(),
agg_job_id: agg_job_id.clone(),
};
let mut helper_state_store_mutex_guard = self
.helper_state_store
.lock()
.map_err(|e| DapError::Fatal(e.to_string()))?;
let helper_state_store = helper_state_store_mutex_guard.deref_mut();
if helper_state_store.contains_key(&helper_state_info) {
return Err(DapError::Fatal(
"overwriting existing helper state".to_string(),
));
}
helper_state_store.insert(helper_state_info, helper_state.clone());
Ok(())
}
async fn get_helper_state(
&self,
task_id: &Id,
agg_job_id: &Id,
) -> Result<Option<DapHelperState>, DapError> {
let helper_state_info = HelperStateInfo {
task_id: task_id.clone(),
agg_job_id: agg_job_id.clone(),
};
let mut helper_state_store_mutex_guard = self
.helper_state_store
.lock()
.map_err(|e| DapError::Fatal(e.to_string()))?;
let helper_state_store = helper_state_store_mutex_guard.deref_mut();
if helper_state_store.contains_key(&helper_state_info) {
let helper_state = helper_state_store.remove(&helper_state_info);
return Ok(helper_state);
}
Ok(None)
}
}
#[async_trait(?Send)]
impl<'srv, 'req> DapLeader<'srv, 'req, BearerToken> for MockAggregator
where
'srv: 'req,
{
type ReportSelector = MockAggregatorReportSelector;
async fn put_report(&self, report: &Report) -> Result<(), DapError> {
let bucket = self
.assign_report_to_bucket(report)
.expect("could not determine batch for report");
if let Some(transition_failure) = self
.check_report_early_fail(&report.task_id, bucket.borrow(), &report.metadata)
.await
{
return Err(DapError::Transition(transition_failure));
};
let mut guard = self
.report_store
.lock()
.expect("report_store: failed to lock");
let queue = guard
.get_mut(&report.task_id)
.expect("report_store: unrecognized task")
.pending
.entry(bucket)
.or_default();
queue.push_back(report.clone());
Ok(())
}
async fn get_reports(
&self,
report_sel: &MockAggregatorReportSelector,
) -> Result<HashMap<Id, HashMap<PartialBatchSelector, Vec<Report>>>, DapError> {
let mut guard = self
.report_store
.lock()
.expect("report_store: failed to lock");
let task_id = &report_sel.0;
let task_config = self.tasks.get(task_id).expect("tasks: unrecognized task");
let report_store = guard.entry(task_id.clone()).or_default();
match task_config.query {
DapQueryConfig::TimeInterval { .. } => {
let mut reports = Vec::new();
for (_bucket, queue) in report_store.pending.iter_mut() {
if !queue.is_empty() {
reports.append(&mut queue.drain(..1).collect());
break;
}
}
return Ok(HashMap::from([(
task_id.clone(),
HashMap::from([(PartialBatchSelector::TimeInterval, reports)]),
)]));
}
DapQueryConfig::FixedSize { .. } => {
let bucket = if let Some(batch_id) = self.current_batch(task_id) {
DapBatchBucketOwned::FixedSize { batch_id }
} else {
return Ok(HashMap::default());
};
let queue = report_store
.pending
.get_mut(&bucket)
.expect("report_store: unknown bucket");
let reports = queue.drain(..1).collect();
return Ok(HashMap::from([(
task_id.clone(),
HashMap::from([(bucket.into(), reports)]),
)]));
}
}
}
async fn init_collect_job(&self, collect_req: &CollectReq) -> Result<Url, DapError> {
let mut rng = thread_rng();
let task_config = self
.get_task_config_for(Cow::Borrowed(&collect_req.task_id))
.await?
.ok_or_else(|| DapError::fatal("task not found"))?;
let mut leader_state_store_mutex_guard = self
.leader_state_store
.lock()
.map_err(|e| DapError::Fatal(e.to_string()))?;
let leader_state_store = leader_state_store_mutex_guard.deref_mut();
let collect_id = Id(rng.gen());
let collect_uri = task_config
.leader_url
.join(&format!(
"collect/task/{}/req/{}",
collect_req.task_id.to_base64url(),
collect_id.to_base64url(),
))
.map_err(|e| DapError::Fatal(e.to_string()))?;
let leader_state = leader_state_store
.entry(collect_req.task_id.clone())
.or_default();
leader_state.collect_ids.push_back(collect_id.clone());
let collect_job_state = CollectJobState::Pending(collect_req.clone());
leader_state
.collect_jobs
.insert(collect_id, collect_job_state);
Ok(collect_uri)
}
async fn poll_collect_job(
&self,
task_id: &Id,
collect_id: &Id,
) -> Result<DapCollectJob, DapError> {
let mut leader_state_store_mutex_guard = self
.leader_state_store
.lock()
.map_err(|e| DapError::Fatal(e.to_string()))?;
let leader_state_store = leader_state_store_mutex_guard.deref_mut();
let leader_state = leader_state_store
.get(task_id)
.ok_or_else(|| DapError::fatal("collect job not found for task_id"))?;
if let Some(collect_job_state) = leader_state.collect_jobs.get(collect_id) {
match collect_job_state {
CollectJobState::Pending(_) => Ok(DapCollectJob::Pending),
CollectJobState::Processed(resp) => Ok(DapCollectJob::Done(resp.clone())),
}
} else {
Ok(DapCollectJob::Unknown)
}
}
async fn get_pending_collect_jobs(&self) -> Result<Vec<(Id, CollectReq)>, DapError> {
let mut leader_state_store_mutex_guard = self
.leader_state_store
.lock()
.map_err(|e| DapError::Fatal(e.to_string()))?;
let leader_state_store = leader_state_store_mutex_guard.deref_mut();
let mut res = Vec::new();
for (_task_id, leader_state) in leader_state_store.iter() {
for collect_id in leader_state.collect_ids.iter() {
if let CollectJobState::Pending(collect_req) =
leader_state.collect_jobs.get(collect_id).unwrap()
{
res.push((collect_id.clone(), collect_req.clone()));
}
}
}
Ok(res)
}
async fn finish_collect_job(
&self,
task_id: &Id,
collect_id: &Id,
collect_resp: &CollectResp,
) -> Result<(), DapError> {
let mut leader_state_store_mutex_guard = self
.leader_state_store
.lock()
.map_err(|e| DapError::Fatal(e.to_string()))?;
let leader_state_store = leader_state_store_mutex_guard.deref_mut();
let leader_state = leader_state_store
.get_mut(task_id)
.ok_or_else(|| DapError::fatal("collect job not found for task_id"))?;
let collect_job = leader_state
.collect_jobs
.get_mut(collect_id)
.ok_or_else(|| DapError::fatal("collect job not found for collect_id"))?;
if let PartialBatchSelector::FixedSize { ref batch_id } = collect_resp.part_batch_sel {
leader_state
.batch_queue
.retain(|(id, _report_count)| id != batch_id);
}
match collect_job {
CollectJobState::Pending(_) => {
*collect_job = CollectJobState::Processed(collect_resp.clone());
let index = leader_state
.collect_ids
.iter()
.position(|r| r == collect_id)
.unwrap();
leader_state.collect_ids.remove(index);
Ok(())
}
CollectJobState::Processed(_) => {
Err(DapError::fatal("tried to overwrite collect response"))
}
}
}
async fn send_http_post(&self, _req: DapRequest<BearerToken>) -> Result<DapResponse, DapError> {
unreachable!("not implemented");
}
}
#[derive(Clone, Eq, Hash, PartialEq, Deserialize, Serialize)]
pub(crate) struct HelperStateInfo {
task_id: Id,
agg_job_id: Id,
}
#[derive(Default)]
pub(crate) struct ReportStore {
pub(crate) pending: HashMap<DapBatchBucketOwned, VecDeque<Report>>,
pub(crate) processed: HashSet<ReportId>,
}
pub(crate) enum CollectJobState {
Pending(CollectReq),
Processed(CollectResp),
}
#[derive(Default)]
pub(crate) struct LeaderState {
collect_ids: VecDeque<Id>,
collect_jobs: HashMap<Id, CollectJobState>,
batch_queue: VecDeque<(Id, u64)>, }
#[derive(Default)]
pub(crate) struct AggStore {
pub(crate) agg_share: DapAggregateShare,
pub(crate) collected: bool,
}