From 21bf859c0b6c37872214debda0b72c92962bafdd Mon Sep 17 00:00:00 2001 From: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 06:30:38 +0000 Subject: [PATCH] fix(python): validate compaction option bounds --- python/python/tests/test_table.py | 22 +++++++++++++ python/src/table.rs | 54 ++++++++++++++++++++++++++++--- 2 files changed, 71 insertions(+), 5 deletions(-) diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index 5fb5ce9c9..d764b9165 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -3404,6 +3404,28 @@ async def test_optimize_compaction_options(mem_db_async: AsyncConnection): await table.optimize(compaction_options={"unknown": 1}) +@pytest.mark.parametrize( + ("option", "value", "message"), + [ + ("target_rows_per_fragment", 0, "must be between 1 and 4294967295"), + ("max_rows_per_group", 0, "must be between 1 and 4294967295"), + ("batch_size", 0, "must be between 1 and 4294967295"), + ("num_threads", 0, "must be greater than 0"), + ("target_rows_per_fragment", 2**32, "must be between 1 and 4294967295"), + ("max_rows_per_group", 2**32, "must be between 1 and 4294967295"), + ("batch_size", 2**32, "must be between 1 and 4294967295"), + ("io_buffer_size", 2**63, "must be at most 9223372036854775807"), + ], +) +@pytest.mark.asyncio +async def test_optimize_compaction_options_validation( + mem_db_async: AsyncConnection, option: str, value: int, message: str +): + table = await mem_db_async.create_table("test", data=[{"x": 1}]) + with pytest.raises(ValueError, match=message): + await table.optimize(compaction_options={option: value}) + + @pytest.mark.asyncio async def test_optimize_delete_unverified(tmp_db_async: AsyncConnection, tmp_path): table = await tmp_db_async.create_table( diff --git a/python/src/table.rs b/python/src/table.rs index 963183dea..a3555bdb6 100644 --- a/python/src/table.rs +++ b/python/src/table.rs @@ -40,6 +40,48 @@ enum PredicateArg { Sql(String), } +fn validate_positive_u32(value: u64, name: &str) -> PyResult { + if !(1..=u32::MAX as u64).contains(&value) { + return Err(PyValueError::new_err(format!( + "{name} must be between 1 and {}", + u32::MAX + ))); + } + Ok(value as usize) +} + +fn positive_u32(value: &Bound<'_, PyAny>, name: &str) -> PyResult { + validate_positive_u32(value.extract()?, name) +} + +fn optional_positive_u32(value: &Bound<'_, PyAny>, name: &str) -> PyResult> { + value + .extract::>()? + .map(|value| validate_positive_u32(value, name)) + .transpose() +} + +fn optional_positive_usize(value: &Bound<'_, PyAny>, name: &str) -> PyResult> { + let value: Option = value.extract()?; + if value == Some(0) { + return Err(PyValueError::new_err(format!( + "{name} must be greater than 0" + ))); + } + Ok(value) +} + +fn optional_i64_bounded_u64(value: &Bound<'_, PyAny>, name: &str) -> PyResult> { + let value: Option = value.extract()?; + if value.is_some_and(|value| value > i64::MAX as u64) { + return Err(PyValueError::new_err(format!( + "{name} must be at most {}", + i64::MAX + ))); + } + Ok(value) +} + fn parse_compaction_options(options: Option<&Bound<'_, PyDict>>) -> PyResult { let mut parsed = CompactionOptions::default(); let Some(options) = options else { @@ -49,16 +91,18 @@ fn parse_compaction_options(options: Option<&Bound<'_, PyDict>>) -> PyResult parsed.target_rows_per_fragment = value.extract()?, - "max_rows_per_group" => parsed.max_rows_per_group = value.extract()?, + "target_rows_per_fragment" => { + parsed.target_rows_per_fragment = positive_u32(&value, &key)? + } + "max_rows_per_group" => parsed.max_rows_per_group = positive_u32(&value, &key)?, "max_bytes_per_file" => parsed.max_bytes_per_file = value.extract()?, "materialize_deletions" => parsed.materialize_deletions = value.extract()?, "materialize_deletions_threshold" => { parsed.materialize_deletions_threshold = value.extract()? } - "num_threads" => parsed.num_threads = value.extract()?, - "batch_size" => parsed.batch_size = value.extract()?, - "io_buffer_size" => parsed.io_buffer_size = value.extract()?, + "num_threads" => parsed.num_threads = optional_positive_usize(&value, &key)?, + "batch_size" => parsed.batch_size = optional_positive_u32(&value, &key)?, + "io_buffer_size" => parsed.io_buffer_size = optional_i64_bounded_u64(&value, &key)?, "defer_index_remap" => parsed.defer_index_remap = value.extract()?, "index_remap_mode" => { let mode: String = value.extract()?;