diff --git a/c_src/sqlite3_drv.c b/c_src/sqlite3_drv.c index f22d661..48bcb20 100644 --- a/c_src/sqlite3_drv.c +++ b/c_src/sqlite3_drv.c @@ -147,6 +147,9 @@ static int control( case CMD_PREPARED_FINALIZE: prepared_finalize(driver_data, buf, len); break; + case CMD_PREPARED_COLUMNS: + prepared_columns(driver_data, buf, len); + break; default: unknown(driver_data, buf, len); } @@ -396,14 +399,13 @@ static int bind_parameters( } static void get_columns( - sqlite3_drv_t *drv, sqlite3_stmt *statement, int column_count, + sqlite3_drv_t *drv, sqlite3_stmt *statement, int column_count, int base, int *p_term_count, int *p_term_allocated, ErlDrvTermData **p_dataset) { int i; - int base = *p_term_count; - *p_term_count += column_count * 3 + 1 + 2 + 2; + *p_term_count += column_count * 3 + 3; if (*p_term_count > *p_term_allocated) { - *p_term_allocated = max(*p_term_count, *p_term_allocated*2); + *p_term_allocated = max(*p_term_count, (*p_term_allocated)*2); *p_dataset = driver_realloc(*p_dataset, sizeof(ErlDrvTermData) * *p_term_allocated); } for (i = 0; i < column_count; i++) { @@ -413,13 +415,13 @@ static void get_columns( fflush(drv->log); #endif - (*p_dataset)[base + 2 + (i * 3)] = ERL_DRV_STRING; - (*p_dataset)[base + 2 + (i * 3) + 1] = (ErlDrvTermData) column_name; - (*p_dataset)[base + 2 + (i * 3) + 2] = strlen(column_name); + (*p_dataset)[base + (i * 3)] = ERL_DRV_STRING; + (*p_dataset)[base + (i * 3) + 1] = (ErlDrvTermData) column_name; + (*p_dataset)[base + (i * 3) + 2] = strlen(column_name); } - (*p_dataset)[base + 2 + column_count * 3 + 0] = ERL_DRV_NIL; - (*p_dataset)[base + 2 + column_count * 3 + 1] = ERL_DRV_LIST; - (*p_dataset)[base + 2 + column_count * 3 + 2] = column_count + 1; + (*p_dataset)[base + column_count * 3 + 0] = ERL_DRV_NIL; + (*p_dataset)[base + column_count * 3 + 1] = ERL_DRV_LIST; + (*p_dataset)[base + column_count * 3 + 2] = column_count + 1; } static int sql_bind_and_exec(sqlite3_drv_t *drv, char *buffer, int buffer_size) { @@ -519,14 +521,19 @@ static void sql_exec_async(void *_async_command) { } dataset[term_count - 2] = ERL_DRV_ATOM; dataset[term_count - 1] = drv->atom_columns; - int base = term_count; + int base_term_count = term_count; get_columns( - drv, statement, column_count, &term_count, &term_allocated, &dataset); - dataset[base + column_count * 3 + 3] = ERL_DRV_TUPLE; - dataset[base + column_count * 3 + 4] = 2; + drv, statement, column_count, base_term_count, &term_count, &term_allocated, &dataset); + term_count += 4; + if (term_count > term_allocated) { + term_allocated = max(term_count, term_allocated*2); + dataset = driver_realloc(dataset, sizeof(ErlDrvTermData) * term_allocated); + } + dataset[base_term_count + column_count * 3 + 3] = ERL_DRV_TUPLE; + dataset[base_term_count + column_count * 3 + 4] = 2; - dataset[base + column_count * 3 + 5] = ERL_DRV_ATOM; - dataset[base + column_count * 3 + 6] = drv->atom_rows; + dataset[base_term_count + column_count * 3 + 5] = ERL_DRV_ATOM; + dataset[base_term_count + column_count * 3 + 6] = drv->atom_rows; } #ifdef DEBUG @@ -954,9 +961,8 @@ static int prepared_bind(sqlite3_drv_t *drv, char *buffer, int buffer_size) { } static int prepared_columns(sqlite3_drv_t *drv, char *buffer, int buffer_size) { - int result; long long_prepared_index; - int index = 0, type, size; + int index = 0, term_count = 0, term_allocated = 0; #ifdef DEBUG fprintf(drv->log, "Finalizing prepared statement: %.*s\n", buffer_size, buffer); @@ -973,13 +979,25 @@ static int prepared_columns(sqlite3_drv_t *drv, char *buffer, int buffer_size) { } sqlite3_stmt *statement = drv->prepared_stmts[prepared_index]; -// result = -// bind_parameters(drv, buffer, buffer_size, &index, statement, &type, &size); -// if (result == SQLITE_OK) { -// return output_ok(drv); -// } else { -// return result; // error has already been output -// } + ErlDrvTermData *dataset = NULL; + + term_count += 2; + if (term_count > term_allocated) { + term_allocated = max(term_count, term_allocated*2); + dataset = driver_realloc(dataset, sizeof(ErlDrvTermData) * term_allocated); + } + dataset[term_count - 2] = ERL_DRV_PORT; + dataset[term_count - 1] = driver_mk_port(drv->port); + + int column_count = sqlite3_column_count(statement); + + get_columns( + drv, statement, column_count, 2, &term_count, &term_allocated, &dataset); + term_count += 2; + dataset[term_count - 2] = ERL_DRV_TUPLE; + dataset[term_count - 1] = 2; + + return driver_output_term(drv->port, dataset, term_count); } static int prepared_step(sqlite3_drv_t *drv, char *buffer, int buffer_size) { diff --git a/c_src/sqlite3_drv.h b/c_src/sqlite3_drv.h index cc038b7..8b8af69 100644 --- a/c_src/sqlite3_drv.h +++ b/c_src/sqlite3_drv.h @@ -26,6 +26,7 @@ #define CMD_PREPARED_RESET 8 #define CMD_PREPARED_CLEAR_BINDINGS 9 #define CMD_PREPARED_FINALIZE 10 +#define CMD_PREPARED_COLUMNS 11 // Number of bytes for each key // (160 bits for SHA1 hash) @@ -82,6 +83,7 @@ static int prepared_step(sqlite3_drv_t *drv, char *buf, int len); static int prepared_reset(sqlite3_drv_t *drv, char *buf, int len); static int prepared_clear_bindings(sqlite3_drv_t *drv, char *buf, int len); static int prepared_finalize(sqlite3_drv_t *drv, char *buf, int len); +static int prepared_columns(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 5a33801..6f860fc 100644 --- a/src/sqlite3.erl +++ b/src/sqlite3.erl @@ -16,7 +16,8 @@ -export([start_link/1, start_link/2]). -export([stop/0, close/1]). -export([sql_exec/1, sql_exec/2, sql_exec/3]). --export([prepare/2, bind/3, next/2, reset/2, clear_bindings/2, finalize/2]). +-export([prepare/2, bind/3, next/2, reset/2, clear_bindings/2, finalize/2, + columns/2]). -export([create_table/2, create_table/3, create_table/4]). -export([list_tables/0, list_tables/1, table_info/1, table_info/2]). -export([write/2, write/3, write_many/2, write_many/3]). @@ -192,6 +193,10 @@ clear_bindings(Db, Ref) -> finalize(Db, Ref) -> gen_server:call(Db, {finalize, Ref}). +-spec columns(atom(), reference()) -> sql_non_query_result(). +columns(Db, Ref) -> + gen_server:call(Db, {columns, Ref}). + %%-------------------------------------------------------------------- %% @spec create_table(Tbl :: atom(), TblInfo :: [{atom(), atom()}]) -> sql_non_query_result() %% @doc @@ -785,6 +790,7 @@ get_priv_dir() -> -define(PREPARED_RESET, 8). -define(PREPARED_CLEAR_BINDINGS, 9). -define(PREPARED_FINALIZE, 10). +-define(PREPARED_COLUMNS, 11). create_port_cmd(DbFile) -> atom_to_list(?DRIVER_NAME) ++ " " ++ DbFile. @@ -824,7 +830,8 @@ exec(Port, {Cmd, Index}) when is_integer(Index) -> next -> ?PREPARED_STEP; reset -> ?PREPARED_RESET; clear_bindings -> ?PREPARED_CLEAR_BINDINGS; - finalize -> ?PREPARED_FINALIZE + finalize -> ?PREPARED_FINALIZE; + columns -> ?PREPARED_COLUMNS end, Bin = term_to_binary(Index), port_control(Port, CmdCode, Bin), diff --git a/test/sqlite3_test.erl b/test/sqlite3_test.erl index e569a4d..6d8f67a 100644 --- a/test/sqlite3_test.erl +++ b/test/sqlite3_test.erl @@ -202,6 +202,7 @@ large_number() -> ?assertNot([{N1 + 1, N2 - 1}] == rows(sqlite3:sql_exec(ct, Query2, [N1 + 1, N2 - 1]))). prepared_test() -> + Columns = ["id", "name", "age", "wage"], Abby = {1, <<"abby">>, 20, 2000}, Marge = {2, <<"marge">>, 30, 2000}, TableInfo = [{id, integer, [primary_key]}, {name, text, [not_null, unique]}, {age, integer}, {wage, integer}], @@ -210,6 +211,7 @@ prepared_test() -> sqlite3:write(prepared, user, [{name, "abby"}, {age, 20}, {wage, 2000}]), sqlite3:write(prepared, user, [{name, "marge"}, {age, 30}, {wage, 2000}]), {ok, Ref} = sqlite3:prepare(prepared, "SELECT * FROM user"), + ?assertEqual(Columns, sqlite3:columns(prepared, Ref)), ?assertEqual(Abby, sqlite3:next(prepared, Ref)), ?assertEqual(ok, sqlite3:reset(prepared, Ref)), ?assertEqual(Abby, sqlite3:next(prepared, Ref)),