diff --git a/client/tests/integration/tasks.rs b/client/tests/integration/tasks.rs index 3b63b4198..e2159fbbc 100644 --- a/client/tests/integration/tasks.rs +++ b/client/tests/integration/tasks.rs @@ -9,7 +9,10 @@ async fn task_list(app: Arc, account: Account, client: DivviupClient fixtures::task(&app, &account).await, ]; let response_tasks = client.tasks(account.id).await?; - assert_same_json_representation(&tasks, &response_tasks); + assert_eq!(tasks.len(), response_tasks.len()); + for (task, response_task) in tasks.iter().zip(response_tasks.iter()) { + assert_same_json_representation_ignoring_query_type(task, response_task); + } Ok(()) } @@ -17,7 +20,7 @@ async fn task_list(app: Arc, account: Account, client: DivviupClient async fn get_task(app: Arc, account: Account, client: DivviupClient) -> TestResult { let task = fixtures::task(&app, &account).await; let response_task = client.task(&task.id).await?; - assert_same_json_representation(&task, &response_task); + assert_same_json_representation_ignoring_query_type(&task, &response_task); Ok(()) } @@ -45,7 +48,7 @@ async fn create_task(app: Arc, account: Account, client: DivviupClie .one(app.db()) .await? .unwrap(); - assert_same_json_representation(&task_from_db, &response_task); + assert_same_json_representation_ignoring_query_type(&task_from_db, &response_task); Ok(()) } @@ -85,7 +88,7 @@ async fn create_task_time_bucketed_fixed_size( .one(app.db()) .await? .unwrap(); - assert_same_json_representation(&task_from_db, &response_task); + assert_same_json_representation_ignoring_query_type(&task_from_db, &response_task); Ok(()) } diff --git a/compose.yaml b/compose.yaml index 85bc0aded..f2d15926e 100644 --- a/compose.yaml +++ b/compose.yaml @@ -104,6 +104,10 @@ services: --name=helper --api-url=http://janus_2_aggregator:8080/aggregator-api \ --bearer-token=0000 && \ touch /tmp/done) + volumes: + - type: volume + source: pair_aggregator_state + target: /tmp network_mode: service:divviup_api depends_on: divviup_api: @@ -277,6 +281,9 @@ services: CONFIG_FILE: /janus_2_garbage_collector.yaml <<: *janus_environment +volumes: + pair_aggregator_state: + configs: postgres_init: content: | diff --git a/documentation/openapi.yml b/documentation/openapi.yml index c6fa8b2a9..1142fc852 100644 --- a/documentation/openapi.yml +++ b/documentation/openapi.yml @@ -276,6 +276,9 @@ paths: format: uuid vdaf: $ref: "#/components/schemas/Vdaf" + query_type: + type: string + enum: [TimeInterval, FixedSize] min_batch_size: type: number max_batch_size: @@ -694,6 +697,9 @@ components: format: uuid vdaf: $ref: "#/components/schemas/Vdaf" + query_type: + type: string + enum: [TimeInterval, FixedSize] min_batch_size: type: number max_batch_size: @@ -929,8 +935,10 @@ components: is_first_party: type: boolean query_types: - type: string - enum: [TimeInterval, FixedSize] + type: array + items: + type: string + enum: [TimeInterval, FixedSize] vdafs: type: string examples: diff --git a/migration/README.md b/migration/README.md index 726d90b7f..c894d7017 100644 --- a/migration/README.md +++ b/migration/README.md @@ -11,7 +11,7 @@ This is the standard migrator CLI that comes with SeaORM. - Generate a new migration file ```sh - cargo run -- migrate generate MIGRATION_NAME + cargo run -- generate MIGRATION_NAME ``` - Apply all pending migrations ```sh diff --git a/migration/src/lib.rs b/migration/src/lib.rs index 3971eede4..63b4b2e0a 100644 --- a/migration/src/lib.rs +++ b/migration/src/lib.rs @@ -26,6 +26,7 @@ mod m20240214_215101_upload_metrics; mod m20240411_195358_time_bucketed_fixed_size; mod m20240416_172920_task_deleted_at; mod m20250801_164739_aggregation_job_metrics; +mod m20260921_223229_add_query_type_column; pub struct Migrator; @@ -59,6 +60,7 @@ impl MigratorTrait for Migrator { Box::new(m20240411_195358_time_bucketed_fixed_size::Migration), Box::new(m20240416_172920_task_deleted_at::Migration), Box::new(m20250801_164739_aggregation_job_metrics::Migration), + Box::new(m20260921_223229_add_query_type_column::Migration), ] } } diff --git a/migration/src/m20260921_223229_add_query_type_column.rs b/migration/src/m20260921_223229_add_query_type_column.rs new file mode 100644 index 000000000..e0e225024 --- /dev/null +++ b/migration/src/m20260921_223229_add_query_type_column.rs @@ -0,0 +1,105 @@ +use sea_orm::{sea_query::extension::postgres::Type, DbBackend}; +use sea_orm_migration::prelude::*; + +#[derive(DeriveMigrationName)] +pub struct Migration; + +#[async_trait::async_trait] +impl MigrationTrait for Migration { + async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> { + let db = manager.get_connection(); + // Use an enum in Postgres, and a string in SQLite. + let mut column_def = if db.get_database_backend() == DbBackend::Postgres { + manager + .create_type( + Type::create() + .as_enum(QueryType::Enum) + .values([QueryType::TimeInterval, QueryType::FixedSize]) + .to_owned(), + ) + .await?; + ColumnDef::new(Task::QueryType) + .custom(QueryType::Enum) + .to_owned() + } else { + ColumnDef::new(Task::QueryType).string().to_owned() + }; + // Add the column as a nullable column. + manager + .alter_table( + Table::alter() + .table(Task::Table) + .add_column(column_def.null().default(Expr::null())) + .to_owned(), + ) + .await?; + // Backfill the column. + manager + .execute( + Query::update() + .table(Task::Table) + .value( + Task::QueryType, + Expr::case( + Expr::column(Task::MaxBatchSize).is_not_null(), + Expr::cast_as( + Expr::value(QueryType::FixedSize.unquoted()), + QueryType::Enum, + ), + ) + .finally(Expr::cast_as( + Expr::value(QueryType::TimeInterval.unquoted()), + QueryType::Enum, + )), + ) + .to_owned(), + ) + .await?; + // Change the column to be not nullable. + manager + .alter_table( + Table::alter() + .table(Task::Table) + .modify_column(column_def.not_null()) + .to_owned(), + ) + .await + } + + async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> { + manager + .alter_table( + Table::alter() + .table(Task::Table) + .drop_column(Task::QueryType) + .to_owned(), + ) + .await?; + let db = manager.get_connection(); + if db.get_database_backend() == DbBackend::Postgres { + manager + .drop_type(Type::drop().name(QueryType::Enum).to_owned()) + .await?; + } + Ok(()) + } +} + +#[derive(Iden)] +enum Task { + Table, + + MaxBatchSize, + QueryType, +} + +#[derive(Iden)] +pub enum QueryType { + #[iden = "query_type"] + Enum, + + #[iden = "TIME_INTERVAL"] + TimeInterval, + #[iden = "FIXED_SIZE"] + FixedSize, +} diff --git a/src/clients/aggregator_client/api_types.rs b/src/clients/aggregator_client/api_types.rs index bf9c4d971..680a8c0b8 100644 --- a/src/clients/aggregator_client/api_types.rs +++ b/src/clients/aggregator_client/api_types.rs @@ -154,7 +154,8 @@ impl From for Vdaf { pub enum QueryType { TimeInterval, FixedSize { - max_batch_size: u64, + #[serde(skip_serializing_if = "Option::is_none")] + max_batch_size: Option, #[serde(skip_serializing_if = "Option::is_none")] batch_time_window_size: Option, }, diff --git a/src/entity/aggregator/query_type_name.rs b/src/entity/aggregator/query_type_name.rs index 89758f259..abe4e171e 100644 --- a/src/entity/aggregator/query_type_name.rs +++ b/src/entity/aggregator/query_type_name.rs @@ -9,6 +9,8 @@ use std::{ str::FromStr, }; +use crate::entity::task; + /// https://www.ietf.org/archive/id/draft-ietf-ppm-dap-05.html#name-queries #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Hash, PartialOrd, Ord)] pub enum QueryTypeName { @@ -61,6 +63,15 @@ impl From<&str> for QueryTypeName { } } +impl From for QueryTypeName { + fn from(value: task::QueryType) -> Self { + match value { + task::QueryType::TimeInterval => Self::TimeInterval, + task::QueryType::FixedSize => Self::FixedSize, + } + } +} + #[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)] pub struct QueryTypeNameSet(Set); diff --git a/src/entity/task.rs b/src/entity/task.rs index 2827ca183..6f3f1957d 100644 --- a/src/entity/task.rs +++ b/src/entity/task.rs @@ -20,6 +20,8 @@ mod provisionable_task; pub use provisionable_task::ProvisionableTask; pub mod model; pub use model::*; +pub mod query_type; +pub use query_type::*; pub const DEFAULT_EXPIRATION_DURATION: Duration = Duration::days(365); diff --git a/src/entity/task/model.rs b/src/entity/task/model.rs index a656a9b46..b9b221254 100644 --- a/src/entity/task/model.rs +++ b/src/entity/task/model.rs @@ -1,8 +1,8 @@ use crate::{ clients::aggregator_client::{api_types::TaskAggregationJobMetrics, TaskUploadMetrics}, entity::{ - account, json::Json, membership, AccountColumn, Accounts, Aggregator, AggregatorColumn, - Aggregators, CollectorCredentialColumn, CollectorCredentials, + account, json::Json, membership, task::QueryType, AccountColumn, Accounts, Aggregator, + AggregatorColumn, Aggregators, CollectorCredentialColumn, CollectorCredentials, }, }; use sea_orm::{ @@ -26,6 +26,7 @@ pub struct Model { pub account_id: Uuid, pub name: String, pub vdaf: Json, + pub query_type: QueryType, pub min_batch_size: i64, pub max_batch_size: Option, pub batch_time_window_size_seconds: Option, diff --git a/src/entity/task/new_task.rs b/src/entity/task/new_task.rs index 52dc89f69..7b0e12a3a 100644 --- a/src/entity/task/new_task.rs +++ b/src/entity/task/new_task.rs @@ -1,9 +1,9 @@ use super::*; use crate::{ - clients::aggregator_client::api_types::{AggregatorVdaf, QueryType}, + clients::aggregator_client::api_types::AggregatorVdaf, entity::{ - aggregator::{Feature, Role}, - Account, CollectorCredential, Protocol, + aggregator::{Feature, QueryTypeName, Role}, + task, Account, CollectorCredential, Protocol, }, handler::Error, }; @@ -29,6 +29,8 @@ pub struct NewTask { #[validate(required, nested)] pub vdaf: Option, + pub query_type: Option, + #[validate(required, range(min = 100))] pub min_batch_size: Option, @@ -98,15 +100,19 @@ impl NewTask { } } - fn validate_batch_time_window_size(&self, errors: &mut ValidationErrors) { + fn validate_batch_time_window_size( + &self, + query_type: QueryType, + errors: &mut ValidationErrors, + ) { let window = self.batch_time_window_size_seconds; if let Some(window) = window { - if self.max_batch_size.is_none() { + if query_type != QueryType::FixedSize { errors.add( "batch_time_window_size_seconds", - ValidationError::new("missing-max-batch-size"), + ValidationError::new("wrong-query-type"), ); - } + }; if let Some(precision) = self.time_precision_seconds { if window % precision != 0 { errors.add( @@ -327,11 +333,30 @@ impl NewTask { &self, leader: &Aggregator, helper: &Aggregator, + query_type: &QueryTypeName, errors: &mut ValidationErrors, ) { - let name = self.query_type().name(); - if !leader.query_types.contains(&name) || !helper.query_types.contains(&name) { - errors.add("max_batch_size", ValidationError::new("not-supported")); + if !leader.query_types.contains(query_type) || !helper.query_types.contains(query_type) { + errors.add("query_type", ValidationError::new("not-supported")); + } + } + + fn validate_query_type( + &self, + query_type: task::QueryType, + protocol: &Protocol, + errors: &mut ValidationErrors, + ) { + // Note DAP draft-09 says that `max_batch_size` is optional. In DAP draft-04, + // `max_batch_size` was mandatory for fixed size tasks, and previous versions of divviup-api + // relied on that linkage. + match (protocol, query_type, self.max_batch_size) { + (Protocol::Dap09, QueryType::TimeInterval, None) + | (Protocol::Dap09, QueryType::FixedSize, Some(_)) + | (Protocol::Dap09, QueryType::FixedSize, None) => {} + (Protocol::Dap09, QueryType::TimeInterval, Some(_)) => { + errors.add("max_batch_size", ValidationError::new("conflict")); + } } } @@ -341,8 +366,19 @@ impl NewTask { db: &impl ConnectionTrait, ) -> Result { let mut errors = Validate::validate(self).err().unwrap_or_default(); + + // Backfill the query type from `max_batch_size` if necessary. + // + // This was previously not part of requests. For the initial support of DAP draft-04, the + // query type was inferred from the presence or absence of a `max_batch_size` value. + let query_type = self.query_type.unwrap_or(match self.max_batch_size { + Some(_) => task::QueryType::FixedSize, + None => task::QueryType::TimeInterval, + }); + let query_type_name = QueryTypeName::from(query_type); + self.validate_min_lte_max(&mut errors); - self.validate_batch_time_window_size(&mut errors); + self.validate_batch_time_window_size(query_type, &mut errors); let aggregators = self.validate_aggregators(&account, db, &mut errors).await; let collector_credential = self .validate_collector_credential( @@ -354,7 +390,8 @@ impl NewTask { .await; let aggregator_vdaf = if let Some((leader, helper, protocol)) = aggregators.as_ref() { - self.validate_query_type_is_supported(leader, helper, &mut errors); + self.validate_query_type(query_type, protocol, &mut errors); + self.validate_query_type_is_supported(leader, helper, &query_type_name, &mut errors); self.populate_chunk_length(protocol); self.validate_vdaf_is_supported(leader, helper, protocol, &mut errors) } else { @@ -380,6 +417,7 @@ impl NewTask { leader_aggregator, helper_aggregator, vdaf: self.vdaf.clone().unwrap(), + query_type, aggregator_vdaf: aggregator_vdaf.unwrap(), min_batch_size: self.min_batch_size.unwrap(), max_batch_size: self.max_batch_size, @@ -394,15 +432,4 @@ impl NewTask { Err(errors) } } - - pub fn query_type(&self) -> QueryType { - if let Some(max_batch_size) = self.max_batch_size { - QueryType::FixedSize { - max_batch_size, - batch_time_window_size: self.batch_time_window_size_seconds, - } - } else { - QueryType::TimeInterval - } - } } diff --git a/src/entity/task/provisionable_task.rs b/src/entity/task/provisionable_task.rs index 3196fda2e..ffa91e4fa 100644 --- a/src/entity/task/provisionable_task.rs +++ b/src/entity/task/provisionable_task.rs @@ -2,7 +2,7 @@ use super::*; use crate::clients::HttpClient; use crate::{ clients::aggregator_client::api_types::{AggregatorVdaf, AuthenticationToken, QueryType}, - entity::{Account, CollectorCredential, Protocol, Task}, + entity::{task, Account, CollectorCredential, Protocol, Task}, handler::Error, Crypter, }; @@ -19,6 +19,7 @@ pub struct ProvisionableTask { pub helper_aggregator: Aggregator, pub vdaf: Vdaf, pub aggregator_vdaf: AggregatorVdaf, + pub query_type: task::QueryType, pub min_batch_size: u64, pub max_batch_size: Option, pub batch_time_window_size_seconds: Option, @@ -88,6 +89,7 @@ impl ProvisionableTask { account_id: self.account.id, name: self.name, vdaf: self.vdaf.into(), + query_type: self.query_type, min_batch_size: self.min_batch_size.try_into()?, max_batch_size: self.max_batch_size.map(TryInto::try_into).transpose()?, batch_time_window_size_seconds: self @@ -127,13 +129,12 @@ impl ProvisionableTask { } pub fn query_type(&self) -> QueryType { - if let Some(max_batch_size) = self.max_batch_size { - QueryType::FixedSize { - max_batch_size, + match self.query_type { + task::QueryType::TimeInterval => QueryType::TimeInterval, + task::QueryType::FixedSize => QueryType::FixedSize { + max_batch_size: self.max_batch_size, batch_time_window_size: self.batch_time_window_size_seconds, - } - } else { - QueryType::TimeInterval + }, } } } diff --git a/src/entity/task/query_type.rs b/src/entity/task/query_type.rs new file mode 100644 index 000000000..44857c0f0 --- /dev/null +++ b/src/entity/task/query_type.rs @@ -0,0 +1,11 @@ +use sea_orm::{DeriveActiveEnum, EnumIter}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, EnumIter, DeriveActiveEnum, Serialize, Deserialize)] +#[sea_orm(rs_type = "String", db_type = "Enum", enum_name = "query_type")] +pub enum QueryType { + #[sea_orm(string_value = "TIME_INTERVAL")] + TimeInterval, + #[sea_orm(string_value = "FIXED_SIZE")] + FixedSize, +} diff --git a/test-support/src/fixtures.rs b/test-support/src/fixtures.rs index 6b7cc735b..c7b3bf8e5 100644 --- a/test-support/src/fixtures.rs +++ b/test-support/src/fixtures.rs @@ -110,6 +110,7 @@ pub async fn task(app: &DivviupApi, account: &Account) -> Task { account_id: account.id, name: random_name(), vdaf: task::vdaf::Vdaf::Count.into(), + query_type: task::QueryType::FixedSize, min_batch_size: 100, max_batch_size: Some(200), batch_time_window_size_seconds: None, diff --git a/test-support/src/lib.rs b/test-support/src/lib.rs index a01c912a0..43debbcb6 100644 --- a/test-support/src/lib.rs +++ b/test-support/src/lib.rs @@ -576,3 +576,16 @@ where serde_json::to_value(expected).unwrap() ); } + +// TODO(#2554): Remove this once divviup-client is updated to reflect query types. +#[track_caller] +pub fn assert_same_json_representation_ignoring_query_type( + actual: &impl Serialize, + expected: &impl Serialize, +) { + let mut modified = serde_json::to_value(actual).unwrap(); + // Ignore the new query type field in divviup-api, as it has not been reflected in + // divviup-client yet. + modified.as_object_mut().unwrap().remove("query_type"); + assert_same_json_representation(&modified, expected); +} diff --git a/tests/integration/new_task.rs b/tests/integration/new_task.rs index 615b61eb7..6804ed21f 100644 --- a/tests/integration/new_task.rs +++ b/tests/integration/new_task.rs @@ -95,7 +95,22 @@ async fn time_bucketed_fixed_size(app: DivviupApi) -> TestResult { ..Default::default() }, "batch_time_window_size_seconds", - &["missing-max-batch-size"], + &["wrong-query-type"], + ) + .await; + + assert_errors( + &app, + &mut NewTask { + leader_aggregator_id: Some(leader.id.to_string()), + helper_aggregator_id: Some(helper.id.to_string()), + time_precision_seconds: Some(300), + query_type: Some(task::QueryType::TimeInterval), + batch_time_window_size_seconds: Some(300), + ..Default::default() + }, + "batch_time_window_size_seconds", + &["wrong-query-type"], ) .await; @@ -130,6 +145,22 @@ async fn time_bucketed_fixed_size(app: DivviupApi) -> TestResult { ) .await; + assert_no_errors( + &app, + &mut NewTask { + leader_aggregator_id: Some(leader.id.to_string()), + helper_aggregator_id: Some(helper.id.to_string()), + time_precision_seconds: Some(123), + query_type: Some(task::QueryType::FixedSize), + min_batch_size: Some(100), + max_batch_size: None, + batch_time_window_size_seconds: Some(300), + ..Default::default() + }, + "leader_aggregator_id", + ) + .await; + let mut leader = fixtures::aggregator(&app, None).await.into_active_model(); leader.role = ActiveValue::Set(Role::Leader); let leader = leader.update(app.db()).await?;