diff --git a/src/api/server/key.rs b/src/api/server/key.rs index 7a4a296c7..76d577485 100644 --- a/src/api/server/key.rs +++ b/src/api/server/key.rs @@ -106,24 +106,32 @@ fn valid_until_ts() -> MilliSecondsSinceUnixEpoch { MilliSecondsSinceUnixEpoch::from_system_time(timepoint).expect("UInt should not overflow") } -const MAX_KEYS_PER_QUERY: usize = 16 * 1024; +const MAX_KEYS_PER_QUERY: usize = 8192; +const MAX_SERVERS_PER_QUERY: usize = 4096; pub(crate) async fn get_remote_server_keys_batch_route( State(services): State, body: Ruma, ) -> Result { let start = MilliSecondsSinceUnixEpoch::now(); + if body.server_keys.is_empty() { + return Ok(get_remote_server_keys_batch::v2::Response::new(Vec::new())); + } else if body.server_keys.len() > MAX_SERVERS_PER_QUERY { + return Err!(Request(TooLarge( + "Too many server keys requested ({} > {MAX_SERVERS_PER_QUERY}", + body.server_keys.len() + ))); + } + let total_queried_keys = body .server_keys .values() .fold(body.server_keys.len(), |acc, q| acc.saturating_add(q.len())); if total_queried_keys > MAX_KEYS_PER_QUERY { - return Err!(Request(Forbidden( + return Err!(Request(TooLarge( "Too many keys requested ({total_queried_keys} > {MAX_KEYS_PER_QUERY})" ))); - } else if total_queried_keys == 0 { - return Ok(get_remote_server_keys_batch::v2::Response::new(Vec::new())); } debug_info!("Fetching {total_queried_keys} keys across {} servers", body.server_keys.len());