From 02e93d6149130de76385852ec016343f9b591307 Mon Sep 17 00:00:00 2001 From: Jacob Quinn Date: Tue, 22 Sep 2026 12:05:26 -0600 Subject: [PATCH] Contain UDF callback errors and clean aggregate state safely --- src/UDF.jl | 213 ++++++++++++++++++------------------- test/runtests.jl | 2 + test/udf_errors.jl | 260 +++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 366 insertions(+), 109 deletions(-) create mode 100644 test/udf_errors.jl diff --git a/src/UDF.jl b/src/UDF.jl index d32d62d..aee18df 100644 --- a/src/UDF.jl +++ b/src/UDF.jl @@ -54,18 +54,40 @@ mutable struct AggregateUDFData final::Function end +# A callback must report errors to SQLite, not unwind through its C caller. +function udf_error(context, err) + try + if err isa OutOfMemoryError + C.sqlite3_result_error_nomem(context) + else + message = sprint(showerror, err) + C.sqlite3_result_error(context, message, sizeof(message)) + end + catch + # Error formatting can itself invoke user-defined methods. + C.sqlite3_result_error( + context, + "Julia exception in SQLite callback", + -1, + ) + end + nothing +end + function wrap_scalarfunc( context::Ptr{Cvoid}, nargs::Cint, values::Ptr{Ptr{Cvoid}}, ) - udf_data = - unsafe_pointer_to_objref(C.sqlite3_user_data(context))::ScalarUDFData - func = udf_data.func - - args = [sqlvalue(values, i) for i in 1:nargs] - ret = func(args...) - sqlreturn(context, ret) + try + udf_data = unsafe_pointer_to_objref( + C.sqlite3_user_data(context), + )::ScalarUDFData + args = [sqlvalue(values, i) for i in 1:nargs] + sqlreturn(context, udf_data.func(args...)) + catch err + udf_error(context, err) + end nothing end @@ -82,120 +104,92 @@ function bytestoint(ptr::Ptr{UInt8}, start::Int, len::Int) return htol(s) end +# Transfer ownership to the callback before any operation that may throw. +# A cleared context also tells xFinal that a failed step has no state to finalize. +function take_aggregate_buffer!(acptr) + valsize = bytestoint(acptr, 1, sizeof(Int)) + valptr = + reinterpret(Ptr{UInt8}, bytestoint(acptr, sizeof(Int) + 1, sizeof(Ptr))) + unsafe_store!(Ptr{Int}(acptr), 0) + unsafe_store!(Ptr{Ptr{UInt8}}(acptr + sizeof(Int)), C_NULL) + return valsize, valptr +end + function wrap_stepfunc( context::Ptr{Cvoid}, nargs::Cint, values::Ptr{Ptr{Cvoid}}, ) - udf_data = - unsafe_pointer_to_objref(C.sqlite3_user_data(context))::AggregateUDFData - init = udf_data.init - func = udf_data.step - - args = [sqlvalue(values, i) for i in 1:nargs] - - intsize = sizeof(Int) - ptrsize = sizeof(Ptr) - acsize = intsize + ptrsize - acptr = convert(Ptr{UInt8}, C.sqlite3_aggregate_context(context, acsize)) - - # acptr will be zeroed-out if this is the first iteration - ret = ccall( - :memcmp, - Cint, - (Ptr{UInt8}, Ptr{UInt8}, Cuint), - zeros(UInt8, acsize), - acptr, - acsize, - ) - if ret == 0 - acval = init - valsize = 256 - # avoid the garbage collector using malloc - valptr = convert(Ptr{UInt8}, Libc.malloc(valsize)) - valptr == C_NULL && throw(SQLiteException("memory error")) - else - # size of serialized value is first sizeof(Int) bytes - valsize = bytestoint(acptr, 1, intsize) - # ptr to serialized value is last sizeof(Ptr) bytes - valptr = - reinterpret(Ptr{UInt8}, bytestoint(acptr, intsize + 1, ptrsize)) - # deserialize the value pointed to by valptr - acvalbuf = zeros(UInt8, valsize) - unsafe_copyto!(pointer(acvalbuf), valptr, valsize) - acval = sqldeserialize(acvalbuf) - end - - local funcret + valptr = Ptr{UInt8}(C_NULL) try - funcret = sqlserialize(func(acval, args...)) - catch - Libc.free(valptr) - rethrow() - end - - newsize = sizeof(funcret) - if newsize > valsize - # TODO: increase this in a cleverer way? - tmp = convert(Ptr{UInt8}, Libc.realloc(valptr, newsize)) - if tmp == C_NULL - Libc.free(valptr) - throw(SQLiteException("memory error")) + acptr = convert( + Ptr{UInt8}, + C.sqlite3_aggregate_context(context, sizeof(Int) + sizeof(Ptr)), + ) + acptr == C_NULL && throw(OutOfMemoryError()) + valsize, valptr = take_aggregate_buffer!(acptr) + + udf_data = unsafe_pointer_to_objref( + C.sqlite3_user_data(context), + )::AggregateUDFData + args = [sqlvalue(values, i) for i in 1:nargs] + if valptr == C_NULL + acval = udf_data.init + valsize = 256 + valptr = convert(Ptr{UInt8}, Libc.malloc(valsize)) + valptr == C_NULL && throw(OutOfMemoryError()) else + acvalbuf = zeros(UInt8, valsize) + unsafe_copyto!(pointer(acvalbuf), valptr, valsize) + acval = sqldeserialize(acvalbuf) + end + + funcret = sqlserialize(udf_data.step(acval, args...)) + newsize = sizeof(funcret) + if newsize > valsize + tmp = convert(Ptr{UInt8}, Libc.realloc(valptr, newsize)) + tmp == C_NULL && throw(OutOfMemoryError()) valptr = tmp end - end - # copy serialized return value - unsafe_copyto!(valptr, pointer(funcret), newsize) - - # copy the size of the serialized value - unsafe_copyto!(acptr, pointer(reinterpret(UInt8, [newsize])), intsize) - # copy the address of the pointer to the serialized value - valarr = reinterpret(UInt8, [valptr]) - for i in 1:length(valarr) - unsafe_store!(acptr, valarr[i], intsize + i) + GC.@preserve funcret unsafe_copyto!(valptr, pointer(funcret), newsize) + + # Publish the new state only after the step and serialization succeed. + unsafe_store!(Ptr{Int}(acptr), newsize) + unsafe_store!(Ptr{Ptr{UInt8}}(acptr + sizeof(Int)), valptr) + valptr = Ptr{UInt8}(C_NULL) + catch err + udf_error(context, err) + finally + Libc.free(valptr) end nothing end -function wrap_finalfunc( - context::Ptr{Cvoid}, - nargs::Cint, - values::Ptr{Ptr{Cvoid}}, -) - udf_data = - unsafe_pointer_to_objref(C.sqlite3_user_data(context))::AggregateUDFData - init = udf_data.init - func = udf_data.final - - acptr = convert(Ptr{UInt8}, C.sqlite3_aggregate_context(context, 0)) - - # step function wasn't run - if acptr == C_NULL - sqlreturn(context, init) - else - intsize = sizeof(Int) - ptrsize = sizeof(Ptr) - acsize = intsize + ptrsize - - # load size - valsize = bytestoint(acptr, 1, intsize) - # load ptr - valptr = - reinterpret(Ptr{UInt8}, bytestoint(acptr, intsize + 1, ptrsize)) - - # load value - acvalbuf = zeros(UInt8, valsize) - unsafe_copyto!(pointer(acvalbuf), valptr, valsize) - acval = sqldeserialize(acvalbuf) - - local ret - try - ret = func(acval) - finally - Libc.free(valptr) +function wrap_finalfunc(context::Ptr{Cvoid}) + valptr = Ptr{UInt8}(C_NULL) + try + acptr = convert(Ptr{UInt8}, C.sqlite3_aggregate_context(context, 0)) + if acptr != C_NULL + valsize, valptr = take_aggregate_buffer!(acptr) + # SQLite still calls xFinal after a failed xStep. Preserve that error. + valptr == C_NULL && return nothing + end + udf_data = unsafe_pointer_to_objref( + C.sqlite3_user_data(context), + )::AggregateUDFData + if acptr == C_NULL + # Preserve the initial value when no rows reached xStep. + sqlreturn(context, udf_data.init) + else + acvalbuf = zeros(UInt8, valsize) + unsafe_copyto!(pointer(acvalbuf), valptr, valsize) + acval = sqldeserialize(acvalbuf) + sqlreturn(context, udf_data.final(acval)) end - sqlreturn(context, ret) + catch err + udf_error(context, err) + finally + Libc.free(valptr) end nothing end @@ -217,7 +211,8 @@ UDF_keep_alive_list = [] SQLite.register(db, init, step_func, final_func; nargs=-1, name=string(step), isdeterm=true) Register a scalar (first method) or aggregate (second method) function -with a [`SQLite.DB`](@ref). +with a [`SQLite.DB`](@ref). Callback errors, including value conversion errors, +are reported as `SQLiteException`s by the query that invokes the function. """ function register( db::DB, @@ -274,7 +269,7 @@ function register( udf_data_ptr = pointer_from_objref(udf_data) cs = @cfunction(wrap_stepfunc, Cvoid, (Ptr{Cvoid}, Cint, Ptr{Ptr{Cvoid}})) - cf = @cfunction(wrap_finalfunc, Cvoid, (Ptr{Cvoid}, Cint, Ptr{Ptr{Cvoid}})) + cf = @cfunction(wrap_finalfunc, Cvoid, (Ptr{Cvoid},)) enc = C.SQLITE_UTF8 enc = isdeterm ? enc | C.SQLITE_DETERMINISTIC : enc diff --git a/test/runtests.jl b/test/runtests.jl index e1237be..f51861c 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -1167,6 +1167,8 @@ end end end + include("udf_errors.jl") + @testset "UDF value marshalling" begin # These tests pin down SQLite.sqlvalue, which unmarshals a C # sqlite3_value* into a Julia value on the UDF argument hot path. A diff --git a/test/udf_errors.jl b/test/udf_errors.jl new file mode 100644 index 0000000..8cd8d56 --- /dev/null +++ b/test/udf_errors.jl @@ -0,0 +1,260 @@ +struct UDFReturnError end +SQLite.sqlreturn(context, ::UDFReturnError) = error("return fixture") + +struct UDFSerializeError end +function SQLite.Serialization.serialize( + ::SQLite.Serialization.AbstractSerializer, + ::UDFSerializeError, +) + error("serialize fixture") +end + +struct UDFDeserializeError + value::Int +end +function SQLite.Serialization.deserialize( + ::SQLite.Serialization.AbstractSerializer, + ::Type{UDFDeserializeError}, +) + error("deserialize fixture") +end + +struct UDFShowError <: Exception end +Base.showerror(::IO, ::UDFShowError) = error("showerror fixture") + +@testset "UDF error boundaries" begin + db = SQLite.DB() + SQLite.execute(db, "CREATE TABLE input(x INTEGER)") + SQLite.execute(db, "INSERT INTO input VALUES (1), (2)") + function value(sql, params = ()) + q = DBInterface.execute(db, sql, params) + try + return first(q)[1] + finally + DBInterface.close!(q) + DBInterface.close!(q.stmt) + end + end + function fails(sql, message, params = ()) + stmt = DBInterface.prepare(db, sql) + try + err = try + DBInterface.execute(stmt, params) + nothing + catch e + e + end + @test err isa SQLite.SQLiteException + @test err isa SQLite.SQLiteException && occursin(message, err.msg) + finally + # Finalization must be safe even after a partially computed aggregate. + DBInterface.close!(stmt) + end + @test value("SELECT 7") == 7 + end + try + @testset "scalar function and result conversion" begin + SQLite.register( + db, + x -> error("scalar é🦆 fixture"); + nargs = 1, + name = "scalar_error", + ) + fails("SELECT scalar_error(x) FROM input", "scalar é🦆 fixture") + SQLite.register( + db, + x -> UDFReturnError(); + nargs = 1, + name = "return_error", + ) + fails("SELECT return_error(1)", "return fixture") + SQLite.register( + db, + x -> UDFSerializeError(); + nargs = 1, + name = "serialize_error", + ) + fails("SELECT serialize_error(1)", "serialize fixture") + SQLite.register( + db, + x -> throw(UDFShowError()); + nargs = 1, + name = "show_error", + ) + fails("SELECT show_error(1)", "Julia exception in SQLite callback") + end + @testset "argument conversion" begin + calls = Ref(0) + SQLite.register( + db, + x -> (calls[] += 1; x); + nargs = 1, + name = "argument_error", + ) + fails( + "SELECT argument_error(?)", + "Error deserializing", + (UDFDeserializeError(1),), + ) + @test calls[] == 0 + end + @testset "aggregate step and state conversion" begin + for failing_row in (1, 2) + final_calls = Ref(0) + SQLite.register( + db, + 0, + (state, x) -> + x == failing_row ? error("step fixture") : state + x, + state -> (final_calls[] += 1; state); + nargs = 1, + name = "step_error", + ) + fails("SELECT step_error(x) FROM input", "step fixture") + @test final_calls[] == 0 + end + SQLite.register( + db, + 0, + (state, x) -> x == 2 ? UDFSerializeError() : state + x; + nargs = 1, + name = "state_serialize_error", + ) + fails( + "SELECT state_serialize_error(x) FROM input", + "serialize fixture", + ) + SQLite.register( + db, + 0, + (state, x) -> UDFDeserializeError(x); + nargs = 1, + name = "state_deserialize_error", + ) + fails( + "SELECT state_deserialize_error(x) FROM input", + "Error deserializing", + ) + # Deserialization also occurs in xFinal when there was only one row. + fails( + "SELECT state_deserialize_error(x) FROM input WHERE x=1", + "Error deserializing", + ) + final_calls = Ref(0) + SQLite.register( + db, + 0, + (state, x) -> state + 1, + state -> (final_calls[] += 1; state); + nargs = 1, + name = "aggregate_argument_error", + ) + fails( + "SELECT aggregate_argument_error(?)", + "Error deserializing", + (UDFDeserializeError(1),), + ) + @test final_calls[] == 0 + end + @testset "aggregate final and empty input" begin + SQLite.register( + db, + 0, + +, + state -> error("final fixture"); + nargs = 1, + name = "final_error", + ) + fails("SELECT final_error(x) FROM input", "final fixture") + @test value("SELECT final_error(x) FROM input WHERE 0") == 0 + SQLite.register( + db, + 0, + +, + state -> UDFReturnError(); + nargs = 1, + name = "final_return_error", + ) + fails("SELECT final_return_error(x) FROM input", "return fixture") + SQLite.register( + db, + 0, + +, + state -> UDFSerializeError(); + nargs = 1, + name = "final_serialize_error", + ) + fails( + "SELECT final_serialize_error(x) FROM input", + "serialize fixture", + ) + SQLite.register( + db, + UDFReturnError(), + (state, x) -> state; + nargs = 1, + name = "empty_return_error", + ) + fails( + "SELECT empty_return_error(x) FROM input WHERE 0", + "return fixture", + ) + end + @testset "successful state growth and statement reuse" begin + SQLite.register( + db, + "", + (state, x) -> state * repeat(string(x), 1024); + nargs = 1, + name = "grow_state", + ) + @test value("SELECT grow_state(x) FROM input") == + repeat("1", 1024) * repeat("2", 1024) + fail = Ref(true) + SQLite.register( + db, + 0, + (state, x) -> + fail[] && x == 2 ? error("reuse fixture") : state + x; + nargs = 1, + name = "reusable", + ) + stmt = DBInterface.prepare(db, "SELECT reusable(x) FROM input") + try + @test_throws SQLite.SQLiteException DBInterface.execute(stmt) + fail[] = false + q = DBInterface.execute(stmt) + @test first(q)[1] == 3 + DBInterface.close!(q) + finally + DBInterface.close!(stmt) + end + end + finally + close(db) + end + @test !isopen(db) + + # Keep the failed prepared statement registered until the DB itself closes. + db = SQLite.DB() + stmt = nothing + try + SQLite.register( + db, + 0, + (state, x) -> x == 2 ? error("close fixture") : state + x; + nargs = 1, + name = "close_error", + ) + stmt = DBInterface.prepare( + db, + "SELECT close_error(x) FROM (SELECT 1 AS x UNION ALL SELECT 2)", + ) + @test_throws SQLite.SQLiteException DBInterface.execute(stmt) + @test SQLite.isready(stmt) + finally + close(db) + end + @test !isopen(db) + @test !SQLite.isready(stmt) +end