diff --git a/google/cloud/odbc/bq_driver/internal/odbc_sql_fetch.cc b/google/cloud/odbc/bq_driver/internal/odbc_sql_fetch.cc index 69618d66f8..2296b3e598 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_sql_fetch.cc +++ b/google/cloud/odbc/bq_driver/internal/odbc_sql_fetch.cc @@ -237,6 +237,85 @@ StatusRecord WriteRowset(ResultSet const& result_set, int const rowset_size, return StatusRecord::Ok(); } +StatusRecord WriteRowset(StatementHandle& stmt_handle, int const rowset_size, + DescriptorHandle& ard, DescriptorHandle& ird) { + if (rowset_size <= 0) { + LOG(ERROR) << "WriteRowset:: rowset_size should not be <= 0"; + return StatusRecord{SQLStates::k_HY000(), "rowset_size should not be <= 0"}; + } + + int row_counter = 0; + SQLUSMALLINT* row_status_ptr = ird.GetHeaderRecord().array_status_ptr; + + while (row_counter < rowset_size) { + ResultSet& result_set = stmt_handle.GetResultSet(); + if (result_set.cursor >= static_cast(result_set.rows.size())) { + StatusRecord next_page_status = FetchNextResultSet(stmt_handle); + if (!next_page_status.ok()) { + if (next_page_status.sql_state == SQLStates::k_SQL_NO_DATA()) { + break; + } + LOG(ERROR) << "WriteRowset::FetchNextResultSet:: " + << next_page_status.message; + if (row_counter == 0) { + return next_page_status; + } + break; + } + ResultSet& updated_rs = stmt_handle.GetResultSet(); + if (updated_rs.rows.empty()) { + break; + } + updated_rs.cursor++; + } + + ResultSet& current_rs = stmt_handle.GetResultSet(); + if (current_rs.cursor < 0 || + current_rs.cursor >= static_cast(current_rs.rows.size())) { + break; + } + + StatusRecord status_record = + WriteDSRow(current_rs.rows[current_rs.cursor], current_rs.row_schema, + ard, row_counter); + if (!status_record.ok()) { + LOG(ERROR) << "WriteRowset::WriteDSRow:: " << status_record.message; + return status_record; + } + + if (row_status_ptr) { + row_status_ptr[row_counter] = SQL_ROW_SUCCESS; + } + + row_counter++; + current_rs.cursor++; + } + + ResultSet& final_rs = stmt_handle.GetResultSet(); + if (row_counter > 0) { + final_rs.cursor--; + } + + // Mark unused rows + if (row_status_ptr) { + for (int i = row_counter; i < rowset_size; i++) { + row_status_ptr[i] = SQL_ROW_NOROW; + } + } + + SQLULEN* rows_processed_ptr = ird.GetHeaderRecord().rows_processed_ptr; + if (rows_processed_ptr) { + *rows_processed_ptr = row_counter; + } + + if (row_counter == 0) { + return StatusRecord( + {SQLStates::k_SQL_NO_DATA(), "No more data to return."}); + } + + return StatusRecord::Ok(); +} + StatusRecord FetchNextResultSet(StatementHandle& stmt_handle) { // We need to return only the top `SQL_ATTR_MAX_ROWS` number of rows auto max_rows_status = stmt_handle.GetAttribute(SQL_ATTR_MAX_ROWS); @@ -258,13 +337,17 @@ StatusRecord FetchNextResultSet(StatementHandle& stmt_handle) { if (stmt_handle.WasHtapiEnabled()) { StatusRecord read_status = ReadNextResultsFromStream(stmt_handle); if (!read_status.ok()) { - LOG(ERROR) << "ReadNextResultsFromStream:: " << read_status.message; + if (read_status.sql_state != SQLStates::k_SQL_NO_DATA()) { + LOG(ERROR) << "ReadNextResultsFromStream:: " << read_status.message; + } return read_status; } } else { StatusRecord read_status = FetchNextPageResultSet(stmt_handle); if (!read_status.ok()) { - LOG(ERROR) << "FetchNextPageResultSet:: " << read_status.message; + if (read_status.sql_state != SQLStates::k_SQL_NO_DATA()) { + LOG(ERROR) << "FetchNextPageResultSet:: " << read_status.message; + } return read_status; } } @@ -272,7 +355,9 @@ StatusRecord FetchNextResultSet(StatementHandle& stmt_handle) { StatusRecord read_status = FetchNextPageResultSet(stmt_handle); if (!read_status.ok()) { - LOG(ERROR) << "FetchNextPageResultSet:: " << read_status.message; + if (read_status.sql_state != SQLStates::k_SQL_NO_DATA()) { + LOG(ERROR) << "FetchNextPageResultSet:: " << read_status.message; + } return read_status; } #endif // (!defined(_WIN32) || defined(_WIN64)) && !defined(NO_ARROW) diff --git a/google/cloud/odbc/bq_driver/internal/odbc_sql_fetch.h b/google/cloud/odbc/bq_driver/internal/odbc_sql_fetch.h index 1a01b551a0..fa5576a6d8 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_sql_fetch.h +++ b/google/cloud/odbc/bq_driver/internal/odbc_sql_fetch.h @@ -21,7 +21,12 @@ namespace google::cloud::odbc_bq_driver_internal { -// Writes rowset_size number of rows to the columns bound by the application +class StatementHandle; +google::cloud::odbc_internal::StatusRecord WriteRowset( + StatementHandle& stmt_handle, int rowset_size, DescriptorHandle& ard, + DescriptorHandle& ird); + +// Overload for unit testing with ResultSet only google::cloud::odbc_internal::StatusRecord WriteRowset( ResultSet const& result_set, int rowset_size, DescriptorHandle& ard, DescriptorHandle& ird); diff --git a/google/cloud/odbc/bq_driver/internal/odbc_stmt_handle.cc b/google/cloud/odbc/bq_driver/internal/odbc_stmt_handle.cc index 701ba5fade..4cc2881b58 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_stmt_handle.cc +++ b/google/cloud/odbc/bq_driver/internal/odbc_stmt_handle.cc @@ -21,6 +21,7 @@ #include "google/cloud/odbc/bq_driver/internal/odbc_sql_type_info.h" #include "google/cloud/odbc/bq_driver/internal/odbc_transactions.h" #include "google/cloud/odbc/bq_driver/internal/trace_utils.h" +#include "google/cloud/odbc/bq_driver/internal/utils.h" #include "google/cloud/odbc/internal/status_record_or.h" namespace google::cloud::odbc_bq_driver_internal { @@ -213,8 +214,14 @@ StatusRecord StatementHandle::PrepareQuery(std::string const& query) { } ConnectionHandle& conn_handle = *GetConnectionHandle(); + std::string processed_query = query; + auto noscan_attr = GetAttribute(SQL_ATTR_NOSCAN); + if (!noscan_attr || *noscan_attr != SQL_NOSCAN_ON) { + processed_query = TranslateOdbcEscapeSequences(query); + } + Job req; - req.configuration.query.query = query; + req.configuration.query.query = processed_query; req.configuration.query.use_query_cache = conn_handle.GetDsn().is_query_cache; req.configuration.dry_run = true; req.configuration.query.use_legacy_sql = @@ -233,7 +240,7 @@ StatusRecord StatementHandle::PrepareQuery(std::string const& query) { // to be used during table creation. Subsequent operations on the table will // automatically use the KMS key without the application sending it. std::string kms_key_name = conn_handle.GetDsn().kms_key_name; - if (IsInsertQuery(query) || IsSelectQuery(query)) { + if (IsInsertQuery(processed_query) || IsSelectQuery(processed_query)) { if (!kms_key_name.empty()) { req.configuration.query.destination_encryption_configuration .kms_key_name = kms_key_name; @@ -243,8 +250,8 @@ StatusRecord StatementHandle::PrepareQuery(std::string const& query) { if (!conn_handle.GetDsn().is_bq_legacy_sql) { // Detect POSITIONAL (`?`) and NAMED (`[:@]\w+`) parameter markers using // RE2 instead of a manual character scan. - bool has_positional = re2::RE2::PartialMatch(query, R"(\?)"); - bool has_named = re2::RE2::PartialMatch(query, R"([:@]\w+)"); + bool has_positional = re2::RE2::PartialMatch(processed_query, R"(\?)"); + bool has_named = re2::RE2::PartialMatch(processed_query, R"([:@]\w+)"); if (has_positional) { req.configuration.query.parameter_mode = "POSITIONAL"; } @@ -330,7 +337,7 @@ StatusRecord StatementHandle::PrepareQuery(std::string const& query) { conn_handle.SetSessionId(response->statistics.session_info.session_id); } - query_str_ = query; + query_str_ = processed_query; prepared_job_ = *response; return StatusRecord::Ok(); } @@ -551,6 +558,10 @@ StatusRecord StatementHandle::PopulateIpd(DescriptorHandle& handle, void StatementHandle::CloseCursor() { ResultSet result_set; result_set_ = result_set; +#if (!defined(_WIN32) || defined(_WIN64)) && !defined(NO_ARROW) + ClearReadRowsStream(); + ClearReadRowsIterator(); +#endif if (StatementPrepared()) { SetStmtState(StmtStates::kStatementPrepared); } else { diff --git a/google/cloud/odbc/bq_driver/internal/odbc_stmt_handle.h b/google/cloud/odbc/bq_driver/internal/odbc_stmt_handle.h index 08e70a3d5e..69c5df367d 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_stmt_handle.h +++ b/google/cloud/odbc/bq_driver/internal/odbc_stmt_handle.h @@ -141,6 +141,7 @@ class StatementHandle : public Handle { StreamRange<::google::cloud::bigquery::storage::v1::ReadRowsResponse> stream_range) { read_rows_stream_ = std::move(stream_range); + read_rows_iterator_.reset(); } std::optional< diff --git a/google/cloud/odbc/bq_driver/internal/utils.cc b/google/cloud/odbc/bq_driver/internal/utils.cc index 32bc9d1c73..b306a77a43 100644 --- a/google/cloud/odbc/bq_driver/internal/utils.cc +++ b/google/cloud/odbc/bq_driver/internal/utils.cc @@ -20,8 +20,10 @@ #include "google/cloud/odbc/bq_driver/internal/trace_utils.h" #include "google/cloud/odbc/bq_driver/internal/utils.h" #include "google/cloud/internal/getenv.h" +#include "absl/strings/match.h" #include #include +#include #include #include #include @@ -1174,9 +1176,11 @@ StatusRecord PopulateOutputConnectionString(SQLCHAR* out_conn_str, auto out_str_len = out_tmp_str.length(); if (out_str_len >= out_conn_str_buflen) { - strncpy(reinterpret_cast(out_conn_str), out_tmp_str.c_str(), - out_conn_str_buflen - 1); - out_conn_str[out_conn_str_buflen - 1] = '\0'; + if (out_conn_str && out_conn_str_buflen > 0) { + strncpy(reinterpret_cast(out_conn_str), out_tmp_str.c_str(), + out_conn_str_buflen - 1); + out_conn_str[out_conn_str_buflen - 1] = '\0'; + } if (out_conn_str_len) { *out_conn_str_len = out_str_len; } @@ -1184,9 +1188,11 @@ StatusRecord PopulateOutputConnectionString(SQLCHAR* out_conn_str, << "PopulateOutputConnectionString:: String data, right truncated"; return StatusRecord{SQLStates::k_01004(), "String data, right truncated"}; } - strncpy(reinterpret_cast(out_conn_str), out_tmp_str.c_str(), - out_tmp_str.length()); - out_conn_str[out_tmp_str.length()] = '\0'; + if (out_conn_str && out_conn_str_buflen > 0) { + strncpy(reinterpret_cast(out_conn_str), out_tmp_str.c_str(), + out_tmp_str.length()); + out_conn_str[out_tmp_str.length()] = '\0'; + } if (out_conn_str_len) { *out_conn_str_len = out_tmp_str.length(); } @@ -1505,4 +1511,267 @@ StatusRecord NormalizeOAuthMechanism(Section& section) { return StatusRecord::Ok(); } #endif // _WIN32 + +namespace { + +std::string_view TrimWhitespace(std::string_view sv) { + while (!sv.empty() && std::isspace(static_cast(sv.front()))) { + sv.remove_prefix(1); + } + while (!sv.empty() && std::isspace(static_cast(sv.back()))) { + sv.remove_suffix(1); + } + return sv; +} + +bool ExtractQuotedLiteral(std::string_view sv, std::string& out_literal) { + sv = TrimWhitespace(sv); + if (sv.size() >= 2) { + char quote = sv.front(); + if ((quote == '\'' || quote == '"') && sv.back() == quote) { + out_literal = std::string(sv.substr(1, sv.size() - 2)); + return true; + } + } + return false; +} + +bool StartsWithIgnoreCase(std::string_view sv, std::string_view prefix) { + if (sv.size() < prefix.size()) return false; + for (size_t i = 0; i < prefix.size(); ++i) { + if (std::tolower(static_cast(sv[i])) != + std::tolower(static_cast(prefix[i]))) { + return false; + } + } + return true; +} + +std::string ProcessEscapeContent(std::string_view content) { + content = TrimWhitespace(content); + if (content.empty()) return "{}"; + + // Check for {ts '...'} / {TS '...'} + if (StartsWithIgnoreCase(content, "ts") && + (content.size() == 2 || + std::isspace(static_cast(content[2])))) { + std::string_view rest = TrimWhitespace(content.substr(2)); + std::string literal; + if (ExtractQuotedLiteral(rest, literal)) { + return "TIMESTAMP '" + literal + "'"; + } + } + + // Check for {d '...'} / {D '...'} + if (StartsWithIgnoreCase(content, "d") && + (content.size() == 1 || + std::isspace(static_cast(content[1])))) { + std::string_view rest = TrimWhitespace(content.substr(1)); + std::string literal; + if (ExtractQuotedLiteral(rest, literal)) { + return "DATE '" + literal + "'"; + } + } + + // Check for {t '...'} / {T '...'} + if (StartsWithIgnoreCase(content, "t") && + (content.size() == 1 || + std::isspace(static_cast(content[1])))) { + std::string_view rest = TrimWhitespace(content.substr(1)); + std::string literal; + if (ExtractQuotedLiteral(rest, literal)) { + return "TIME '" + literal + "'"; + } + } + + // Check for {escape '...'} + if (StartsWithIgnoreCase(content, "escape") && + (content.size() == 6 || + std::isspace(static_cast(content[6])))) { + std::string_view rest = TrimWhitespace(content.substr(6)); + std::string literal; + if (ExtractQuotedLiteral(rest, literal)) { + return "ESCAPE '" + literal + "'"; + } + } + + // Check for {guid '...'} + if (StartsWithIgnoreCase(content, "guid") && + (content.size() == 4 || + std::isspace(static_cast(content[4])))) { + std::string_view rest = TrimWhitespace(content.substr(4)); + std::string literal; + if (ExtractQuotedLiteral(rest, literal)) { + return "'" + literal + "'"; + } + } + + // Check for {oj ...} -> outer join + if (StartsWithIgnoreCase(content, "oj") && + (content.size() == 2 || + std::isspace(static_cast(content[2])))) { + std::string_view rest = TrimWhitespace(content.substr(2)); + return std::string(rest); + } + + // Check for {fn ...} -> scalar function + if (StartsWithIgnoreCase(content, "fn") && + (content.size() == 2 || + std::isspace(static_cast(content[2])))) { + std::string_view rest = TrimWhitespace(content.substr(2)); + return std::string(rest); + } + + // If none matched, return original braced text + return "{" + std::string(content) + "}"; +} + +} // namespace + +std::string TranslateOdbcEscapeSequences(std::string const& sql) { + // Fast path: if there are no braces, no ODBC escape sequence can be present. + if (!absl::StrContains(sql, '{')) { + return sql; + } + + std::string current = sql; + constexpr int kMaxPasses = 10; + for (int pass = 0; pass < kMaxPasses; ++pass) { + std::string result; + result.reserve(current.size()); + bool changed = false; + + size_t i = 0; + size_t const n = current.size(); + + while (i < n) { + char c = current[i]; + + // Check for single-line comment: -- or # + if ((c == '-' && i + 1 < n && current[i + 1] == '-') || c == '#') { + size_t comment_end = current.find('\n', i); + if (comment_end == std::string::npos) { + result.append(current, i, n - i); + break; + } + result.append(current, i, comment_end - i + 1); + i = comment_end + 1; + continue; + } + + // Check for multi-line comment: /* ... */ + if (c == '/' && i + 1 < n && current[i + 1] == '*') { + size_t comment_end = current.find("*/", i + 2); + if (comment_end == std::string::npos) { + result.append(current, i, n - i); + break; + } + result.append(current, i, comment_end + 2 - i); + i = comment_end + 2; + continue; + } + + // Check for string literals and quoted identifiers: '...', "...", `...` + if (c == '\'' || c == '"' || c == '`') { + char const quote_char = c; + result.push_back(c); + ++i; + while (i < n) { + char sc = current[i]; + result.push_back(sc); + if (sc == quote_char) { + if (i + 1 < n && current[i + 1] == quote_char) { + // Escaped quote (e.g. '') + ++i; + result.push_back(current[i]); + ++i; + } else { + ++i; + break; + } + } else if (sc == '\\' && i + 1 < n && current[i + 1] == '\\') { + ++i; + result.push_back(current[i]); + ++i; + } else { + ++i; + } + } + continue; + } + + // Check for opening brace `{` + if (c == '{') { + // Find matching `}` while respecting quotes inside + size_t start_brace = i; + size_t j = i + 1; + int brace_depth = 1; + bool matched = false; + + while (j < n && brace_depth > 0) { + char jc = current[j]; + if (jc == '\'' || jc == '"' || jc == '`') { + char const quote_char = jc; + ++j; + while (j < n) { + if (current[j] == quote_char) { + if (j + 1 < n && current[j + 1] == quote_char) { + j += 2; + } else { + ++j; + break; + } + } else if (current[j] == '\\' && j + 1 < n && + current[j + 1] == '\\') { + j += 2; + } else { + ++j; + } + } + } else if (jc == '{') { + ++brace_depth; + ++j; + } else if (jc == '}') { + --brace_depth; + if (brace_depth == 0) { + matched = true; + break; + } + ++j; + } else { + ++j; + } + } + + if (matched) { + std::string_view const inner = std::string_view{current}.substr( + start_brace + 1, j - start_brace - 1); + std::string replaced = ProcessEscapeContent(inner); + if (replaced != current.substr(start_brace, j - start_brace + 1)) { + changed = true; + } + result.append(replaced); + i = j + 1; + continue; + } + + // No matching brace found, output '{' + result.push_back(c); + ++i; + continue; + } + + result.push_back(c); + ++i; + } + + if (!changed) { + return result; + } + current = std::move(result); + } + + return current; +} + } // namespace google::cloud::odbc_bq_driver_internal diff --git a/google/cloud/odbc/bq_driver/internal/utils.h b/google/cloud/odbc/bq_driver/internal/utils.h index fa0c78e56b..32585da845 100644 --- a/google/cloud/odbc/bq_driver/internal/utils.h +++ b/google/cloud/odbc/bq_driver/internal/utils.h @@ -486,6 +486,8 @@ odbc_internal::StatusRecordOr ParseStringToInteger( std::string const& input); std::string GetLocationfromPSC(std::string const& psc); + +std::string TranslateOdbcEscapeSequences(std::string const& sql); } // namespace google::cloud::odbc_bq_driver_internal #endif // CPP_BIGQUERY_ODBC_GOOGLE_CLOUD_ODBC_BQ_DRIVER_INTERNAL_UTILS_H diff --git a/google/cloud/odbc/bq_driver/internal/utils_test.cc b/google/cloud/odbc/bq_driver/internal/utils_test.cc index 9af16db60e..45ad70726a 100644 --- a/google/cloud/odbc/bq_driver/internal/utils_test.cc +++ b/google/cloud/odbc/bq_driver/internal/utils_test.cc @@ -1057,4 +1057,50 @@ TEST(EscapeOdbcPattern, EscapedNameMatchesOnlyItself) { EXPECT_TRUE(re2::RE2::FullMatch("ODBCxTESTyDATASET", *unescaped)); } +TEST(TranslateOdbcEscapeSequences, DatetimeLiterals) { + EXPECT_EQ(TranslateOdbcEscapeSequences( + "SELECT {ts '2014-02-20 09:34:06.000'} AS ts_col"), + "SELECT TIMESTAMP '2014-02-20 09:34:06.000' AS ts_col"); + EXPECT_EQ(TranslateOdbcEscapeSequences( + "SELECT {TS '2014-02-20 09:34:06'} AS ts_col"), + "SELECT TIMESTAMP '2014-02-20 09:34:06' AS ts_col"); + EXPECT_EQ(TranslateOdbcEscapeSequences( + "SELECT * FROM t WHERE d = {d '2023-01-01'}"), + "SELECT * FROM t WHERE d = DATE '2023-01-01'"); + EXPECT_EQ(TranslateOdbcEscapeSequences( + "SELECT * FROM t WHERE d = {D '2023-01-01'}"), + "SELECT * FROM t WHERE d = DATE '2023-01-01'"); + EXPECT_EQ( + TranslateOdbcEscapeSequences("SELECT * FROM t WHERE tm = {t '12:30:00'}"), + "SELECT * FROM t WHERE tm = TIME '12:30:00'"); + EXPECT_EQ( + TranslateOdbcEscapeSequences("SELECT * FROM t WHERE tm = {T '12:30:00'}"), + "SELECT * FROM t WHERE tm = TIME '12:30:00'"); +} + +TEST(TranslateOdbcEscapeSequences, EscapeAndOuterJoin) { + EXPECT_EQ(TranslateOdbcEscapeSequences( + "SELECT * FROM t WHERE name LIKE '\\%AAA%' {escape '\\'}"), + "SELECT * FROM t WHERE name LIKE '\\%AAA%' ESCAPE '\\'"); + EXPECT_EQ( + TranslateOdbcEscapeSequences( + "SELECT * FROM {oj Customers LEFT OUTER JOIN Orders ON c.id=o.id}"), + "SELECT * FROM Customers LEFT OUTER JOIN Orders ON c.id=o.id"); +} + +TEST(TranslateOdbcEscapeSequences, PreservesStringsAndCommentsAndStructs) { + EXPECT_EQ( + TranslateOdbcEscapeSequences( + "SELECT '{ts 2020-01-01}' AS col, {ts '2020-01-01 00:00:00'} AS ts"), + "SELECT '{ts 2020-01-01}' AS col, TIMESTAMP '2020-01-01 00:00:00' AS ts"); + EXPECT_EQ(TranslateOdbcEscapeSequences("SELECT `table_{ts}` FROM tbl"), + "SELECT `table_{ts}` FROM tbl"); + EXPECT_EQ(TranslateOdbcEscapeSequences("SELECT /* {ts '2020'} */ 1"), + "SELECT /* {ts '2020'} */ 1"); + EXPECT_EQ(TranslateOdbcEscapeSequences("SELECT -- {ts '2020'}\n 1"), + "SELECT -- {ts '2020'}\n 1"); + EXPECT_EQ(TranslateOdbcEscapeSequences("SELECT {'a': 1} AS struct_col"), + "SELECT {'a': 1} AS struct_col"); +} + } // namespace google::cloud::odbc_bq_driver_internal diff --git a/google/cloud/odbc/bq_driver/odbc_api.cc b/google/cloud/odbc/bq_driver/odbc_api.cc index d10c3d72f2..df450efcfc 100644 --- a/google/cloud/odbc/bq_driver/odbc_api.cc +++ b/google/cloud/odbc/bq_driver/odbc_api.cc @@ -268,33 +268,44 @@ SQLRETURN SQL_API SQLDriverConnectW( if (inConnectionStringLen && inConnectionStringLen != SQL_NTS) inConnectionStringLen = utf8_in_connection_str->length(); } - // outConnectionString is an output value that is not populated by the user. - // This should not be unicode converted if it is empty. Instead we send a - // SQLCHAR empty value directly to the internal function. - SQLCHAR* out_conn_str = reinterpret_cast(outConnectionString); + SQLCHAR out_conn_str_buf[kBufferLength] = {0}; SQLSMALLINT out_conn_str_len = 0; // Call to internal common function for SQLDriverConnect and // SQLDriverConnectW in odbc_connection.h. rc = google::cloud::odbc_bq_driver::SQLDriverConnectInternal( connectionHandle, windowHandle, sqlchar_in_connection_str, - inConnectionStringLen, out_conn_str, outConnectionStringBufferLen, - &out_conn_str_len, driverCompletion); + inConnectionStringLen, out_conn_str_buf, + static_cast(sizeof(out_conn_str_buf)), &out_conn_str_len, + driverCompletion); // Handle Unicode conversion of output parameters. - if (SQL_SUCCEEDED(rc) && outConnectionString) { + if (SQL_SUCCEEDED(rc) && outConnectionString && + outConnectionStringBufferLen > 0) { StatusRecordOr utf16_out_conn_str; if (out_conn_str_len > 0) { - utf16_out_conn_str = Utf8ToUtf16((char*)out_conn_str); + utf16_out_conn_str = Utf8ToUtf16((char*)out_conn_str_buf); } else { - std::string val(ToCharStr(out_conn_str)); + std::string val(ToCharStr(out_conn_str_buf)); utf16_out_conn_str = Utf8ToUtf16(val); } if (!utf16_out_conn_str) { return utf16_out_conn_str.GetCalculatedReturnCode(); } - WriteWideToWireBuffer(*utf16_out_conn_str, outConnectionString, - out_conn_str_len); + size_t const dest_chars = static_cast(outConnectionStringBufferLen); + size_t const to_copy = + std::min(utf16_out_conn_str->size(), dest_chars - 1); + WriteWideToWireBuffer(*utf16_out_conn_str, outConnectionString, to_copy, + /*null_terminate=*/true); + + if (utf16_out_conn_str->size() >= dest_chars) { + rc = SQL_SUCCESS_WITH_INFO; + auto* dbc_handle = reinterpret_cast(connectionHandle); + if (dbc_handle) { + dbc_handle->GetDiagnostics().AddStatusRecord( + StatusRecord{SQLStates::k_01004(), "String data, right truncated"}); + } + } } if (outConnectionStringLen) *outConnectionStringLen = out_conn_str_len; @@ -544,33 +555,6 @@ SQLRETURN SQL_API SQLConnectW(SQLHDBC connectionHandle, SQLWCHAR* serverName, ToSqlChar(""), w_user_name_len, ToSqlChar(""), w_auth_str_len); } - // Handle Unicode conversion of output parameters. - StatusRecordOr utf16_server_name = - Utf8ToUtf16(*utf8_server_name); - if (!utf16_server_name) { - return utf16_server_name.GetCalculatedReturnCode(); - } - serverNameLen = utf16_server_name->length(); - WriteWideToWireBuffer(*utf16_server_name, serverName, serverNameLen); - - if (w_user_name_len > 0) { - StatusRecordOr utf16_user_name = Utf8ToUtf16(*utf8_user_name); - if (!utf16_user_name) { - return utf16_user_name.GetCalculatedReturnCode(); - } - userNameLen = utf16_user_name->length(); - WriteWideToWireBuffer(*utf16_user_name, userName, userNameLen); - } - - if (w_auth_str_len > 0) { - StatusRecordOr utf16_auth_str = Utf8ToUtf16(*utf8_auth_str); - if (!utf16_auth_str) { - return utf16_auth_str.GetCalculatedReturnCode(); - } - authStringLen = utf16_auth_str->length(); - WriteWideToWireBuffer(*utf16_auth_str, authString, authStringLen); - } - return rc; } @@ -2590,7 +2574,7 @@ SQLRETURN SQL_API SQLGetDiagRecW(SQLSMALLINT handleType, SQLHANDLE handle, SQLRETURN rc = SQL_SUCCESS; SQLRETURN status; SQLCHAR sql_state_buffer[kBufferLength] = {0}; - SQLCHAR* message_text_buffer = reinterpret_cast(messageText); + SQLCHAR message_text_buffer[kBufferLength] = {0}; SQLSMALLINT message_text_buffer_len = 0; InitializeTracing("SQLGetDiagRecW"); @@ -2604,7 +2588,8 @@ SQLRETURN SQL_API SQLGetDiagRecW(SQLSMALLINT handleType, SQLHANDLE handle, // in odbc_diagnostics.h. rc = google::cloud::odbc_bq_driver::SQLGetDiagRecInternal( handleType, handle, recNumber, sql_state_buffer, nativeError, - message_text_buffer, messageTextBufferLen, &message_text_buffer_len); + message_text_buffer, sizeof(message_text_buffer), + &message_text_buffer_len); // Handle Unicode conversion of output parameters. @@ -2614,7 +2599,9 @@ SQLRETURN SQL_API SQLGetDiagRecW(SQLSMALLINT handleType, SQLHANDLE handle, if (!utf16_sql_state) { return utf16_sql_state.GetCalculatedReturnCode(); } - WriteWideToWireBuffer(*utf16_sql_state, sqlState, utf16_sql_state->size(), + std::memset(sqlState, '\0', 6 * WireWcharSize()); + WriteWideToWireBuffer(*utf16_sql_state, sqlState, + std::min(utf16_sql_state->size(), 5), /*null_terminate=*/true); } diff --git a/google/cloud/odbc/bq_driver/odbc_sql_results.cc b/google/cloud/odbc/bq_driver/odbc_sql_results.cc index 5db5dca66e..10d1c91bd1 100644 --- a/google/cloud/odbc/bq_driver/odbc_sql_results.cc +++ b/google/cloud/odbc/bq_driver/odbc_sql_results.cc @@ -197,12 +197,14 @@ SQLRETURN SQLFetchInternal(SQLHSTMT statement_handle) { result_set.translated_data.row_offset = 0; result_set.translated_data.data.clear(); result_set.translated_data.last_target_c_type = 0; - if (result_set.cursor >= result_set.rows.size()) { + if (result_set.cursor >= static_cast(result_set.rows.size())) { LOG(INFO) << "SQLFetch:: cursor: " << result_set.cursor << " is >= result set size: " << result_set.rows.size(); StatusRecord next_page_status = FetchNextResultSet(handle); if (!next_page_status.ok()) { - LOG(ERROR) << "SQLFetch:: " << next_page_status.message; + if (next_page_status.sql_state != SQLStates::k_SQL_NO_DATA()) { + LOG(ERROR) << "SQLFetch:: " << next_page_status.message; + } return LogAndReturnCode(handle, next_page_status); } result_set = handle.GetResultSet(); @@ -214,7 +216,7 @@ SQLRETURN SQLFetchInternal(SQLHSTMT statement_handle) { rowset_size = 1; } DescriptorHandle& ird = handle.GetDescriptorHandle(DescriptorType::kIRD); - StatusRecord status_record = WriteRowset(result_set, rowset_size, ard, ird); + StatusRecord status_record = WriteRowset(handle, rowset_size, ard, ird); return LogAndReturnCode(handle, status_record); } @@ -262,26 +264,25 @@ SQLRETURN SQLFetchScrollInternal(SQLHSTMT statement_handle, DescriptorHandle& ard = handle.GetDescriptorHandle(DescriptorType::kARD); ResultSet& result_set = handle.GetResultSet(); - if (result_set.cursor >= result_set.rows.size() - 1 && - !handle.GetPagingInfo().page_token.empty()) { - StatusRecord status_record = FetchNextResultSet(handle); - if (!status_record.ok()) { - LOG(ERROR) << "SQLFetchScroll::FetchNextResultSet:: " - << status_record.message; - return LogAndReturnCode(handle, status_record); - } - } result_set.translated_data.row_offset = 0; + result_set.translated_data.data.clear(); + result_set.translated_data.last_target_c_type = 0; // Compute new row position based on fetch type switch (fetch_orientation) { case SQL_FETCH_NEXT: result_set.cursor++; - if (result_set.cursor >= result_set.rows.size() && - handle.GetPagingInfo().page_token.empty()) { - LOG(INFO) << "SQLFetch:: cursor: " << result_set.cursor - << " is >= result set size: " << result_set.rows.size(); - return SQL_NO_DATA; + if (result_set.cursor >= static_cast(result_set.rows.size())) { + StatusRecord next_page_status = FetchNextResultSet(handle); + if (!next_page_status.ok()) { + if (next_page_status.sql_state != SQLStates::k_SQL_NO_DATA()) { + LOG(ERROR) << "SQLFetchScroll::FetchNextResultSet:: " + << next_page_status.message; + } + return LogAndReturnCode(handle, next_page_status); + } + result_set = handle.GetResultSet(); + result_set.cursor++; } break; case SQL_FETCH_PRIOR: @@ -302,7 +303,7 @@ SQLRETURN SQLFetchScrollInternal(SQLHSTMT statement_handle, rowset_size = 1; } DescriptorHandle& ird = handle.GetDescriptorHandle(DescriptorType::kIRD); - status_record = WriteRowset(result_set, rowset_size, ard, ird); + status_record = WriteRowset(handle, rowset_size, ard, ird); return LogAndReturnCode(handle, status_record); } diff --git a/google/cloud/odbc/integration_tests/odbc_driver_tests/connection_test.cc b/google/cloud/odbc/integration_tests/odbc_driver_tests/connection_test.cc index 2798bef220..9ed97ee748 100644 --- a/google/cloud/odbc/integration_tests/odbc_driver_tests/connection_test.cc +++ b/google/cloud/odbc/integration_tests/odbc_driver_tests/connection_test.cc @@ -1874,7 +1874,8 @@ TEST(ConnectionTest, SQLDriverConnectW_Utf16EncodingOverride) { for (char c : conn_str) { utf16_conn.push_back(static_cast(c)); } - utf16_conn.push_back(0); // NUL terminator + utf16_conn.insert(utf16_conn.end(), 8, + 0); // Ensure safe NUL termination for 4-byte wchar auto run_connect_attempt = [&](char const* ini_path) -> bool { if (ini_path && ini_path[0] != '\0') { @@ -1982,6 +1983,9 @@ TEST(ConnectionTest, SQLDriverConnectW_Utf8EncodingOverride) { // Construct a UTF-8 connection string buffer (1 byte per char). std::string conn_str = kDefaultConnectionString; + std::vector utf8_conn(conn_str.begin(), conn_str.end()); + utf8_conn.insert(utf8_conn.end(), 16, + '\0'); // Ensure safe NUL termination for any wchar size auto run_connect_attempt = [&](char const* ini_path) -> bool { if (ini_path && ini_path[0] != '\0') { @@ -2000,8 +2004,7 @@ TEST(ConnectionTest, SQLDriverConnectW_Utf8EncodingOverride) { if (sql_set_env_attr(henv, SQL_ATTR_ODBC_VERSION, (SQLPOINTER)SQL_OV_ODBC3, 0) == SQL_SUCCESS) { if (sql_alloc_handle(SQL_HANDLE_DBC, henv, &hdbc) == SQL_SUCCESS) { - SQLWCHAR* in_str = - reinterpret_cast(const_cast(conn_str.data())); + SQLWCHAR* in_str = reinterpret_cast(utf8_conn.data()); SQLRETURN rc = sql_driver_connect_w(hdbc, nullptr, in_str, SQL_NTS, nullptr, 0, nullptr, SQL_DRIVER_COMPLETE); diff --git a/google/cloud/odbc/integration_tests/odbc_driver_tests/statement_test.cc b/google/cloud/odbc/integration_tests/odbc_driver_tests/statement_test.cc index f4f93443dd..f0b497b601 100644 --- a/google/cloud/odbc/integration_tests/odbc_driver_tests/statement_test.cc +++ b/google/cloud/odbc/integration_tests/odbc_driver_tests/statement_test.cc @@ -771,6 +771,104 @@ TEST_P(HTAPIParameterizedTest, SQLExecDirect_with_pagination) { EXPECT_EQ(Disconnect(conn), SQL_SUCCESS); } +TEST_P(HTAPIParameterizedTest, SQLExecDirect_pagination_with_row_array_size) { + bool is_htapi = GetParam(); + auto conn = std::make_shared(); + std::string connection_string = kDefaultConnectionString; + int limit = 3000; + if (is_htapi) { + connection_string = + kDefaultConnectionString + + ";AllowHtapiForLargeResults=1;UseDefaultLargeResultsDataset=0;" + "LargeResultsDataSetId=_bqodbc_temp_tables"; + } + EXPECT_EQ(Connect(connection_string, conn), SQL_SUCCESS); + + // Configure Block Cursor / Array Fetching (as SAP HANA SDA does) + const SQLULEN kRowArraySize = 5000; + SQLULEN rows_fetched = 0; + std::vector row_status(kRowArraySize, 0); + + SQLRETURN status = SQLSetStmtAttr(conn->hstmt, SQL_ATTR_ROW_ARRAY_SIZE, + (SQLPOINTER)kRowArraySize, 0); + ASSERT_EQ(status, SQL_SUCCESS); + + status = + SQLSetStmtAttr(conn->hstmt, SQL_ATTR_ROWS_FETCHED_PTR, &rows_fetched, 0); + ASSERT_EQ(status, SQL_SUCCESS); + + status = SQLSetStmtAttr(conn->hstmt, SQL_ATTR_ROW_STATUS_PTR, + row_status.data(), 0); + ASSERT_EQ(status, SQL_SUCCESS); + + std::string query = + "SELECT * EXCEPT (index) FROM " + "ODBC_HTAPI_TESTING.300_columns_string " + "ORDER BY index LIMIT " + + std::to_string(limit) + ";"; + + status = SQLExecDirect(conn->hstmt, (SQLCHAR*)query.c_str(), SQL_NTS); + CheckError(status, "SQLExecDirect", conn); + + SQLSMALLINT num_cols; + status = SQLNumResultCols(conn->hstmt, &num_cols); + ASSERT_EQ(status, SQL_SUCCESS); + ASSERT_EQ(num_cols, 300); + + // Column-wise array buffers for binding + size_t const kColBufferLen = 64; + std::vector> col_buffers( + num_cols, std::vector(kRowArraySize * kColBufferLen)); + std::vector> col_ind_buffers( + num_cols, std::vector(kRowArraySize)); + + for (int i = 0; i < num_cols; ++i) { + status = SQLBindCol(conn->hstmt, static_cast(i + 1), + SQL_C_CHAR, col_buffers[i].data(), kColBufferLen, + col_ind_buffers[i].data()); + CheckError(status, "SQLBindCol", conn); + } + + // Fetch all batches + SQLULEN total_rows_fetched = 0; + int fetch_call_count = 0; + std::vector batch_sizes; + while (true) { + status = SQLFetch(conn->hstmt); + if (status == SQL_NO_DATA) { + break; + } + fetch_call_count++; + batch_sizes.push_back(rows_fetched); + ASSERT_TRUE(status == SQL_SUCCESS || status == SQL_SUCCESS_WITH_INFO); + EXPECT_GT(rows_fetched, 0); + total_rows_fetched += rows_fetched; + } + + for (size_t b = 0; b < batch_sizes.size(); ++b) { + std::cout << "Fetch call " << b + 1 << " returned " << batch_sizes[b] + << " rows" << std::endl; + } + + // Assert all rows across all pages were fetched + EXPECT_EQ(total_rows_fetched, limit) + << "Row count mismatch when using block fetch across multiple pages."; + + // Validation to catch if a driver is not populating the full requested array + // size across pages: + ASSERT_FALSE(batch_sizes.empty()); + EXPECT_EQ(batch_sizes[0], static_cast(limit)) + << "Driver failed to populate the full requested size (" << limit + << " rows) across page boundaries in the first fetch call (got " + << batch_sizes[0] << " instead)."; + EXPECT_EQ(batch_sizes.size(), 1U) + << "Driver took " << batch_sizes.size() + << " fetch calls instead of populating all " << limit + << " rows in 1 call when ROW_ARRAY_SIZE is " << kRowArraySize; + + EXPECT_EQ(Disconnect(conn), SQL_SUCCESS); +} + TEST_P(HTAPIParameterizedTest, SQLExecDirect_with_empty_result_set) { bool is_htapi = GetParam(); auto conn = std::make_shared(); @@ -1683,6 +1781,11 @@ TEST_P(StatementParameterizedTest, FreeExplicitDescriptor) { status = SQLSetStmtAttr(conn->hstmt, SQL_ATTR_APP_PARAM_DESC, conn->apd, 0); CheckError(status, "SQLSetStmtAttr", conn); + // Disassociate explicit descriptor by reverting statement to SQL_NULL_HDESC + status = + SQLSetStmtAttr(conn->hstmt, SQL_ATTR_APP_PARAM_DESC, SQL_NULL_HDESC, 0); + CheckError(status, "SQLSetStmtAttr(SQL_NULL_HDESC)", conn); + // Free explicit descriptor EXPECT_EQ(SQLFreeHandle(SQL_HANDLE_DESC, conn->apd), SQL_SUCCESS); @@ -4471,4 +4574,47 @@ INSTANTIATE_TEST_SUITE_P( SQL_ROLLBACK, static_cast(1)), std::make_tuple("1", "ODBC_IGNORE_TRANSACTIONS_ON_COMMIT", SQL_COMMIT, static_cast(1)))); + +TEST(StatementTest, OdbcEscapeTimestampLiteral) { + auto conn = std::make_shared(); + ASSERT_EQ(Connect(kDefaultConnectionString, conn), SQL_SUCCESS); + + // Test 1: SQLExecDirect with ODBC timestamp escape clause {ts '...'} in WHERE + // filter (as in SAP HANA SDA) + std::string const query = + "SELECT 1 AS res FROM (SELECT TIMESTAMP '2014-02-20 09:00:00' AS " + "created_date) WHERE created_date < {ts '2014-02-20 09:34:06.000'}"; + SQLRETURN status = + SQLExecDirect(conn->hstmt, (SQLCHAR*)query.c_str(), SQL_NTS); + CheckError(status, "SQLExecDirect({ts '...'})", conn); + + ASSERT_EQ(SQLFetch(conn->hstmt), SQL_SUCCESS); + SQLBIGINT res = 0; + SQLLEN ind = 0; + ASSERT_EQ(SQLGetData(conn->hstmt, 1, SQL_C_SBIGINT, &res, sizeof(res), &ind), + SQL_SUCCESS); + EXPECT_EQ(res, 1); + + SQLFreeStmt(conn->hstmt, SQL_CLOSE); + + // Test 2: SQLPrepare & SQLExecute with ODBC timestamp escape clause in + // projection + std::string const prepare_query = + "SELECT {ts '2014-02-20 09:34:06.000'} AS ts_col"; + status = SQLPrepare(conn->hstmt, (SQLCHAR*)prepare_query.c_str(), SQL_NTS); + CheckError(status, "SQLPrepare({ts '...'})", conn); + status = SQLExecute(conn->hstmt); + CheckError(status, "SQLExecute({ts '...'})", conn); + + ASSERT_EQ(SQLFetch(conn->hstmt), SQL_SUCCESS); + char ts_buf[64] = {0}; + ASSERT_EQ( + SQLGetData(conn->hstmt, 1, SQL_C_CHAR, ts_buf, sizeof(ts_buf), &ind), + SQL_SUCCESS); + EXPECT_THAT(std::string(ts_buf), HasSubstr("2014-02-20 09:34:06")); + + SQLFreeStmt(conn->hstmt, SQL_CLOSE); + EXPECT_EQ(Disconnect(conn), SQL_SUCCESS); +} + } // namespace google::cloud::odbc_tests diff --git a/google/cloud/odbc/testing/odbc_utils/commons.h b/google/cloud/odbc/testing/odbc_utils/commons.h index e5946a7b87..a78e965870 100644 --- a/google/cloud/odbc/testing/odbc_utils/commons.h +++ b/google/cloud/odbc/testing/odbc_utils/commons.h @@ -141,7 +141,7 @@ struct ODBCHandles { SQLHDESC apd; // Application parameter descriptor SQLHDESC ipd; // Implementation parameter descriptor bool connected; - SQLCHAR outdsn[4096]; + alignas(SQLWCHAR) SQLCHAR outdsn[4096]; Metadata metadata; }; diff --git a/google/cloud/odbc/testing/odbc_utils/connection.cc b/google/cloud/odbc/testing/odbc_utils/connection.cc index da32dda607..45bc57da80 100644 --- a/google/cloud/odbc/testing/odbc_utils/connection.cc +++ b/google/cloud/odbc/testing/odbc_utils/connection.cc @@ -220,16 +220,15 @@ SQLRETURN Connect(std::wstring dsn, std::shared_ptr const& conn, sql_w_str.emplace_back(L'\0'); if (is_driver_connect) { - status = - SQLDriverConnectW(conn->hdbc, nullptr, sql_w_str.data(), SQL_NTS, - reinterpret_cast(conn->outdsn), - sizeof(conn->outdsn), &buflen, SQL_DRIVER_COMPLETE); + status = SQLDriverConnectW(conn->hdbc, nullptr, sql_w_str.data(), SQL_NTS, + reinterpret_cast(conn->outdsn), + sizeof(conn->outdsn) / sizeof(SQLWCHAR), &buflen, + SQL_DRIVER_COMPLETE); CheckError(status, "SQLDriverConnectW", conn); } else { - status = SQLConnectW(conn->hdbc, sql_w_str.data(), SQL_NTS, - reinterpret_cast(conn->outdsn), - NumSqlChar(conn->outdsn), nullptr, 0); + status = SQLConnectW(conn->hdbc, sql_w_str.data(), SQL_NTS, nullptr, 0, + nullptr, 0); CheckError(status, "SQLConnectW", conn); }