Skip to content

Commit d3f84f2

Browse files
committed
sqlite3: validate narg before creating functions/aggregates
CPython 3.14 added check_num_params() which raises ProgrammingError when narg/n_arg/num_params is out of range (-1..=SQLITE_LIMIT_FUNCTION_ARG). RustPython passed invalid values directly to SQLite, resulting in an OperationalError instead. Add check_num_params() helper and call it in create_function(), create_aggregate(), and create_window_function(). Assisted-by: GitHub Copilot:claude-sonnet-4-6
1 parent 2274cef commit d3f84f2

2 files changed

Lines changed: 19 additions & 3 deletions

File tree

Lib/test/test_sqlite3/test_userfunctions.py

Lines changed: 0 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -170,7 +170,6 @@ def setUp(self):
170170
def tearDown(self):
171171
self.con.close()
172172

173-
@unittest.expectedFailure # TODO: RUSTPYTHON; error message differs for invalid num args
174173
def test_func_error_on_create(self):
175174
with self.assertRaisesRegex(sqlite.ProgrammingError, "not -100"):
176175
self.con.create_function("bla", -100, lambda x: 2*x)
@@ -514,7 +513,6 @@ def test_win_sum_int(self):
514513
self.cur.execute(self.query % "sumint")
515514
self.assertEqual(self.cur.fetchall(), self.expected)
516515

517-
@unittest.expectedFailure # TODO: RUSTPYTHON; error message differs for invalid num args
518516
def test_win_error_on_create(self):
519517
with self.assertRaisesRegex(sqlite.ProgrammingError, "not -100"):
520518
self.con.create_window_function("shouldfail", -100, WindowSumInt)
@@ -649,7 +647,6 @@ def setUp(self):
649647
def tearDown(self):
650648
self.con.close()
651649

652-
@unittest.expectedFailure # TODO: RUSTPYTHON; error message differs for invalid num args
653650
def test_aggr_error_on_create(self):
654651
with self.assertRaisesRegex(sqlite.ProgrammingError, "not -100"):
655652
self.con.create_function("bla", -100, AggrSum)

crates/stdlib/src/_sqlite3.rs

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1279,6 +1279,7 @@ mod _sqlite3 {
12791279
SQLITE_UTF8
12801280
};
12811281
let db = self.db_lock(vm)?;
1282+
check_num_params(&db, args.narg, "narg", vm)?;
12821283
let Some(data) = CallbackData::new(args.func, vm) else {
12831284
return db.create_function(
12841285
name.as_ptr(),
@@ -1310,6 +1311,7 @@ mod _sqlite3 {
13101311
fn create_aggregate(&self, args: CreateAggregateArgs, vm: &VirtualMachine) -> PyResult<()> {
13111312
let name = args.name.to_cstring(vm)?;
13121313
let db = self.db_lock(vm)?;
1314+
check_num_params(&db, args.narg, "n_arg", vm)?;
13131315
let Some(data) = CallbackData::new(args.aggregate_class, vm) else {
13141316
return db.create_function(
13151317
name.as_ptr(),
@@ -1392,6 +1394,7 @@ mod _sqlite3 {
13921394
) -> PyResult<()> {
13931395
let name = name.to_cstring(vm)?;
13941396
let db = self.db_lock(vm)?;
1397+
check_num_params(&db, narg, "num_params", vm)?;
13951398
let Some(data) = CallbackData::new(aggregate_class, vm) else {
13961399
unsafe {
13971400
sqlite3_create_window_function(
@@ -3475,6 +3478,22 @@ mod _sqlite3 {
34753478
Ok(obj)
34763479
}
34773480

3481+
fn check_num_params(
3482+
db: &Sqlite,
3483+
n: c_int,
3484+
param_name: &str,
3485+
vm: &VirtualMachine,
3486+
) -> PyResult<()> {
3487+
let limit = unsafe { sqlite3_limit(db.db, SQLITE_LIMIT_FUNCTION_ARG, -1) };
3488+
if n < -1 || n > limit {
3489+
return Err(new_programming_error(
3490+
vm,
3491+
format!("'{param_name}' must be between -1 and {limit}, not {n}"),
3492+
));
3493+
}
3494+
Ok(())
3495+
}
3496+
34783497
fn ptr_to_str<'a>(p: *const libc::c_char, vm: &VirtualMachine) -> PyResult<&'a str> {
34793498
if p.is_null() {
34803499
return Err(vm.new_memory_error("string pointer is null"));

0 commit comments

Comments
 (0)