diff --git a/docs/grpc/docs.md b/docs/grpc/docs.md index 8a759e6a23..45515225a5 100644 --- a/docs/grpc/docs.md +++ b/docs/grpc/docs.md @@ -1432,6 +1432,8 @@ Note: 1kB = 1 vector of size 256. | | search_max_oversampling | [float](#float) | optional | | | upsert_max_batchsize | [uint64](#uint64) | optional | | | max_collection_vector_size_bytes | [uint64](#uint64) | optional | | +| read_rate_limit_per_sec | [uint32](#uint32) | optional | | +| write_rate_limit_per_sec | [uint32](#uint32) | optional | | diff --git a/docs/redoc/master/openapi.json b/docs/redoc/master/openapi.json index b17b517b32..68ee12d92d 100644 --- a/docs/redoc/master/openapi.json +++ b/docs/redoc/master/openapi.json @@ -7298,6 +7298,20 @@ "format": "uint", "minimum": 0, "nullable": true + }, + "read_rate_limit_per_sec": { + "description": "Max number of read operations per second per shard per peer", + "type": "integer", + "format": "uint", + "minimum": 0, + "nullable": true + }, + "write_rate_limit_per_sec": { + "description": "Max number of write operations per second per shard per peer", + "type": "integer", + "format": "uint", + "minimum": 0, + "nullable": true } } }, diff --git a/lib/api/src/grpc/conversions.rs b/lib/api/src/grpc/conversions.rs index 4537e65e5a..a38b39d778 100644 --- a/lib/api/src/grpc/conversions.rs +++ b/lib/api/src/grpc/conversions.rs @@ -1635,6 +1635,8 @@ impl From for segment::types::StrictModeConfig { max_collection_vector_size_bytes: value .max_collection_vector_size_bytes .map(|i| i as usize), + read_rate_limit_per_sec: value.read_rate_limit_per_sec.map(|i| i as usize), + write_rate_limit_per_sec: value.write_rate_limit_per_sec.map(|i| i as usize), } } } @@ -1654,6 +1656,8 @@ impl From for StrictModeConfig { max_collection_vector_size_bytes: value .max_collection_vector_size_bytes .map(|i| i as u64), + read_rate_limit_per_sec: value.read_rate_limit_per_sec.map(|i| i as u32), + write_rate_limit_per_sec: value.write_rate_limit_per_sec.map(|i| i as u32), } } } diff --git a/lib/api/src/grpc/proto/collections.proto b/lib/api/src/grpc/proto/collections.proto index fd10cdd2d6..24e9492957 100644 --- a/lib/api/src/grpc/proto/collections.proto +++ b/lib/api/src/grpc/proto/collections.proto @@ -322,6 +322,8 @@ message StrictModeConfig { optional float search_max_oversampling = 8; optional uint64 upsert_max_batchsize = 9; optional uint64 max_collection_vector_size_bytes = 10; + optional uint32 read_rate_limit_per_sec = 11; + optional uint32 write_rate_limit_per_sec = 12; } message CreateCollection { diff --git a/lib/api/src/grpc/qdrant.rs b/lib/api/src/grpc/qdrant.rs index 1e338e9469..654d6d3211 100644 --- a/lib/api/src/grpc/qdrant.rs +++ b/lib/api/src/grpc/qdrant.rs @@ -456,6 +456,10 @@ pub struct StrictModeConfig { pub upsert_max_batchsize: ::core::option::Option, #[prost(uint64, optional, tag = "10")] pub max_collection_vector_size_bytes: ::core::option::Option, + #[prost(uint32, optional, tag = "11")] + pub read_rate_limit_per_sec: ::core::option::Option, + #[prost(uint32, optional, tag = "12")] + pub write_rate_limit_per_sec: ::core::option::Option, } #[derive(validator::Validate)] #[derive(serde::Serialize)] diff --git a/lib/collection/src/collection/collection_ops.rs b/lib/collection/src/collection/collection_ops.rs index b384db2392..3011052806 100644 --- a/lib/collection/src/collection/collection_ops.rs +++ b/lib/collection/src/collection/collection_ops.rs @@ -176,7 +176,14 @@ impl Collection { config.strict_mode_config = Some(strict_mode_diff); } } + // update collection config self.collection_config.read().await.save(&self.path)?; + // apply config change to all shards + let shard_holder = self.shards_holder.read().await; + let updates = shard_holder + .all_shards() + .map(|replica_set| replica_set.on_strict_mode_config_update()); + future::try_join_all(updates).await?; Ok(()) } diff --git a/lib/collection/src/operations/types.rs b/lib/collection/src/operations/types.rs index f1f752e916..5ef0dda595 100644 --- a/lib/collection/src/operations/types.rs +++ b/lib/collection/src/operations/types.rs @@ -1107,7 +1107,6 @@ impl CollectionError { Self::Cancelled { .. } => true, Self::OutOfMemory { .. } => true, Self::PreConditionFailed { .. } => true, - Self::RateLimitExceeded { .. } => true, // Not transient Self::BadInput { .. } => false, Self::NotFound { .. } => false, @@ -1119,6 +1118,7 @@ impl CollectionError { Self::ObjectStoreError { .. } => false, Self::StrictMode { .. } => false, Self::InferenceError { .. } => false, + Self::RateLimitExceeded { .. } => false, } } diff --git a/lib/collection/src/operations/verification/mod.rs b/lib/collection/src/operations/verification/mod.rs index 925f007f99..7abfee4907 100644 --- a/lib/collection/src/operations/verification/mod.rs +++ b/lib/collection/src/operations/verification/mod.rs @@ -409,6 +409,8 @@ mod test { search_max_oversampling: Some(0.2), upsert_max_batchsize: None, max_collection_vector_size_bytes: None, + read_rate_limit_per_sec: None, + write_rate_limit_per_sec: None, }; fixture_collection(&strict_mode_config).await diff --git a/lib/collection/src/shards/dummy_shard.rs b/lib/collection/src/shards/dummy_shard.rs index c7e46d4b08..63152eedab 100644 --- a/lib/collection/src/shards/dummy_shard.rs +++ b/lib/collection/src/shards/dummy_shard.rs @@ -49,6 +49,8 @@ impl DummyShard { self.dummy() } + pub async fn on_strict_mode_config_update(&self) {} + pub fn get_telemetry_data(&self) -> LocalShardTelemetry { LocalShardTelemetry { variant_name: Some("dummy shard".into()), diff --git a/lib/collection/src/shards/forward_proxy_shard.rs b/lib/collection/src/shards/forward_proxy_shard.rs index 9a919dd02b..04088cdb18 100644 --- a/lib/collection/src/shards/forward_proxy_shard.rs +++ b/lib/collection/src/shards/forward_proxy_shard.rs @@ -206,6 +206,10 @@ impl ForwardProxyShard { self.wrapped_shard.on_optimizer_config_update().await } + pub async fn on_strict_mode_config_update(&self) { + self.wrapped_shard.on_strict_mode_config_update().await + } + pub fn trigger_optimizers(&self) { self.wrapped_shard.trigger_optimizers(); } diff --git a/lib/collection/src/shards/local_shard/mod.rs b/lib/collection/src/shards/local_shard/mod.rs index ca87713577..f95a5fa7b1 100644 --- a/lib/collection/src/shards/local_shard/mod.rs +++ b/lib/collection/src/shards/local_shard/mod.rs @@ -96,8 +96,8 @@ pub struct LocalShard { update_runtime: Handle, pub(super) search_runtime: Handle, disk_usage_watcher: DiskUsageWatcher, - read_rate_limiter: Option>, - write_rate_limiter: Option>, + read_rate_limiter: ParkingMutex>, + write_rate_limiter: ParkingMutex>, } /// Shard holds information about segments and WAL. @@ -196,6 +196,20 @@ impl LocalShard { let update_tracker = segment_holder.read().update_tracker(); + let read_rate_limiter = config.strict_mode_config.as_ref().and_then(|strict_mode| { + strict_mode + .read_rate_limit_per_sec + .map(RateLimiter::with_rate_per_sec) + }); + let read_rate_limiter = ParkingMutex::new(read_rate_limiter); + + let write_rate_limiter = config.strict_mode_config.as_ref().and_then(|strict_mode| { + strict_mode + .write_rate_limit_per_sec + .map(RateLimiter::with_rate_per_sec) + }); + let write_rate_limiter = ParkingMutex::new(write_rate_limiter); + drop(config); // release `shared_config` from borrow checker Self { @@ -214,8 +228,8 @@ impl LocalShard { optimizers_log, total_optimized_points, disk_usage_watcher, - read_rate_limiter: None, // TODO initialize rate limiter from config - write_rate_limiter: None, // TODO initialize rate limiter from config + read_rate_limiter, + write_rate_limiter, } } @@ -741,6 +755,28 @@ impl LocalShard { Ok(()) } + /// Apply shard's strict mode configuration update + /// - Update read and write rate limiters + pub async fn on_strict_mode_config_update(&self) { + let config = self.collection_config.read().await; + + if let Some(strict_mode_config) = &config.strict_mode_config { + // Update read rate limiter + if let Some(read_rate_limit_per_sec) = strict_mode_config.read_rate_limit_per_sec { + let mut read_rate_limiter_guard = self.read_rate_limiter.lock(); + read_rate_limiter_guard + .replace(RateLimiter::with_rate_per_sec(read_rate_limit_per_sec)); + } + + // update write rate limiter + if let Some(write_rate_limit_per_sec) = strict_mode_config.write_rate_limit_per_sec { + let mut write_rate_limiter_guard = self.write_rate_limiter.lock(); + write_rate_limiter_guard + .replace(RateLimiter::with_rate_per_sec(write_rate_limit_per_sec)); + } + } + } + pub fn trigger_optimizers(&self) { // Send a trigger signal and ignore errors because all error cases are acceptable: // - If receiver is already dead - we do not care @@ -1107,8 +1143,8 @@ impl LocalShard { /// /// Returns an error if the rate limit is exceeded. fn check_write_rate_limiter(&self) -> CollectionResult<()> { - if let Some(rate_limiter) = &self.write_rate_limiter { - if !rate_limiter.lock().check() { + if let Some(rate_limiter) = self.write_rate_limiter.lock().as_mut() { + if !rate_limiter.check() { return Err(CollectionError::RateLimitExceeded { description: "Write rate limit exceeded, retry later".to_string(), }); @@ -1121,8 +1157,8 @@ impl LocalShard { /// /// Returns an error if the rate limit is exceeded. fn check_read_rate_limiter(&self) -> CollectionResult<()> { - if let Some(rate_limiter) = &self.read_rate_limiter { - if !rate_limiter.lock().check() { + if let Some(rate_limiter) = self.read_rate_limiter.lock().as_mut() { + if !rate_limiter.check() { return Err(CollectionError::RateLimitExceeded { description: "Read rate limit exceeded, retry later".to_string(), }); diff --git a/lib/collection/src/shards/proxy_shard.rs b/lib/collection/src/shards/proxy_shard.rs index a8c6f979b9..16f13f33ea 100644 --- a/lib/collection/src/shards/proxy_shard.rs +++ b/lib/collection/src/shards/proxy_shard.rs @@ -85,6 +85,10 @@ impl ProxyShard { self.wrapped_shard.on_optimizer_config_update().await } + pub async fn on_strict_mode_config_update(&self) { + self.wrapped_shard.on_strict_mode_config_update().await; + } + pub fn trigger_optimizers(&self) { // TODO: we might want to defer this trigger until we unproxy self.wrapped_shard.trigger_optimizers(); diff --git a/lib/collection/src/shards/queue_proxy_shard.rs b/lib/collection/src/shards/queue_proxy_shard.rs index f94a2e9c29..c9ed750665 100644 --- a/lib/collection/src/shards/queue_proxy_shard.rs +++ b/lib/collection/src/shards/queue_proxy_shard.rs @@ -159,6 +159,13 @@ impl QueueProxyShard { .await } + pub async fn on_strict_mode_config_update(&self) { + self.inner_unchecked() + .wrapped_shard + .on_strict_mode_config_update() + .await + } + pub fn trigger_optimizers(&self) { self.inner_unchecked().wrapped_shard.trigger_optimizers(); } diff --git a/lib/collection/src/shards/replica_set/mod.rs b/lib/collection/src/shards/replica_set/mod.rs index b75d420e24..fd1e488d78 100644 --- a/lib/collection/src/shards/replica_set/mod.rs +++ b/lib/collection/src/shards/replica_set/mod.rs @@ -730,7 +730,15 @@ impl ShardReplicaSet { } } - /// Check if the are any locally disabled peers + pub(crate) async fn on_strict_mode_config_update(&self) -> CollectionResult<()> { + let read_local = self.local.read().await; + if let Some(shard) = &*read_local { + shard.on_strict_mode_config_update().await + } + Ok(()) + } + + /// Check if there are any locally disabled peers /// And if so, report them to the consensus pub fn sync_local_state(&self, get_shard_transfers: F) -> CollectionResult<()> where diff --git a/lib/collection/src/shards/shard.rs b/lib/collection/src/shards/shard.rs index f7f93a44b4..3dae04ec98 100644 --- a/lib/collection/src/shards/shard.rs +++ b/lib/collection/src/shards/shard.rs @@ -130,6 +130,16 @@ impl Shard { } } + pub async fn on_strict_mode_config_update(&self) { + match self { + Shard::Local(local_shard) => local_shard.on_strict_mode_config_update().await, + Shard::Proxy(proxy_shard) => proxy_shard.on_strict_mode_config_update().await, + Shard::ForwardProxy(proxy_shard) => proxy_shard.on_strict_mode_config_update().await, + Shard::QueueProxy(proxy_shard) => proxy_shard.on_strict_mode_config_update().await, + Shard::Dummy(dummy_shard) => dummy_shard.on_strict_mode_config_update().await, + } + } + pub fn trigger_optimizers(&self) { match self { Shard::Local(local_shard) => local_shard.trigger_optimizers(), diff --git a/lib/common/common/src/rate_limiting.rs b/lib/common/common/src/rate_limiting.rs index cea14a7b54..c3ade6507f 100644 --- a/lib/common/common/src/rate_limiting.rs +++ b/lib/common/common/src/rate_limiting.rs @@ -58,6 +58,12 @@ impl RateLimiter { } } + /// Create a new rate limiter from a rate per second. + pub fn with_rate_per_sec(rate_per_sec: usize) -> Self { + let rate = Rate::new(rate_per_sec as u64, Duration::from_secs(1)); + Self::new(rate) + } + /// Attempt to consume a token. Returns `true` if allowed, `false` otherwise. pub fn check(&mut self) -> bool { let now = Instant::now(); diff --git a/lib/segment/src/types.rs b/lib/segment/src/types.rs index e50a841c3d..9991a769d4 100644 --- a/lib/segment/src/types.rs +++ b/lib/segment/src/types.rs @@ -711,6 +711,14 @@ pub struct StrictModeConfig { /// Max size of a collections vector storage in bytes #[serde(skip_serializing_if = "Option::is_none")] pub max_collection_vector_size_bytes: Option, + + /// Max number of read operations per second per shard per peer + #[serde(skip_serializing_if = "Option::is_none")] + pub read_rate_limit_per_sec: Option, + + /// Max number of write operations per second per shard per peer + #[serde(skip_serializing_if = "Option::is_none")] + pub write_rate_limit_per_sec: Option, } impl Eq for StrictModeConfig {} @@ -729,6 +737,8 @@ impl Hash for StrictModeConfig { search_max_oversampling: _, upsert_max_batchsize, max_collection_vector_size_bytes, + read_rate_limit_per_sec, + write_rate_limit_per_sec, } = self; ( enabled, @@ -740,6 +750,8 @@ impl Hash for StrictModeConfig { search_allow_exact, upsert_max_batchsize, max_collection_vector_size_bytes, + read_rate_limit_per_sec, + write_rate_limit_per_sec, ) .hash(state); } diff --git a/lib/storage/src/content_manager/conversions.rs b/lib/storage/src/content_manager/conversions.rs index 40252f7697..3d4b5b6a26 100644 --- a/lib/storage/src/content_manager/conversions.rs +++ b/lib/storage/src/content_manager/conversions.rs @@ -89,6 +89,8 @@ pub fn strict_mode_from_api(value: api::grpc::qdrant::StrictModeConfig) -> Stric max_collection_vector_size_bytes: value .max_collection_vector_size_bytes .map(|i| i as usize), + read_rate_limit_per_sec: value.write_rate_limit_per_sec.map(|i| i as usize), + write_rate_limit_per_sec: value.write_rate_limit_per_sec.map(|i| i as usize), } } diff --git a/tests/openapi/test_strictmode.py b/tests/openapi/test_strictmode.py index a3974e7eac..01fd52e4c3 100644 --- a/tests/openapi/test_strictmode.py +++ b/tests/openapi/test_strictmode.py @@ -603,3 +603,84 @@ def test_strict_mode_max_collection_size_upsert_batch(collection_name): assert False, "Upserting should have failed but didn't" + +def test_strict_mode_read_rate_limiting(collection_name): + set_strict_mode(collection_name, { + "enabled": True, + "read_rate_limit_per_sec": 1, + }) + + response = request_with_validation( + api='/collections/{collection_name}', + method="GET", + path_params={'collection_name': collection_name}, + ) + + assert response.ok + new_strict_mode_config = response.json()['result']['config']['strict_mode_config'] + assert new_strict_mode_config['enabled'] + assert new_strict_mode_config['read_rate_limit_per_sec'] == 1 + + failed_count = 0 + + for _ in range(10): + response = request_with_validation( + api='/collections/{collection_name}/points/search', + method="POST", + path_params={'collection_name': collection_name}, + body={ + "vector": [0.2, 0.1, 0.9, 0.7], + "limit": 4 + } + ) + if not response.ok: + failed_count += 1 + assert response.status_code == 429 + assert "Rate limiting exceeded: Read rate limit exceeded, retry later" in response.json()['status']['error'] + + # loose check, as the rate limiting might not be exact + assert failed_count > 5, "Rate limiting did not work" + + +def test_strict_mode_write_rate_limiting(collection_name): + set_strict_mode(collection_name, { + "enabled": True, + "write_rate_limit_per_sec": 1, + }) + + response = request_with_validation( + api='/collections/{collection_name}', + method="GET", + path_params={'collection_name': collection_name}, + ) + + assert response.ok + new_strict_mode_config = response.json()['result']['config']['strict_mode_config'] + assert new_strict_mode_config['enabled'] + assert new_strict_mode_config['write_rate_limit_per_sec'] == 1 + + failed_count = 0 + + for _ in range(10): + response = request_with_validation( + api='/collections/{collection_name}/points', + method="PUT", + path_params={'collection_name': collection_name}, + query_params={'wait': 'true'}, + body={ + "points": [ + { + "id": 1, + "vector": [0.05, 0.61, 0.76, 0.74], + }, + ] + } + ) + + if not response.ok: + failed_count += 1 + assert response.status_code == 429 + assert "Rate limiting exceeded: Write rate limit exceeded, retry later" in response.json()['status']['error'] + + # loose check, as the rate limiting might not be exact + assert failed_count > 5, "Rate limiting did not work"