From: Aki Tuomi Date: Fri, 29 May 2026 11:22:57 +0000 (+0000) Subject: [PATCH 5/6] lib-sql: Make escape_string return int with error_r, fail instead of... X-Git-Tag: archive/raspbian/1%2.4.1+dfsg1-6+rpi1+deb13u7^2~77 X-Git-Url: https://dgit.raspbian.org/?a=commitdiff_plain;h=72d591863baabdfb55787984bfa241b866db3d21;p=dovecot.git [PATCH 5/6] lib-sql: Make escape_string return int with error_r, fail instead of unsafe fallback Change escape_string driver vfunc and sql_escape_string() to return int with separate error_r output parameter. On failure (e.g. not connected), return -1 instead of falling back to unsafe escaping. Move escaping from sql_statement_bind_str() to sql_statement_get_query() so errors can be propagated to callers. Add failed_error field to sql_transaction_context for deferred error reporting at commit time. Gbp-Pq: Name 0005-lib-sql-Make-escape_string-return-int-with-error_r-f.patch --- diff --git a/src/auth/passdb-sql.c b/src/auth/passdb-sql.c index 3745942..e9e5115 100644 --- a/src/auth/passdb-sql.c +++ b/src/auth/passdb-sql.c @@ -168,11 +168,10 @@ static void sql_query_callback(struct sql_result *result, } 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) diff --git a/src/auth/userdb-sql.c b/src/auth/userdb-sql.c index 39caf38..3b197af 100644 --- a/src/auth/userdb-sql.c +++ b/src/auth/userdb-sql.c @@ -113,11 +113,10 @@ static void sql_query_callback(struct sql_result *sql_result, } 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, diff --git a/src/lib-sql/driver-cassandra.c b/src/lib-sql/driver-cassandra.c index 5e04975..735ba2b 100644 --- a/src/lib-sql/driver-cassandra.c +++ b/src/lib-sql/driver-cassandra.c @@ -830,22 +830,26 @@ static void driver_cassandra_disconnect(struct sql_db *_db) 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 @@ -2192,8 +2196,13 @@ static void cassandra_transaction_finish(struct cassandra_transaction_context *c 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); @@ -2289,11 +2298,15 @@ driver_cassandra_try_commit_s(struct cassandra_transaction_context *ctx) 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) @@ -2796,8 +2809,14 @@ driver_cassandra_statement_query(struct sql_statement *_stmt, } 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, @@ -2826,8 +2845,13 @@ driver_cassandra_update_stmt(struct sql_transaction_context *_ctx, 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; } diff --git a/src/lib-sql/driver-mysql.c b/src/lib-sql/driver-mysql.c index d161318..6880177 100644 --- a/src/lib-sql/driver-mysql.c +++ b/src/lib-sql/driver-mysql.c @@ -462,8 +462,9 @@ static int driver_mysql_do_query(struct mysql_db *db, const char *query, 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); @@ -475,22 +476,15 @@ driver_mysql_escape_string(struct sql_db *_db, const char *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) diff --git a/src/lib-sql/driver-pgsql.c b/src/lib-sql/driver-pgsql.c index c687258..f9c8e6a 100644 --- a/src/lib-sql/driver-pgsql.c +++ b/src/lib-sql/driver-pgsql.c @@ -718,8 +718,9 @@ static void do_query(struct pgsql_result *result, 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); @@ -735,14 +736,24 @@ driver_pgsql_escape_string(struct sql_db *_db, const char *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, diff --git a/src/lib-sql/driver-sqlite.c b/src/lib-sql/driver-sqlite.c index 06903b8..5f51650 100644 --- a/src/lib-sql/driver-sqlite.c +++ b/src/lib-sql/driver-sqlite.c @@ -221,15 +221,17 @@ static void driver_sqlite_deinit_v(struct sql_db *_db) 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 * diff --git a/src/lib-sql/driver-sqlpool.c b/src/lib-sql/driver-sqlpool.c index fba4ea7..c2ff42c 100644 --- a/src/lib-sql/driver-sqlpool.c +++ b/src/lib-sql/driver-sqlpool.c @@ -573,8 +573,9 @@ static void driver_sqlpool_disconnect(struct sql_db *_db) 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; @@ -584,11 +585,12 @@ driver_sqlpool_escape_string(struct sql_db *_db, const char *string) 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) diff --git a/src/lib-sql/driver-test.c b/src/lib-sql/driver-test.c index af748e7..b1ee39c 100644 --- a/src/lib-sql/driver-test.c +++ b/src/lib-sql/driver-test.c @@ -33,10 +33,12 @@ static int driver_test_sqlite_init(struct event *event, struct sql_db **db_r, 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); @@ -237,9 +239,10 @@ static int driver_test_connect(struct sql_db *_db ATTR_UNUSED) 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++) { @@ -248,13 +251,17 @@ driver_test_mysql_escape_string(struct sql_db *_db ATTR_UNUSED, 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) diff --git a/src/lib-sql/sql-api-private.h b/src/lib-sql/sql-api-private.h index fb41db7..aa3053f 100644 --- a/src/lib-sql/sql-api-private.h +++ b/src/lib-sql/sql-api-private.h @@ -81,7 +81,8 @@ struct sql_db_vfuncs { 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, @@ -215,6 +216,7 @@ struct sql_statement { 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. */ @@ -252,6 +254,7 @@ struct sql_transaction_context { /* commit() must use this query list if head is non-NULL. */ struct sql_transaction_query *head, *tail; + char *failed_error; bool non_atomic; }; @@ -277,7 +280,8 @@ inline static const char *sql_db_table_prefix(struct sql_db *db) { 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 * diff --git a/src/lib-sql/sql-api.c b/src/lib-sql/sql-api.c index 1560c6f..6fe2307 100644 --- a/src/lib-sql/sql-api.c +++ b/src/lib-sql/sql-api.c @@ -98,7 +98,6 @@ static const struct sql_result_vfuncs sql_result_error_vfuncs = { .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); @@ -327,9 +326,10 @@ void sql_disconnect(struct sql_db *db) 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, @@ -391,19 +391,25 @@ default_sql_statement_init_prepared(struct sql_prepared_statement *stmt) 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 ? */ @@ -413,7 +419,31 @@ const char *sql_statement_get_query(struct sql_statement *stmt) 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); @@ -422,23 +452,37 @@ const char *sql_statement_get_query(struct sql_statement *stmt) 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; } @@ -447,8 +491,14 @@ static void default_sql_update_stmt(struct sql_transaction_context *ctx, 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); } @@ -487,6 +537,7 @@ sql_statement_init_fields(struct sql_statement *stmt, struct sql_db *db) { stmt->db = db; p_array_init(&stmt->args, stmt->pool, 8); + p_array_init(&stmt->args_need_escaping, stmt->pool, 8); } struct sql_statement * @@ -545,10 +596,10 @@ void sql_statement_set_no_log_expanded_values(struct sql_statement *stmt, 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); @@ -908,6 +959,13 @@ void sql_transaction_commit(struct sql_transaction_context **_ctx, 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; @@ -924,6 +982,12 @@ int sql_transaction_commit_s(struct sql_transaction_context **_ctx, 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); } @@ -932,6 +996,7 @@ void sql_transaction_rollback(struct sql_transaction_context **_ctx) struct sql_transaction_context *ctx = *_ctx; *_ctx = NULL; + i_free(ctx->failed_error); ctx->db->v.transaction_rollback(ctx); } diff --git a/src/lib-sql/sql-api.h b/src/lib-sql/sql-api.h index 4ec89b0..5129bef 100644 --- a/src/lib-sql/sql-api.h +++ b/src/lib-sql/sql-api.h @@ -114,7 +114,8 @@ int sql_connect(struct sql_db *db); 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);