First implementation for enable_load_extension

TODO: Wait for the reply of the port driver
This commit is contained in:
booo
2012-07-24 16:14:33 +02:00
parent 08955394ab
commit 20298aad29
4 changed files with 33 additions and 2 deletions

View File

@@ -125,6 +125,16 @@ static ErlDrvData start(ErlDrvPort port, char* cmd) {
return (ErlDrvData) retval; 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 // Driver Stop
static void stop(ErlDrvData handle) { static void stop(ErlDrvData handle) {
sqlite3_drv_t* driver_data = (sqlite3_drv_t*) handle; sqlite3_drv_t* driver_data = (sqlite3_drv_t*) handle;
@@ -181,6 +191,9 @@ static int control(
case CMD_SQL_EXEC_SCRIPT: case CMD_SQL_EXEC_SCRIPT:
sql_exec_script(driver_data, buf, len); sql_exec_script(driver_data, buf, len);
break; break;
case CMD_ENABLE_LOAD_EXTENSION:
enable_load_extension(driver_data, buf, len);
break;
default: default:
unknown(driver_data, buf, len); unknown(driver_data, buf, len);
} }

View File

@@ -39,6 +39,7 @@
#define CMD_PREPARED_FINALIZE 10 #define CMD_PREPARED_FINALIZE 10
#define CMD_PREPARED_COLUMNS 11 #define CMD_PREPARED_COLUMNS 11
#define CMD_SQL_EXEC_SCRIPT 12 #define CMD_SQL_EXEC_SCRIPT 12
#define CMD_ENABLE_LOAD_EXTENSION 13
// Number of bytes for each key // Number of bytes for each key
// (160 bits for SHA1 hash) // (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 sql_free_async(void *async_command);
static void ready_async(ErlDrvData drv_data, ErlDrvThreadData thread_data); static void ready_async(ErlDrvData drv_data, ErlDrvThreadData thread_data);
static int unknown(sqlite3_drv_t *bdb_drv, char *buf, int len); 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) #if defined(_MSC_VER)
#pragma warning(default: 4201) #pragma warning(default: 4201)

View File

@@ -19,6 +19,7 @@
-export([open/1, open/2]). -export([open/1, open/2]).
-export([start_link/1, start_link/2]). -export([start_link/1, start_link/2]).
-export([stop/0, close/1, close_timeout/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, -export([sql_exec/1, sql_exec/2, sql_exec_timeout/3,
sql_exec_script/2, sql_exec_script_timeout/3, sql_exec_script/2, sql_exec_script_timeout/3,
sql_exec/3, sql_exec_timeout/4]). sql_exec/3, sql_exec_timeout/4]).
@@ -152,6 +153,9 @@ close_timeout(Db, Timeout) ->
stop() -> stop() ->
close(?MODULE). close(?MODULE).
enable_load_extension(Db, Value) ->
gen_server:call(Db, {enable_load_extension, Value}).
%%-------------------------------------------------------------------- %%--------------------------------------------------------------------
%% @doc %% @doc
%% Executes the Sql statement directly. %% Executes the Sql statement directly.
@@ -913,6 +917,10 @@ handle_call({finalize, Ref}, _From, State = #state{port = Port, refs = Refs}) ->
NewState = State NewState = State
end, end,
{reply, Reply, NewState}; {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}) -> handle_call({Cmd, Ref}, _From, State = #state{port = Port, refs = Refs}) ->
Reply = case dict:find(Ref, Refs) of Reply = case dict:find(Ref, Refs) of
{ok, Index} -> {ok, Index} ->
@@ -1011,6 +1019,7 @@ get_priv_dir() ->
-define(PREPARED_FINALIZE, 10). -define(PREPARED_FINALIZE, 10).
-define(PREPARED_COLUMNS, 11). -define(PREPARED_COLUMNS, 11).
-define(SQL_EXEC_SCRIPT, 12). -define(SQL_EXEC_SCRIPT, 12).
-define(ENABLE_LOAD_EXTENSION, 13).
create_port_cmd(DbFile) -> create_port_cmd(DbFile) ->
atom_to_list(?DRIVER_NAME) ++ " " ++ DbFile. atom_to_list(?DRIVER_NAME) ++ " " ++ DbFile.
@@ -1052,6 +1061,10 @@ exec(Port, {bind, Index, Params}) ->
Bin = term_to_binary({Index, Params}), Bin = term_to_binary({Index, Params}),
port_control(Port, ?PREPARED_BIND, Bin), port_control(Port, ?PREPARED_BIND, Bin),
wait_result(Port); 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) -> exec(Port, {Cmd, Index}) when is_integer(Index) ->
CmdCode = case Cmd of CmdCode = case Cmd of
next -> ?PREPARED_STEP; next -> ?PREPARED_STEP;

View File

@@ -50,7 +50,8 @@ all_test_() ->
?FuncTest(large_number), ?FuncTest(large_number),
?FuncTest(unicode), ?FuncTest(unicode),
?FuncTest(acc_string_encoding), ?FuncTest(acc_string_encoding),
?FuncTest(large_offset)]}. ?FuncTest(large_offset),
?FuncTest(enable_load_extension)]}.
open_db() -> open_db() ->
sqlite3:open(ct, [in_memory]). sqlite3:open(ct, [in_memory]).
@@ -303,6 +304,9 @@ large_offset() ->
[{columns, ["id"]}, {rows, []}, {error, 20, "datatype mismatch"}], [{columns, ["id"]}, {rows, []}, {error, 20, "datatype mismatch"}],
sqlite3:sql_exec(ct, "select * from large_offset limit 1 offset 9223372036854775808")). 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 % create, read, update, delete
%%==================================================================== %%====================================================================
%% Internal functions %% Internal functions