From 20298aad29ef3499a72fae81a5502817fe65049d Mon Sep 17 00:00:00 2001 From: booo Date: Tue, 24 Jul 2012 16:14:33 +0200 Subject: [PATCH] First implementation for enable_load_extension TODO: Wait for the reply of the port driver --- c_src/sqlite3_drv.c | 13 +++++++++++++ c_src/sqlite3_drv.h | 3 ++- src/sqlite3.erl | 13 +++++++++++++ test/sqlite3_test.erl | 6 +++++- 4 files changed, 33 insertions(+), 2 deletions(-) diff --git a/c_src/sqlite3_drv.c b/c_src/sqlite3_drv.c index a7f4a34..0c05743 100644 --- a/c_src/sqlite3_drv.c +++ b/c_src/sqlite3_drv.c @@ -125,6 +125,16 @@ static ErlDrvData start(ErlDrvPort port, char* cmd) { return (ErlDrvData) retval; } +static int enable_load_extension(sqlite3_drv_t* drv, char *buf, + int len) { + if (buf[1]) { + sqlite3_enable_load_extension(drv->db, 1); + } else { + sqlite3_enable_load_extension(drv->db, 0); + } + return 0; +} + // Driver Stop static void stop(ErlDrvData handle) { sqlite3_drv_t* driver_data = (sqlite3_drv_t*) handle; @@ -181,6 +191,9 @@ static int control( case CMD_SQL_EXEC_SCRIPT: sql_exec_script(driver_data, buf, len); break; + case CMD_ENABLE_LOAD_EXTENSION: + enable_load_extension(driver_data, buf, len); + break; default: unknown(driver_data, buf, len); } diff --git a/c_src/sqlite3_drv.h b/c_src/sqlite3_drv.h index 3b21c1b..103750e 100644 --- a/c_src/sqlite3_drv.h +++ b/c_src/sqlite3_drv.h @@ -39,6 +39,7 @@ #define CMD_PREPARED_FINALIZE 10 #define CMD_PREPARED_COLUMNS 11 #define CMD_SQL_EXEC_SCRIPT 12 +#define CMD_ENABLE_LOAD_EXTENSION 13 // Number of bytes for each key // (160 bits for SHA1 hash) @@ -111,7 +112,7 @@ static void sql_exec_async(void *async_command); static void sql_free_async(void *async_command); static void ready_async(ErlDrvData drv_data, ErlDrvThreadData thread_data); static int unknown(sqlite3_drv_t *bdb_drv, char *buf, int len); - +static int enable_load_extension(sqlite3_drv_t *drv, char *buf, int len); #if defined(_MSC_VER) #pragma warning(default: 4201) diff --git a/src/sqlite3.erl b/src/sqlite3.erl index fb80b74..dabb2b1 100644 --- a/src/sqlite3.erl +++ b/src/sqlite3.erl @@ -19,6 +19,7 @@ -export([open/1, open/2]). -export([start_link/1, start_link/2]). -export([stop/0, close/1, close_timeout/2]). +-export([enable_load_extension/2]). -export([sql_exec/1, sql_exec/2, sql_exec_timeout/3, sql_exec_script/2, sql_exec_script_timeout/3, sql_exec/3, sql_exec_timeout/4]). @@ -152,6 +153,9 @@ close_timeout(Db, Timeout) -> stop() -> close(?MODULE). +enable_load_extension(Db, Value) -> + gen_server:call(Db, {enable_load_extension, Value}). + %%-------------------------------------------------------------------- %% @doc %% Executes the Sql statement directly. @@ -913,6 +917,10 @@ handle_call({finalize, Ref}, _From, State = #state{port = Port, refs = Refs}) -> NewState = State end, {reply, Reply, NewState}; +handle_call({enable_load_extension, _Value} = Payload, _From, State = #state{port = + Port, refs = _Refs}) -> + Reply = exec(Port, Payload), + {reply, Reply, State}; handle_call({Cmd, Ref}, _From, State = #state{port = Port, refs = Refs}) -> Reply = case dict:find(Ref, Refs) of {ok, Index} -> @@ -1011,6 +1019,7 @@ get_priv_dir() -> -define(PREPARED_FINALIZE, 10). -define(PREPARED_COLUMNS, 11). -define(SQL_EXEC_SCRIPT, 12). +-define(ENABLE_LOAD_EXTENSION, 13). create_port_cmd(DbFile) -> atom_to_list(?DRIVER_NAME) ++ " " ++ DbFile. @@ -1052,6 +1061,10 @@ exec(Port, {bind, Index, Params}) -> Bin = term_to_binary({Index, Params}), port_control(Port, ?PREPARED_BIND, Bin), wait_result(Port); +exec(Port, {enable_load_extension, Value}) -> + %TODO wait for the return of the port driver + port_control(Port, ?ENABLE_LOAD_EXTENSION, [1, Value]), + ok; exec(Port, {Cmd, Index}) when is_integer(Index) -> CmdCode = case Cmd of next -> ?PREPARED_STEP; diff --git a/test/sqlite3_test.erl b/test/sqlite3_test.erl index b0cd89b..ccb8bc7 100644 --- a/test/sqlite3_test.erl +++ b/test/sqlite3_test.erl @@ -50,7 +50,8 @@ all_test_() -> ?FuncTest(large_number), ?FuncTest(unicode), ?FuncTest(acc_string_encoding), - ?FuncTest(large_offset)]}. + ?FuncTest(large_offset), + ?FuncTest(enable_load_extension)]}. open_db() -> sqlite3:open(ct, [in_memory]). @@ -303,6 +304,9 @@ large_offset() -> [{columns, ["id"]}, {rows, []}, {error, 20, "datatype mismatch"}], sqlite3:sql_exec(ct, "select * from large_offset limit 1 offset 9223372036854775808")). +enable_load_extension() -> + ?assertEqual(ok, sqlite3:enable_load_extension(ct, 1)). + % create, read, update, delete %%==================================================================== %% Internal functions