diff --git a/c_src/sqlite3_drv.c b/c_src/sqlite3_drv.c index 2aa133c..733c051 100644 --- a/c_src/sqlite3_drv.c +++ b/c_src/sqlite3_drv.c @@ -1,5 +1,6 @@ #include "sqlite3_drv.h" #include "ei.h" +#include "assert.h" // MSVC needs "__inline" instead of "inline" in C-source files. #if defined(_MSC_VER) @@ -107,6 +108,9 @@ static int control(ErlDrvData drv_data, unsigned int command, char *buf, case CMD_SQL_EXEC: sql_exec(driver_data, buf, len); break; + case CMD_SQL_BIND_AND_EXEC: + sql_bind_and_exec(driver_data, buf, len); + break; default: unknown(driver_data, buf, len); } @@ -129,23 +133,19 @@ static inline int return_error(sqlite3_drv_t *drv, const char *error, return 0; } -static int sql_exec(sqlite3_drv_t *drv, char *command, int command_size) { - int result, next_row; - char *rest = NULL; - sqlite3_stmt *statement; +static inline int output_error(sqlite3_drv_t *drv, const char *error) { + ErlDrvTermData *dataset; + int term_count; + return_error(drv, error, &dataset, &term_count); + driver_output_term(drv->port, dataset, term_count); + return 0; +} - // fprintf(drv->log, "Preexec: %.*s\n", command_size, command); - // fflush(drv->log); - result = sqlite3_prepare_v2(drv->db, command, command_size, &statement, - (const char **) &rest); - if (result != SQLITE_OK) { - ErlDrvTermData *dataset; - int term_count; - return_error(drv, sqlite3_errmsg(drv->db), &dataset, &term_count); - driver_output_term(drv->port, dataset, term_count); - return 0; - } +static inline int output_db_error(sqlite3_drv_t *drv) { + return output_error(drv, sqlite3_errmsg(drv->db)); +} +static inline int sql_exec_statement(sqlite3_drv_t *drv, sqlite3_stmt *statement) { async_sqlite3_command *async_command = (async_sqlite3_command *) calloc(1, sizeof(async_sqlite3_command)); async_command->driver_data = drv; @@ -157,7 +157,7 @@ static int sql_exec(sqlite3_drv_t *drv, char *command, int command_size) { if (sqlite3_threadsafe()) { drv->async_handle = driver_async(drv->port, &drv->key, sql_exec_async, - async_command, sql_free_async); + async_command, sql_free_async); } else { sql_exec_async(async_command); ready_async((ErlDrvData) drv, (ErlDrvThreadData) async_command); @@ -166,6 +166,193 @@ static int sql_exec(sqlite3_drv_t *drv, char *command, int command_size) { return 0; } +static int sql_exec(sqlite3_drv_t *drv, char *command, int command_size) { + int result; + char *rest = NULL; + sqlite3_stmt *statement; + + // fprintf(drv->log, "Preexec: %.*s\n", command_size, command); + // fflush(drv->log); + result = sqlite3_prepare_v2(drv->db, command, command_size, &statement, + (const char **) &rest); + if (result != SQLITE_OK) { + return output_db_error(drv); + } + + return sql_exec_statement(drv, statement); +} + +static inline int decode_and_bind_param( + sqlite3_drv_t *drv, char *buffer, int *index, sqlite3_stmt *statement, int param_index, int *type, int *size) { + int result; + ei_get_type(buffer, index, type, size); + long long_val; + sqlite3_int64 int64_val; + double double_val; + char* char_buf_val; + switch (*type) { + case ERL_SMALL_INTEGER_EXT: + ei_decode_long(buffer, index, &long_val); + result = sqlite3_bind_int(statement, param_index, long_val); + break; + case ERL_INTEGER_EXT: + ei_decode_longlong(buffer, index, &int64_val); + result = sqlite3_bind_int64(statement, param_index, int64_val); + break; + case ERL_FLOAT_EXT: + case NEW_FLOAT_EXT: // what's the difference? + ei_decode_double(buffer, index, &double_val); + result = sqlite3_bind_double(statement, param_index, double_val); + break; + case ERL_ATOM_EXT: + char_buf_val = malloc((*size + 1) * sizeof(char)); + ei_decode_atom(buffer, index, char_buf_val); + if (strncmp(char_buf_val, "null", 5) == 0) { + result = sqlite3_bind_null(statement, param_index); + } + else { + output_error(drv, "Non-null atom as parameter"); + return 1; + } + break; + case ERL_STRING_EXT: + char_buf_val = malloc((*size + 1) * sizeof(char)); // space for null separator + ei_decode_string(buffer, index, char_buf_val); + result = sqlite3_bind_text(statement, param_index, char_buf_val, *size, &free); + break; + case ERL_BINARY_EXT: + char_buf_val = malloc(*size * sizeof(char)); + ei_decode_binary(buffer, index, char_buf_val, size); + result = sqlite3_bind_text(statement, param_index, char_buf_val, *size, &free); + break; + case ERL_SMALL_TUPLE_EXT: + // assume this is {blob, Blob} + ei_get_type(buffer, index, type, size); + ei_decode_tuple_header(buffer, index, size); + assert (*size == 2); + ei_skip_term(buffer, index); // skipped the atom 'blob' + ei_get_type(buffer, index, type, size); + assert (*type == ERL_BINARY_EXT); + char_buf_val = malloc(*size * sizeof(char)); + ei_decode_binary(buffer, index, char_buf_val, size); + result = sqlite3_bind_blob(statement, param_index, char_buf_val, *size, &free); + break; + default: + output_error(drv, "bad parameter type"); + return 1; + } + if (result != SQLITE_OK) { + output_db_error(drv); + return result; + } + return SQLITE_OK; +} + +static int sql_bind_and_exec(sqlite3_drv_t *drv, char *buffer, int buffer_size) { + int result; + int index = 0; + int type, size; + char *rest = NULL; + sqlite3_stmt *statement; + + // fprintf(drv->log, "Preexec: %.*s\n", command_size, command); + // fflush(drv->log); + + ei_decode_version(buffer, &index, NULL); + result = ei_decode_tuple_header(buffer, &index, &size); + if (size != 2) { + return output_error(drv, "bad argument"); + } + + // decode SQL statement + ei_get_type(buffer, &index, &type, &size); + if (type != ERL_BINARY_EXT) { + return output_error(drv, "bad argument"); + } + + char *command = malloc(size * sizeof(char)); + ei_decode_binary(buffer, &index, command, &size); + // printf("size: %d, command: %.*s\n", size, size, command); + result = sqlite3_prepare_v2(drv->db, command, size, &statement, + (const char **) &rest); + free(command); + + if (result != SQLITE_OK) { + return output_db_error(drv); + } + + // decoding parameters + int i, cur_list_size = -1, param_index = 1, param_indices_are_explicit = 0; + char param_name[MAXATOMLEN + 1]; // parameter names shouldn't be longer than 256! + while (index < buffer_size) { + ei_decode_list_header(buffer, &index, &cur_list_size); + // note the finish condition; the last element is the tail and we shouldn't decode it! + for (i = 0; i < cur_list_size; i++) { + ei_get_type(buffer, &index, &type, &size); + if (type == ERL_SMALL_TUPLE_EXT) { + int old_index = index; + // param with name or explicit index + param_indices_are_explicit = 1; + if (size != 2) { + return output_error(drv, "bad argument"); + } + ei_decode_tuple_header(buffer, &index, &size); + ei_get_type(buffer, &index, &type, &size); + // first element of tuple is int (index), atom, or string (name) + switch (type) { + case ERL_SMALL_INTEGER_EXT: + ei_decode_long(buffer, &index, ¶m_index); + break; + case ERL_ATOM_EXT: + ei_decode_atom(buffer, &index, param_name); + // insert zero terminator + param_name[size] = '\0'; + if (strncmp(param_name, "blob", 5) == 0) { + // this isn't really a parameter name! + index = old_index; + param_indices_are_explicit = 0; + goto IMPLICIT_INDEX; // yuck + } + else { + param_index = sqlite3_bind_parameter_index(statement, param_name); + } + break; + case ERL_STRING_EXT: + if (size >= MAXATOMLEN) { + return output_error(drv, "parameter name too long"); + } + ei_decode_string(buffer, &index, param_name); + // insert zero terminator + param_name[size] = '\0'; + param_index = sqlite3_bind_parameter_index(statement, param_name); + break; + default: + return output_error(drv, "parameter index must be given as integer, atom, or string"); + } + result = decode_and_bind_param(drv, buffer, &index, statement, param_index, &type, &size); + if (result != SQLITE_OK) { + return result; // error has already been output + } + } + else { + IMPLICIT_INDEX: + if (param_indices_are_explicit) { + return output_error(drv, + "parameters without indices shouldn't follow indexed or named parameters"); + } + + result = decode_and_bind_param(drv, buffer, &index, statement, param_index, &type, &size); + if (result != SQLITE_OK) { + return result; // error has already been output + } + ++param_index; + } + } + } + + return sql_exec_statement(drv, statement); +} + static void sql_free_async(void *_async_command) { int i; async_sqlite3_command *async_command = diff --git a/c_src/sqlite3_drv.h b/c_src/sqlite3_drv.h index 5991e06..39fd003 100644 --- a/c_src/sqlite3_drv.h +++ b/c_src/sqlite3_drv.h @@ -18,6 +18,7 @@ // Binary commands between Erlang VM and Driver #define CMD_SQL_EXEC 2 // #define CMD_DEL 3 +#define CMD_SQL_BIND_AND_EXEC 4 // Number of bytes for each key // (160 bits for SHA1 hash) @@ -61,6 +62,7 @@ static void stop(ErlDrvData handle); static int control(ErlDrvData drv_data, unsigned int command, char *buf, int len, char **rbuf, int rlen); static int sql_exec(sqlite3_drv_t *drv, char *buf, int len); +static int sql_bind_and_exec(sqlite3_drv_t *drv, char *buf, int len); 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); diff --git a/src/sqlite3.erl b/src/sqlite3.erl index 215b15f..dca5e55 100644 --- a/src/sqlite3.erl +++ b/src/sqlite3.erl @@ -733,7 +733,6 @@ exec(Port, {sql_exec, Cmd}) -> exec(Port, {sql_bind_and_exec, SQL, Params}) -> Bin = term_to_binary({iolist_to_binary(SQL), Params}), port_control(Port, ?SQL_BIND_AND_EXEC_COMMAND, Bin), - io:format(user, "Sending sql_bind_and_exec to port", []), wait_result(Port). wait_result(Port) -> diff --git a/test/sqlite3_test.erl b/test/sqlite3_test.erl index 2566477..85e7cc0 100644 --- a/test/sqlite3_test.erl +++ b/test/sqlite3_test.erl @@ -107,9 +107,21 @@ parametrized() -> drop_table_if_exists(ct, user1), sqlite3:create_table(ct, user1, [{id, integer}, {name, text}]), sqlite3:sql_exec(ct, "INSERT INTO user1 (id, name) VALUES (?, ?)", [{1, 1}, {2, "john"}]), + sqlite3:sql_exec(ct, "INSERT INTO user1 (id, name) VALUES (?3, ?5)", [{3, 2}, {5, "joe"}]), + sqlite3:sql_exec(ct, "INSERT INTO user1 (id, name) VALUES (:id, @name)", [{":id", 3}, {'@name', <<"jack">>}]), + sqlite3:sql_exec(ct, "INSERT INTO user1 (id, name) VALUES (?, ?)", [4, "james"]), ?assertEqual( - [{columns, ["id", "name"]}, {rows, [{1, <<"john">>}]}], - sqlite3:sql_exec(ct, "select * from user1;")). + [{columns, ["id", "name"]}, + {rows, [{1, <<"john">>}, {2, <<"joe">>}, {3, <<"jack">>}, {4, <<"james">>}]}], + sqlite3:read_all(ct, user1)), + sqlite3:drop_table(ct, user1), + sqlite3:create_table(ct, user1, [{i, integer}, {d, double}, {b, blob}]), + sqlite3:sql_exec(ct, "INSERT INTO user1 (i, d, b) VALUES (?, ?, ?)", + [null, 1.0, {blob, <<1,0,0>>}]), + ?assertEqual( + [{columns, ["i", "d", "b"]}, + {rows, [{null, 1.0, <<1,0,0>>}]}], + sqlite3:read_all(ct, user1)). blob() -> drop_table_if_exists(ct, blobs), @@ -148,21 +160,21 @@ select_many_records() -> N, length(rows(sqlite3:sql_exec(ct, "select * from many_records;")))). -%% %% note that inserts are actually serialized by gen_server -%% concurrent_inserts_test() -> -%% N = 1024, -%% sqlite3:open(concurrent, [in_memory]), %% doing this test not in memory is much slower! -%% drop_table_if_exists(concurrent, t), -%% sqlite3:create_table(concurrent, t, [{id0, integer}]), -%% Self = self(), -%% [spawn(fun () -> -%% sqlite3:write(concurrent, t, [{id0, X}]), -%% Self ! {finished, N} -%% end) || X <- lists:seq(1, N)], -%% loop_concurrent_inserts(N), -%% ?assertEqual( -%% N, length(rows(sqlite3:read_all(concurrent, t)))), -%% sqlite3:close(concurrent). +%% note that inserts are actually serialized by gen_server +concurrent_inserts_test() -> + N = 1024, + sqlite3:open(concurrent, [in_memory]), %% doing this test not in memory is much slower! + drop_table_if_exists(concurrent, t), + sqlite3:create_table(concurrent, t, [{id0, integer}]), + Self = self(), + [spawn(fun () -> + sqlite3:write(concurrent, t, [{id0, X}]), + Self ! {finished, N} + end) || X <- lists:seq(1, N)], + loop_concurrent_inserts(N), + ?assertEqual( + N, length(rows(sqlite3:read_all(concurrent, t)))), + sqlite3:close(concurrent). loop_concurrent_inserts(0) -> ok;