]> dgit.raspbian.org Git - dovecot.git/commitdiff
[PATCH 5/6] lib-sql: Make escape_string return int with error_r, fail instead of...
authorAki Tuomi <aki.tuomi@open-xchange.com>
Fri, 29 May 2026 11:22:57 +0000 (11:22 +0000)
committerNoah Meyerhans <noahm@debian.org>
Wed, 16 Sep 2026 19:06:35 +0000 (15:06 -0400)
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

src/auth/passdb-sql.c
src/auth/userdb-sql.c
src/lib-sql/driver-cassandra.c
src/lib-sql/driver-mysql.c
src/lib-sql/driver-pgsql.c
src/lib-sql/driver-sqlite.c
src/lib-sql/driver-sqlpool.c
src/lib-sql/driver-test.c
src/lib-sql/sql-api-private.h
src/lib-sql/sql-api.c
src/lib-sql/sql-api.h

index 3745942056e3246079ce4f8558f3043d54370c20..e9e51159eed0dad07fc3e95eba863d828426aa37 100644 (file)
@@ -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)
index 39caf3862e597e2632879fa8c08b6f644f70b5de..3b197af8b7360a44805288ab48a4b192e9b338fe 100644 (file)
@@ -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,
index 5e0497514e41f886bd2adbfe7b1c4d1fb8b982cc..735ba2b96fdd91440e8a8cf8fbfc82cd173b4316 100644 (file)
@@ -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;
        }
index d16131858e87a4e4cbf1eb85332fe67fb7243d18..6880177dc2de71ea71e44710793f6354ff50f661 100644 (file)
@@ -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)
index c6872589bed7990d9fa20bb4dfc0041805325592..f9c8e6acd2e4cc8581415eeef06a029a9e13bab2 100644 (file)
@@ -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,
index 06903b8d120c321b4010e603e4de655897163a46..5f516503937eec41a65664b15f8551a4e4771c2a 100644 (file)
@@ -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 *
index fba4ea7c6c1548b7ee67200a3655cdf89438ce7a..c2ff42c0a17fa752b9007bf7af36b7c21871f56a 100644 (file)
@@ -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)
index af748e7fffa7ade6f18d4c1ab0d0720a6f1231a4..b1ee39c7572dbde595be558c4b07d08eff400b77 100644 (file)
@@ -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)
index fb41db775b42453fe28e464090d818a24e4e3ab7..aa3053fd1a1de3a07d171e65af5a773525fddd2c 100644 (file)
@@ -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 *
index 1560c6fdd88708d387c254dbd95aa405f2491e90..6fe23074c3f98714283433388aa7d3bf6a3947ba 100644 (file)
@@ -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);
 }
 
index 4ec89b006f0d3066aa803f9e7c4ae4f51eb28959..5129befb0ea3f3dcfb7cdb04d5f8c84f0d21dc11 100644 (file)
@@ -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);