From ffa488ab0e403dc4cdb2b0ae0350b7b3f781a6c5 Mon Sep 17 00:00:00 2001 From: LucaCappelletti94 Date: Wed, 30 Sep 2026 17:24:38 +0200 Subject: [PATCH] Return SQLITE_NOMEM on allocation failures in initialization and open --- .github/workflows/ci4sqlite3mc.yml | 5 ++ Makefile.am | 8 +- src/cipher_common.c | 60 ++++--------- src/sqlite3mc.c | 107 +++++++++++----------- test/oomtest.c | 138 +++++++++++++++++++++++++++++ 5 files changed, 223 insertions(+), 95 deletions(-) create mode 100644 test/oomtest.c diff --git a/.github/workflows/ci4sqlite3mc.yml b/.github/workflows/ci4sqlite3mc.yml index 54bc061d..0050c55d 100644 --- a/.github/workflows/ci4sqlite3mc.yml +++ b/.github/workflows/ci4sqlite3mc.yml @@ -57,6 +57,11 @@ jobs: ./tempfiletest 20000 1 db4096 ./tempfiletest 20000 1 chunk faults + - name: Allocation failure tests + run: | + make oomtest + ./oomtest + android_build: name: Android (API ${{ matrix.api-level }}) runs-on: ubuntu-latest diff --git a/Makefile.am b/Makefile.am index fae8fa40..11ac4abe 100644 --- a/Makefile.am +++ b/Makefile.am @@ -164,8 +164,8 @@ sqlite3shell_LDFLAGS += -no-install endif -# Tests for the cryptographic primitives (not installed, built like the shell). -check_PROGRAMS = cryptotest tempfiletest +# Test programs (not installed, built like the shell). +check_PROGRAMS = cryptotest tempfiletest oomtest cryptotest_SOURCES = test/cryptotest.c cryptotest_CFLAGS = $(sqlite3shell_CFLAGS) cryptotest_LDADD = $(sqlite3shell_LDADD) @@ -174,3 +174,7 @@ tempfiletest_SOURCES = test/tempfiletest.c tempfiletest_CFLAGS = $(sqlite3shell_CFLAGS) tempfiletest_LDADD = $(sqlite3shell_LDADD) tempfiletest_LDFLAGS = $(sqlite3shell_LDFLAGS) +oomtest_SOURCES = test/oomtest.c +oomtest_CFLAGS = $(sqlite3shell_CFLAGS) +oomtest_LDADD = $(sqlite3shell_LDADD) +oomtest_LDFLAGS = $(sqlite3shell_LDFLAGS) diff --git a/src/cipher_common.c b/src/cipher_common.c index 9dc495a5..94473ddc 100644 --- a/src/cipher_common.c +++ b/src/cipher_common.c @@ -67,63 +67,39 @@ sqlite3mcCloneCodecParameterTable() /* Count number of codecs and cipher parameters */ int nTables = 0; int nParams = 0; - int j, k, n; + int j, n; + int offset = 0; CipherParams* cloneCipherParams; CodecParameter* cloneCodecParams; for (j = 0; globalCodecParameterTable[j].m_name[0] != 0; ++j) { CipherParams* params = globalCodecParameterTable[j].m_params; - for (k = 0; params[k].m_name[0] != 0; ++k); - nParams += k; + for (n = 0; params[n].m_name[0] != 0; ++n); + nParams += n; } nTables = j; - /* Allocate memory for cloned codec parameter tables (including sentinel for each table) */ - cloneCipherParams = (CipherParams*) sqlite3_malloc((nParams + nTables) * sizeof(CipherParams)); - cloneCodecParams = (CodecParameter*) sqlite3_malloc((nTables + 1) * sizeof(CodecParameter)); + /* The table array and all parameter arrays, each with its sentinel, share one allocation */ + cloneCodecParams = (CodecParameter*) sqlite3_malloc((nTables + 1) * sizeof(CodecParameter) + + (nParams + nTables) * sizeof(CipherParams)); + if (cloneCodecParams == NULL) + return NULL; + cloneCipherParams = (CipherParams*) &cloneCodecParams[nTables + 1]; - /* Create copy of tables */ - if (cloneCodecParams != NULL) - { - int offset = 0; - for (j = 0; j < nTables; ++j) - { - CipherParams* params = globalCodecParameterTable[j].m_params; - cloneCodecParams[j].m_name = globalCodecParameterTable[j].m_name; - cloneCodecParams[j].m_id = globalCodecParameterTable[j].m_id; - cloneCodecParams[j].m_params = &cloneCipherParams[offset]; - for (n = 0; params[n].m_name[0] != 0; ++n); - /* Copy all parameters of the current table (including sentinel) */ - for (k = 0; k <= n; ++k) - { - cloneCipherParams[offset + k].m_name = params[k].m_name; - cloneCipherParams[offset + k].m_value = params[k].m_value; - cloneCipherParams[offset + k].m_default = params[k].m_default; - cloneCipherParams[offset + k].m_minValue = params[k].m_minValue; - cloneCipherParams[offset + k].m_maxValue = params[k].m_maxValue; - } - offset += (n + 1); - } - cloneCodecParams[nTables].m_name = globalCodecParameterTable[nTables].m_name; - cloneCodecParams[nTables].m_id = globalCodecParameterTable[nTables].m_id; - cloneCodecParams[nTables].m_params = NULL; - } - else + for (j = 0; j < nTables; ++j) { - sqlite3_free(cloneCipherParams); + CipherParams* params = globalCodecParameterTable[j].m_params; + for (n = 0; params[n].m_name[0] != 0; ++n); + memcpy(&cloneCipherParams[offset], params, (n + 1) * sizeof(CipherParams)); + cloneCodecParams[j] = globalCodecParameterTable[j]; + cloneCodecParams[j].m_params = &cloneCipherParams[offset]; + offset += n + 1; } + cloneCodecParams[nTables] = globalCodecParameterTable[nTables]; return cloneCodecParams; } -SQLITE_PRIVATE void -sqlite3mcFreeCodecParameterTable(void* ptr) -{ - CodecParameter* codecParams = (CodecParameter*)ptr; - sqlite3_free(codecParams[0].m_params); - sqlite3_free(codecParams); -} - static const CipherDescriptor mcSentinelDescriptor = { "", NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL diff --git a/src/sqlite3mc.c b/src/sqlite3mc.c index 477f0417..26be7b3b 100644 --- a/src/sqlite3mc.c +++ b/src/sqlite3mc.c @@ -426,10 +426,9 @@ mcRegisterCodecExtensions(sqlite3* db, char** pzErrMsg, const sqlite3_api_routin rc = (codecParameterTable != NULL) ? SQLITE_OK : SQLITE_NOMEM; if (rc == SQLITE_OK) { - sqlite3_set_clientdata(db, globalConfigTableName, codecParameterTable, sqlite3mcFreeCodecParameterTable); + /* On failure this frees the table */ + rc = sqlite3_set_clientdata(db, globalConfigTableName, codecParameterTable, sqlite3_free); } - - rc = (codecParameterTable != NULL) ? SQLITE_OK : SQLITE_NOMEM; if (rc == SQLITE_OK) { rc = sqlite3_create_function(db, "sqlite3mc_config", 1, SQLITE_UTF8 | SQLITE_DETERMINISTIC, @@ -513,11 +512,23 @@ sqlite3mcGetGlobalCipherCount() return cipherCount; } +static void +mcFreeCipherParams(CipherParams* params) +{ + int k; + for (k = 0; params[k].m_name[0] != 0; ++k) + { + sqlite3_free((char*) params[k].m_name); + } + sqlite3_free(params); +} + static int sqlite3mcRegisterCipher(const CipherDescriptor* desc, const CipherParams* params, int makeDefault) { - int rc; int np; + int n; + char* cipherName; CipherParams* cipherParams; /* Sanity checks */ @@ -567,60 +578,55 @@ sqlite3mcRegisterCipher(const CipherDescriptor* desc, const CipherParams* params /* Sanity checks were successful, now register cipher */ + if (globalCipherCount >= CODEC_COUNT_MAX) + return SQLITE_NOMEM; + cipherParams = (CipherParams*) sqlite3_malloc((np+1) * sizeof(CipherParams)); if (!cipherParams) return SQLITE_NOMEM; - /* Check for */ - if (globalCipherCount < CODEC_COUNT_MAX) + /* Copy parameters before touching the global tables, so that a failure leaves them unchanged */ + for (n = 0; n < np; ++n) { - int n; - char* cipherName; - ++globalCipherCount; - cipherName = globalCipherNameTable[globalCipherCount].m_name; - strcpy(cipherName, desc->m_name); + cipherParams[n] = params[n]; + cipherParams[n].m_name = sqlite3_mprintf("%s", params[n].m_name); + if (!cipherParams[n].m_name) + { + cipherParams[n].m_name = globalSentinelName; + mcFreeCipherParams(cipherParams); + return SQLITE_NOMEM; + } + } + /* Add sentinel */ + cipherParams[n] = params[n]; + cipherParams[n].m_name = globalSentinelName; + + ++globalCipherCount; + cipherName = globalCipherNameTable[globalCipherCount].m_name; + strcpy(cipherName, desc->m_name); - globalCodecDescriptorTable[globalCipherCount - 1] = *desc; - globalCodecDescriptorTable[globalCipherCount - 1].m_name = cipherName; + globalCodecDescriptorTable[globalCipherCount - 1] = *desc; + globalCodecDescriptorTable[globalCipherCount - 1].m_name = cipherName; - globalCodecParameterTable[globalCipherCount].m_name = cipherName; - globalCodecParameterTable[globalCipherCount].m_id = globalCipherCount; - globalCodecParameterTable[globalCipherCount].m_params = cipherParams; + globalCodecParameterTable[globalCipherCount].m_name = cipherName; + globalCodecParameterTable[globalCipherCount].m_id = globalCipherCount; + globalCodecParameterTable[globalCipherCount].m_params = cipherParams; - /* Copy parameters */ - for (n = 0; n < np; ++n) + /* Make cipher default, if requested */ + if (makeDefault) + { + CipherParams* param = globalCodecParameterTable[0].m_params; + for (; param->m_name[0] != 0; ++param) { - char* paramName = (char*) sqlite3_malloc((int)strlen(params[n].m_name) + 1); - strcpy(paramName, params[n].m_name); - cipherParams[n] = params[n]; - cipherParams[n].m_name = paramName; + if (sqlite3_stricmp("cipher", param->m_name) == 0) break; } - /* Add sentinel */ - cipherParams[n] = params[n]; - cipherParams[n].m_name = globalSentinelName; - - /* Make cipher default, if requested */ - if (makeDefault) + if (param->m_name[0] != 0) { - CipherParams* param = globalCodecParameterTable[0].m_params; - for (; param->m_name[0] != 0; ++param) - { - if (sqlite3_stricmp("cipher", param->m_name) == 0) break; - } - if (param->m_name[0] != 0) - { - param->m_value = param->m_default = globalCipherCount; - } + param->m_value = param->m_default = globalCipherCount; } - - rc = SQLITE_OK; - } - else - { - rc = SQLITE_NOMEM; } - return rc; + return SQLITE_OK; } SQLITE_API int @@ -682,13 +688,7 @@ sqlite3mcTermCipherTables() { if (globalCodecParameterTable[n].m_name[0] != 0) { - int k; - CipherParams* params = globalCodecParameterTable[n].m_params; - for (k = 0; params[k].m_name[0] != 0; ++k) - { - sqlite3_free((char*) params[k].m_name); - } - sqlite3_free(globalCodecParameterTable[n].m_params); + mcFreeCipherParams(globalCodecParameterTable[n].m_params); } } globalCipherCount = 0; @@ -755,6 +755,11 @@ sqlite3mc_initialize(const char* arg) rc = sqlite3mc_vfs_create(NULL, 1); } } + if (rc != SQLITE_OK) + { + /* SQLite does not call sqlite3mc_shutdown after a failed initialization */ + sqlite3mcTermCipherTables(); + } return rc; } diff --git a/test/oomtest.c b/test/oomtest.c new file mode 100644 index 00000000..4a977d74 --- /dev/null +++ b/test/oomtest.c @@ -0,0 +1,138 @@ +/* +** Test for allocation failures: every allocation made by sqlite3_initialize +** and by sqlite3_open fails in turn. Each call must return SQLITE_OK or +** SQLITE_NOMEM, a sqlite3_initialize retried after a failure must build the +** same cipher tables as a clean start, and a connection that opened must +** carry its cipher configuration. +*/ + +#include "sqlite3mc.c" +#include +#include +#include + +static sqlite3_mem_methods realMethods; +static int countdown = -1; /* the allocation that reaches 0 fails */ + +static int failNow(void) +{ + return countdown > 0 && --countdown == 0; +} + +static void* faultMalloc(int n) +{ + return failNow() ? 0 : realMethods.xMalloc(n); +} + +static void* faultRealloc(void* p, int n) +{ + return failNow() ? 0 : realMethods.xRealloc(p, n); +} + +static void installFaultMethods(void) +{ + sqlite3_mem_methods methods = realMethods; + methods.xMalloc = faultMalloc; + methods.xRealloc = faultRealloc; + sqlite3_config(SQLITE_CONFIG_MALLOC, &methods); +} + +static char* fingerprint(void) +{ + int j, k; + sqlite3_str* str = sqlite3_str_new(NULL); + sqlite3_str_appendf(str, "%d ciphers", globalCipherCount); + for (j = 0; globalCodecParameterTable[j].m_name[0] != 0; ++j) + { + CipherParams* params = globalCodecParameterTable[j].m_params; + sqlite3_str_appendf(str, "; %s:", globalCodecParameterTable[j].m_name); + for (k = 0; params[k].m_name[0] != 0; ++k) + { + sqlite3_str_appendf(str, " %s=%d/%d", params[k].m_name, params[k].m_value, params[k].m_default); + } + } + return sqlite3_str_finish(str); +} + +static int unexpectedResult(const char* call, int nth, int rc) +{ + if (rc == SQLITE_OK || rc == SQLITE_NOMEM) return 0; + printf("%s, allocation %d failed: returned %d\n", call, nth, rc); + return 1; +} + +int main(void) +{ + char* baseline; + char* expected; + int nth; + int reached; + int rc; + int nInit; + int nOpen; + int failures = 0; + + sqlite3_config(SQLITE_CONFIG_GETMALLOC, &realMethods); + installFaultMethods(); + if (sqlite3_initialize() != SQLITE_OK) return 1; + baseline = fingerprint(); + if (baseline == NULL) return 1; + /* SQLite must not hold allocations across the sqlite3_shutdown calls below */ + expected = (char*) malloc(strlen(baseline) + 1); + if (expected == NULL) return 1; + strcpy(expected, baseline); + sqlite3_free(baseline); + + for (nth = 1;; ++nth) + { + char* actual; + sqlite3_shutdown(); + installFaultMethods(); + countdown = nth; + rc = sqlite3_initialize(); + reached = countdown == 0; + countdown = -1; + if (!reached) break; + failures += unexpectedResult("sqlite3_initialize", nth, rc); + if (rc != SQLITE_OK && sqlite3_initialize() != SQLITE_OK) + { + printf("sqlite3_initialize, allocation %d failed: retry failed\n", nth); + ++failures; + continue; + } + actual = fingerprint(); + if (actual == NULL || strcmp(actual, expected) != 0) + { + printf("sqlite3_initialize, allocation %d failed: cipher tables differ\n expected %s\n actual %s\n", + nth, expected, actual ? actual : "(out of memory)"); + ++failures; + } + sqlite3_free(actual); + } + /* The last sqlite3_initialize succeeded, so the library stays initialized for the next loop */ + nInit = nth - 1; + + for (nth = 1;; ++nth) + { + sqlite3* db = 0; + countdown = nth; + rc = sqlite3_open(":memory:", &db); + reached = countdown == 0; + countdown = -1; + failures += unexpectedResult("sqlite3_open", nth, rc); + if (rc == SQLITE_OK && sqlite3mc_config(db, "cipher", -1) != sqlite3mc_config(NULL, "cipher", -1)) + { + printf("sqlite3_open, allocation %d failed: connection has no cipher configuration\n", nth); + ++failures; + } + sqlite3_close(db); + if (!reached) break; + } + nOpen = nth - 1; + + sqlite3_shutdown(); + free(expected); + printf("%d allocations in sqlite3_initialize and %d in sqlite3_open failed in turn, %d failures\n", + nInit, nOpen, failures); + return failures != 0; +}