diff --git a/c_src/sqlite3_drv.c b/c_src/sqlite3_drv.c index dc02fe8..7706031 100644 --- a/c_src/sqlite3_drv.c +++ b/c_src/sqlite3_drv.c @@ -124,6 +124,68 @@ static inline unsigned int sql_async_key(char *db_name, ErlDrvPort port) { } } +static inline int return_error( + sqlite3_drv_t *drv, int error_code, const char *error, + ErlDrvTermData **dataset_p, int *term_count_p, int *term_allocated_p, + int* error_code_p) { + if (error_code_p) { + *error_code_p = error_code; + } + EXTEND_DATASET_PTR(9); + append_to_dataset(9, *dataset_p, *term_count_p, + ERL_DRV_ATOM, drv->atom_error, + ERL_DRV_INT, (ErlDrvTermData) error_code, + ERL_DRV_STRING, (ErlDrvTermData) error, (ErlDrvTermData) strlen(error), + ERL_DRV_TUPLE, (ErlDrvTermData) 3); +// int i; +// for (i = 0; i < *term_count_p; i++) { +// printf("%d\n", (*dataset_p)[i]); +// } + return 0; +} + +static inline int output_error( + sqlite3_drv_t *drv, int error_code, const char *error) { + int term_count = 2, term_allocated = 13; + // for some reason breaks if allocated as an array on stack + // even though it shouldn't be extended + ErlDrvTermData *dataset = driver_alloc(sizeof(ErlDrvTermData) * term_allocated); + dataset[0] = ERL_DRV_PORT; + dataset[1] = driver_mk_port(drv->port); + return_error(drv, error_code, error, &dataset, &term_count, &term_allocated, NULL); + term_count += 2; + dataset[11] = ERL_DRV_TUPLE; + dataset[12] = 2; + #ifdef PRE_R16B + driver_output_term(drv->port, + #else + erl_drv_output_term(dataset[1], + #endif + dataset, term_count); + driver_free(dataset); + return 0; +} + +static inline int output_db_error(sqlite3_drv_t *drv) { + return output_error(drv, sqlite3_errcode(drv->db), sqlite3_errmsg(drv->db)); +} + +static inline int output_ok(sqlite3_drv_t *drv) { + // Return {Port, ok} + ErlDrvTermData spec[] = { + ERL_DRV_PORT, driver_mk_port(drv->port), + ERL_DRV_ATOM, drv->atom_ok, + ERL_DRV_TUPLE, 2 + }; + return + #ifdef PRE_R16B + driver_output_term(drv->port, + #else + erl_drv_output_term(spec[1], + #endif + spec, sizeof(spec) / sizeof(spec[0])); +} + static ErlDrvEntry sqlite3_driver_entry = { NULL, /* init */ start, /* startup (defined below) */ @@ -196,16 +258,6 @@ static ErlDrvData start(ErlDrvPort port, char* cmd) { // Create and open the database status = sqlite3_open(db_name, &db); - - if (status != SQLITE_OK) { - LOG_ERROR("Unable to open file: %s because %s\n\n", db_name, sqlite3_errmsg(db)); - // We don't do this because there's no way to pass the error to Erlang - // sqlite3_close(db); - // driver_free(drv); - // return ERL_DRV_ERROR_GENERAL; - } else { - LOG_DEBUG("Opened file %s\n", db_name); - } #if defined(_MSC_VER) #pragma warning(default: 4306) #endif @@ -232,6 +284,14 @@ static ErlDrvData start(ErlDrvPort port, char* cmd) { drv->atom_done = driver_mk_atom("done"); drv->atom_unknown_cmd = driver_mk_atom("unknown_command"); + if (status != SQLITE_OK) { + LOG_ERROR("Unable to open file %s: \"%s\"\n\n", db_name, sqlite3_errmsg(db)); + output_db_error(drv); + } else { + LOG_DEBUG("Opened file %s\n", db_name); + output_ok(drv); + } + return (ErlDrvData) drv; } @@ -261,8 +321,6 @@ static void stop(ErlDrvData handle) { driver_free(drv); } -static inline int output_error(sqlite3_drv_t *drv, int error_code, const char *error); - // Handle input from Erlang VM static ErlDrvSSizeT control( ErlDrvData drv_data, unsigned int command, char *buf, @@ -312,68 +370,6 @@ static ErlDrvSSizeT control( return 0; } -static inline int return_error( - sqlite3_drv_t *drv, int error_code, const char *error, - ErlDrvTermData **dataset_p, int *term_count_p, int *term_allocated_p, - int* error_code_p) { - if (error_code_p) { - *error_code_p = error_code; - } - EXTEND_DATASET_PTR(9); - append_to_dataset(9, *dataset_p, *term_count_p, - ERL_DRV_ATOM, drv->atom_error, - ERL_DRV_INT, (ErlDrvTermData) error_code, - ERL_DRV_STRING, (ErlDrvTermData) error, (ErlDrvTermData) strlen(error), - ERL_DRV_TUPLE, (ErlDrvTermData) 3); -// int i; -// for (i = 0; i < *term_count_p; i++) { -// printf("%d\n", (*dataset_p)[i]); -// } - return 0; -} - -static inline int output_error( - sqlite3_drv_t *drv, int error_code, const char *error) { - int term_count = 2, term_allocated = 13; - // for some reason breaks if allocated as an array on stack - // even though it shouldn't be extended - ErlDrvTermData *dataset = driver_alloc(sizeof(ErlDrvTermData) * term_allocated); - dataset[0] = ERL_DRV_PORT; - dataset[1] = driver_mk_port(drv->port); - return_error(drv, error_code, error, &dataset, &term_count, &term_allocated, NULL); - term_count += 2; - dataset[11] = ERL_DRV_TUPLE; - dataset[12] = 2; - #ifdef PRE_R16B - driver_output_term(drv->port, - #else - erl_drv_output_term(dataset[1], - #endif - dataset, term_count); - driver_free(dataset); - return 0; -} - -static inline int output_db_error(sqlite3_drv_t *drv) { - return output_error(drv, sqlite3_errcode(drv->db), sqlite3_errmsg(drv->db)); -} - -static inline int output_ok(sqlite3_drv_t *drv) { - // Return {Port, ok} - ErlDrvTermData spec[] = { - ERL_DRV_PORT, driver_mk_port(drv->port), - ERL_DRV_ATOM, drv->atom_ok, - ERL_DRV_TUPLE, 2 - }; - return - #ifdef PRE_R16B - driver_output_term(drv->port, - #else - erl_drv_output_term(spec[1], - #endif - spec, sizeof(spec) / sizeof(spec[0])); -} - static int enable_load_extension(sqlite3_drv_t* drv, char *buf, int len) { #ifdef ERLANG_SQLITE3_LOAD_EXTENSION char enable = buf[0]; diff --git a/src/sqlite3.erl b/src/sqlite3.erl index e79401d..23f23fd 100644 --- a/src/sqlite3.erl +++ b/src/sqlite3.erl @@ -778,21 +778,31 @@ value_to_sql(X) -> sqlite3_lib:value_to_sql(X). -spec init([any()]) -> {'ok', #state{}} | {'stop', string()}. init(Options) -> - DbFile = proplists:get_value(file, Options), PrivDir = get_priv_dir(), case erl_ddll:load(PrivDir, atom_to_list(?DRIVER_NAME)) of ok -> - Port = open_port({spawn, create_port_cmd(DbFile)}, [binary]), - {ok, #state{port = Port, ops = Options}}; + do_init(Options); {error, permanent} -> %% already loaded! - Port = open_port({spawn, create_port_cmd(DbFile)}, [binary]), - {ok, #state{port = Port, ops = Options}}; + do_init(Options); {error, Error} -> Msg = io_lib:format("Error loading ~p: ~s", [?DRIVER_NAME, erl_ddll:format_error(Error)]), {stop, lists:flatten(Msg)} end. +-spec do_init([any()]) -> {'ok', #state{}} | {'stop', string()}. +do_init(Options) -> + DbFile = proplists:get_value(file, Options), + Port = open_port({spawn, create_port_cmd(DbFile)}, [binary]), + receive + {Port, ok} -> + {ok, #state{port = Port, ops = Options}}; + {Port, {error, Code, Message}} -> + Msg = io_lib:format("Error opening DB file ~p: code ~B, message '~s'", + [DbFile, Code, Message]), + {stop, lists:flatten(Msg)} + end. + %%-------------------------------------------------------------------- %% @doc Handling call messages %% @end diff --git a/test/sqlite3_test.erl b/test/sqlite3_test.erl index d281e25..e5027fb 100644 --- a/test/sqlite3_test.erl +++ b/test/sqlite3_test.erl @@ -344,6 +344,16 @@ issue23() -> ?assertEqual([SingleStmtResult, SingleStmtResult], ScriptResult), sqlite3:close(issue23). +non_db_file_test() -> + process_flag(trap_exit, true), + ?assertMatch({error, _}, + sqlite3:start_link(bad_file, [{file, "/"}])), + receive + {'EXIT', _, _} -> ok; + _ -> ?assert(false) + end. + + % create, read, update, delete %%==================================================================== %% Internal functions