}
static int passdb_sql_escape(const char *str, const char **output_r,
- void *context, const char **error_r ATTR_UNUSED)
+ void *context, const char **error_r)
{
struct sql_db *db = context;
- *output_r = sql_escape_string(db, str);
- return 0;
+ return sql_escape_string(db, str, output_r, error_r);
}
static void sql_lookup_pass(struct passdb_sql_request *sql_request)
}
static int userdb_sql_escape(const char *str, const char **output_r,
- void *context, const char **error_r ATTR_UNUSED)
+ void *context, const char **error_r)
{
struct sql_db *db = context;
- *output_r = sql_escape_string(db, str);
- return 0;
+ return sql_escape_string(db, str, output_r, error_r);
}
static void userdb_sql_lookup(struct auth_request *auth_request,
driver_cassandra_close(db, "Disconnected");
}
-static const char *
+static int
driver_cassandra_escape_string(struct sql_db *db ATTR_UNUSED,
- const char *string)
+ const char *string, const char **output_r,
+ const char **error_r ATTR_UNUSED)
{
string_t *escaped;
unsigned int i;
- if (strchr(string, '\'') == NULL)
- return string;
+ if (strchr(string, '\'') == NULL) {
+ *output_r = string;
+ return 0;
+ }
escaped = t_str_new(strlen(string)+10);
for (i = 0; string[i] != '\0'; i++) {
if (string[i] == '\'')
str_append_c(escaped, '\'');
str_append_c(escaped, string[i]);
}
- return str_c(escaped);
+ *output_r = str_c(escaped);
+ return 0;
}
static void
if (stmt->prep == NULL)
have_nonprepared = TRUE;
if (stmt->cass_stmt == NULL) {
- stmt->cass_stmt = cass_statement_new(
- sql_statement_get_query(&stmt->stmt), 0);
+ const char *query, *error;
+ if (sql_statement_get_query(&stmt->stmt,
+ &query, &error) < 0) {
+ cassandra_transaction_finish(ctx, error);
+ return;
+ }
+ stmt->cass_stmt = cass_statement_new(query, 0);
if (stmt->timestamp != 0) {
cass_statement_set_timestamp(stmt->cass_stmt,
stmt->timestamp);
i_panic("cassandra: sql_transaction_commit_s() not supported for prepared statements");
}
+ const char *query, *error;
+ if (sql_statement_get_query(&stmt->stmt, &query, &error) < 0) {
+ transaction_set_failed(ctx, error);
+ return;
+ }
+
/* just a single query, send it */
driver_cassandra_sync_init(db);
- result = driver_cassandra_sync_query(db,
- sql_statement_get_query(&stmt->stmt),
- ctx->query_type);
+ result = driver_cassandra_sync_query(db, query, ctx->query_type);
driver_cassandra_sync_deinit(db);
if (sql_result_next_row(result) < 0)
} else {
/* Not a prepared statement. Generate a statement from
the query string. */
- stmt->result->statement =
- cass_statement_new(sql_statement_get_query(_stmt), 0);
+ const char *query, *error;
+ if (sql_statement_get_query(_stmt, &query, &error) < 0) {
+ stmt->result->error = i_strdup(error);
+ result_finish(stmt->result);
+ cassandra_sql_statement_free(stmt);
+ return;
+ }
+ stmt->result->statement = cass_statement_new(query, 0);
stmt->result->timestamp = stmt->timestamp;
if (stmt->timestamp != 0) {
cass_statement_set_timestamp(stmt->result->statement,
i_assert(affected_rows == NULL);
- if (!driver_cassandra_update_query_type(ctx,
- sql_statement_get_query(_stmt))) {
+ const char *query, *error;
+ if (sql_statement_get_query(_stmt, &query, &error) < 0) {
+ transaction_set_failed(ctx, error);
+ cassandra_sql_statement_free(stmt);
+ return;
+ }
+ if (!driver_cassandra_update_query_type(ctx, query)) {
cassandra_sql_statement_free(stmt);
return;
}
return -1;
}
-static const char *
-driver_mysql_escape_string(struct sql_db *_db, const char *string)
+static int
+driver_mysql_escape_string(struct sql_db *_db, const char *string,
+ const char **output_r, const char **error_r)
{
struct mysql_db *db = container_of(_db, struct mysql_db, api);
size_t len = strlen(string);
}
if (_db->state == SQL_DB_STATE_DISCONNECTED) {
- /* FIXME: we don't have a valid connection, so fallback
- to using default escaping. the next query will most
- likely fail anyway so it shouldn't matter that much
- what we return here.. Anyway, this API needs
- changing so that the escaping function could already
- fail the query reliably. */
- to = t_buffer_get(len * 2 + 1);
- len = mysql_escape_string(to, string, len);
- t_buffer_alloc(len + 1);
- return to;
+ *error_r = SQL_ERRSTR_NOT_CONNECTED;
+ return -1;
}
to = t_buffer_get(len * 2 + 1);
len = mysql_real_escape_string(db->mysql, to, string, len);
t_buffer_alloc(len + 1);
- return to;
+ *output_r = to;
+ return 0;
}
static void driver_mysql_exec(struct sql_db *_db, const char *query)
}
}
-static const char *
-driver_pgsql_escape_string(struct sql_db *_db, const char *string)
+static int
+driver_pgsql_escape_string(struct sql_db *_db, const char *string,
+ const char **output_r, const char **error_r)
{
struct pgsql_db *db = (struct pgsql_db *)_db;
size_t len = strlen(string);
to = t_buffer_get(len * 2 + 1);
len = PQescapeStringConn(db->pg, to, string, len, &error);
- } else
-#endif
- {
- to = t_buffer_get(len * 2 + 1);
- len = PQescapeString(to, string, len);
+ if (error != 0) {
+ *error_r = last_error(db);
+ return -1;
+ }
+ t_buffer_alloc(len + 1);
+ *output_r = to;
+ return 0;
+ } else {
+ *error_r = SQL_ERRSTR_NOT_CONNECTED;
+ return -1;
}
+#else
+ to = t_buffer_get(len * 2 + 1);
+ len = PQescapeString(to, string, len);
t_buffer_alloc(len + 1);
- return to;
+ *output_r = to;
+ return 0;
+#endif
}
static void exec_callback(struct sql_result *_result,
i_free(db);
}
-static const char *
+static int
driver_sqlite_escape_string(struct sql_db *_db ATTR_UNUSED,
- const char *string)
+ const char *string, const char **output_r,
+ const char **error_r ATTR_UNUSED)
{
const size_t len = strlen(string) * 2 + 1;
char *escaped = t_malloc_no0(len);
if (sqlite3_snprintf(len, escaped, "%q", string) == NULL)
i_unreached();
- return escaped;
+ *output_r = escaped;
+ return 0;
}
static const char *
driver_sqlpool_abort_requests(db);
}
-static const char *
-driver_sqlpool_escape_string(struct sql_db *_db, const char *string)
+static int
+driver_sqlpool_escape_string(struct sql_db *_db, const char *string,
+ const char **output_r, const char **error_r)
{
struct sqlpool_db *db = (struct sqlpool_db *)_db;
const struct sqlpool_connection *conns;
conns = array_get(&db->all_connections, &count);
for (i = 0; i < count; i++) {
if (SQL_DB_IS_READY(conns[i].db))
- return sql_escape_string(conns[i].db, string);
+ return sql_escape_string(conns[i].db, string,
+ output_r, error_r);
}
/* no ready connections. just use the first one (we're guaranteed
to always have one) */
- return sql_escape_string(conns[0].db, string);
+ return sql_escape_string(conns[0].db, string, output_r, error_r);
}
static void driver_sqlpool_timeout(struct sqlpool_db *db)
static void driver_test_deinit(struct sql_db *_db);
static int driver_test_connect(struct sql_db *_db);
static void driver_test_disconnect(struct sql_db *_db);
-static const char *
-driver_test_mysql_escape_string(struct sql_db *_db, const char *string);
-static const char *
-driver_test_escape_string(struct sql_db *_db, const char *string);
+static int
+driver_test_mysql_escape_string(struct sql_db *_db, const char *string,
+ const char **output_r, const char **error_r);
+static int
+driver_test_escape_string(struct sql_db *_db, const char *string,
+ const char **output_r, const char **error_r);
static void driver_test_exec(struct sql_db *_db, const char *query);
static void driver_test_query(struct sql_db *_db, const char *query,
sql_query_callback_t *callback, void *context);
static void driver_test_disconnect(struct sql_db *_db ATTR_UNUSED)
{ }
-static const char *
+static int
driver_test_mysql_escape_string(struct sql_db *_db ATTR_UNUSED,
- const char *string)
+ const char *string, const char **output_r,
+ const char **error_r ATTR_UNUSED)
{
string_t *esc = t_str_new(strlen(string));
for(const char *ptr = string; *ptr != '\0'; ptr++) {
str_append_c(esc, '\\');
str_append_c(esc, *ptr);
}
- return str_c(esc);
+ *output_r = str_c(esc);
+ return 0;
}
-static const char *
-driver_test_escape_string(struct sql_db *_db ATTR_UNUSED, const char *string)
+static int
+driver_test_escape_string(struct sql_db *_db ATTR_UNUSED, const char *string,
+ const char **output_r,
+ const char **error_r ATTR_UNUSED)
{
- return string;
+ *output_r = string;
+ return 0;
}
static void driver_test_exec(struct sql_db *_db, const char *query)
int (*connect)(struct sql_db *db);
void (*disconnect)(struct sql_db *db);
- const char *(*escape_string)(struct sql_db *db, const char *string);
+ int (*escape_string)(struct sql_db *db, const char *string,
+ const char **output_r, const char **error_r);
void (*exec)(struct sql_db *db, const char *query);
/* Only implement this if the driver can really do asynchronous callbacks,
pool_t pool;
const char *query_template;
ARRAY_TYPE(const_string) args;
+ ARRAY(bool) args_need_escaping;
/* Tell the driver to not log this query with expanded values.
This works only for prepared statements. */
/* commit() must use this query list if head is non-NULL. */
struct sql_transaction_query *head, *tail;
+ char *failed_error;
bool non_atomic;
};
void sql_transaction_add_query(struct sql_transaction_context *ctx, pool_t pool,
const char *query, unsigned int *affected_rows);
const char *sql_statement_get_log_query(struct sql_statement *stmt);
-const char *sql_statement_get_query(struct sql_statement *stmt);
+int sql_statement_get_query(struct sql_statement *stmt,
+ const char **query_r, const char **error_r);
void sql_connection_log_finished(struct sql_db *db);
struct event_passthrough *
.get_error = sql_result_error_get_error,
};
-static struct sql_result *sql_result_new_error(const char *error) ATTR_UNUSED;
static struct sql_result *sql_result_new_error(const char *error)
{
struct sql_result_error *result = i_new(struct sql_result_error, 1);
db->v.disconnect(db);
}
-const char *sql_escape_string(struct sql_db *db, const char *string)
+int sql_escape_string(struct sql_db *db, const char *string,
+ const char **output_r, const char **error_r)
{
- return db->v.escape_string(db, string);
+ return db->v.escape_string(db, string, output_r, error_r);
}
const char *sql_escape_blob(struct sql_db *db,
const char *sql_statement_get_log_query(struct sql_statement *stmt)
{
+ const char *query, *error;
if (stmt->no_log_expanded_values)
return stmt->query_template;
- return sql_statement_get_query(stmt);
+ if (sql_statement_get_query(stmt, &query, &error) < 0)
+ return stmt->query_template;
+ return query;
}
-const char *sql_statement_get_query(struct sql_statement *stmt)
+int sql_statement_get_query(struct sql_statement *stmt,
+ const char **query_r, const char **error_r)
{
- string_t *query = t_str_new(128);
+ string_t *query = str_new(default_pool, 128);
const char *const *args;
- unsigned int args_count, arg_pos = 0;
+ const bool *need_escaping_flags;
+ unsigned int args_count, need_escaping_count, arg_pos = 0;
const char *p0, *p1;
args = array_get(&stmt->args, &args_count);
+ need_escaping_flags = array_get(&stmt->args_need_escaping, &need_escaping_count);
p0 = stmt->query_template;
while ((p1 = strchr(p0, '?')) != NULL) {
/* append until ? */
i_panic("lib-sql: Missing bind for arg #%u in statement: %s",
arg_pos, stmt->query_template);
}
- str_append(query, args[arg_pos++]);
+ if (arg_pos < need_escaping_count && need_escaping_flags[arg_pos]) {
+ const char *escaped;
+
+ /* Escape in a nested data stack frame so the
+ driver's temporary escape buffer is freed
+ immediately. The escaped value is appended to the
+ heap-allocated query before the frame is popped. */
+ T_BEGIN {
+ if (sql_escape_string(stmt->db, args[arg_pos],
+ &escaped, error_r) < 0)
+ escaped = NULL;
+ else {
+ str_append_c(query, '\'');
+ str_append(query, escaped);
+ str_append_c(query, '\'');
+ }
+ } T_END_PASS_STR_IF(escaped == NULL, error_r);
+ if (escaped == NULL) {
+ str_free(&query);
+ return -1;
+ }
+ } else {
+ str_append(query, args[arg_pos]);
+ }
+ arg_pos++;
p0 = p1 + 1;
}
str_append(query, p0);
i_panic("lib-sql: Too many bind args (%u) for statement: %s",
args_count, stmt->query_template);
}
- return str_c(query);
+ *query_r = t_strdup(str_c(query));
+ str_free(&query);
+ return 0;
}
static void
default_sql_statement_query(struct sql_statement *stmt,
sql_query_callback_t *callback, void *context)
{
- sql_query(stmt->db, sql_statement_get_query(stmt),
- callback, context);
+ const char *query, *error;
+ if (sql_statement_get_query(stmt, &query, &error) < 0) {
+ sql_query_callback_delayed(stmt->db,
+ sql_result_new_error(error),
+ callback, context);
+ pool_unref(&stmt->pool);
+ return;
+ }
+ sql_query(stmt->db, query, callback, context);
pool_unref(&stmt->pool);
}
static struct sql_result *
default_sql_statement_query_s(struct sql_statement *stmt)
{
- struct sql_result *result =
- sql_query_s(stmt->db, sql_statement_get_query(stmt));
+ const char *query, *error;
+ if (sql_statement_get_query(stmt, &query, &error) < 0) {
+ struct sql_result *result = sql_result_new_error(error);
+ pool_unref(&stmt->pool);
+ return result;
+ }
+ struct sql_result *result = sql_query_s(stmt->db, query);
pool_unref(&stmt->pool);
return result;
}
struct sql_statement *stmt,
unsigned int *affected_rows)
{
- ctx->db->v.update(ctx, sql_statement_get_query(stmt),
- affected_rows);
+ const char *query, *error;
+ if (sql_statement_get_query(stmt, &query, &error) < 0) {
+ if (ctx->failed_error == NULL)
+ ctx->failed_error = i_strdup(error);
+ pool_unref(&stmt->pool);
+ return;
+ }
+ ctx->db->v.update(ctx, query, affected_rows);
pool_unref(&stmt->pool);
}
{
stmt->db = db;
p_array_init(&stmt->args, stmt->pool, 8);
+ p_array_init(&stmt->args_need_escaping, stmt->pool, 8);
}
struct sql_statement *
void sql_statement_bind_str(struct sql_statement *stmt,
unsigned int column_idx, const char *value)
{
- const char *escaped_value =
- p_strdup_printf(stmt->pool, "'%s'",
- sql_escape_string(stmt->db, value));
- array_idx_set(&stmt->args, column_idx, &escaped_value);
+ const char *value_dup = p_strdup(stmt->pool, value);
+ array_idx_set(&stmt->args, column_idx, &value_dup);
+ bool needs_escaping = TRUE;
+ array_idx_set(&stmt->args_need_escaping, column_idx, &needs_escaping);
if (stmt->db->v.statement_bind_str != NULL)
stmt->db->v.statement_bind_str(stmt, column_idx, value);
struct sql_db *db = ctx->db;
*_ctx = NULL;
+ if (ctx->failed_error != NULL) {
+ sql_commit_schedule_delayed(db, ctx->failed_error, callback, context);
+ i_free(ctx->failed_error);
+ ctx->db->v.transaction_rollback(ctx);
+ return;
+ }
+
if (ctx->db->v.transaction_commit != NULL) {
ctx->db->v.transaction_commit(ctx, callback, context);
return;
struct sql_transaction_context *ctx = *_ctx;
*_ctx = NULL;
+ if (ctx->failed_error != NULL) {
+ *error_r = t_strdup(ctx->failed_error);
+ i_free(ctx->failed_error);
+ ctx->db->v.transaction_rollback(ctx);
+ return -1;
+ }
return ctx->db->v.transaction_commit_s(ctx, error_r);
}
struct sql_transaction_context *ctx = *_ctx;
*_ctx = NULL;
+ i_free(ctx->failed_error);
ctx->db->v.transaction_rollback(ctx);
}
void sql_disconnect(struct sql_db *db);
/* Escape the given string if needed and return it. */
-const char *sql_escape_string(struct sql_db *db, const char *string);
+int sql_escape_string(struct sql_db *db, const char *string,
+ const char **output_r, const char **error_r);
/* Escape the given data as a string. */
const char *sql_escape_blob(struct sql_db *db,
const unsigned char *data, size_t size);