diff --git a/ci/gha/builds/lib/google.googlebigqueryodbc.ini b/ci/gha/builds/lib/google.googlebigqueryodbc.ini index cd13ac1435..e49210b3c6 100644 --- a/ci/gha/builds/lib/google.googlebigqueryodbc.ini +++ b/ci/gha/builds/lib/google.googlebigqueryodbc.ini @@ -1,3 +1,12 @@ [Driver] LogLevel=0 LogPath= +# WcharEncoding sets the wire encoding of SQLWCHAR buffers on Linux/macOS +# when the driver is built against iODBC headers (sizeof(SQLWCHAR) == 4). +# +# Accepted values: +# UTF-16LE - 2-byte UTF-16LE per code unit (unixODBC-loaded driver) +# UTF-32LE - 4-byte UTF-32LE per code unit (native iODBC) +# (empty) - use sizeof(SQLWCHAR) as-is (default) +# +WcharEncoding= diff --git a/google/cloud/odbc/CMakeLists.txt b/google/cloud/odbc/CMakeLists.txt index 57882661fd..ec04725342 100644 --- a/google/cloud/odbc/CMakeLists.txt +++ b/google/cloud/odbc/CMakeLists.txt @@ -159,6 +159,15 @@ else () find_package(google_cloud_cpp_serviceusage REQUIRED) endif () +find_package(re2 CONFIG QUIET) +if (NOT TARGET re2::re2) + if (TARGET re2) + add_library(re2::re2 ALIAS re2) + else () + find_package(re2 REQUIRED) + endif () +endif () + # Restore the original BUILD_SHARED_LIBS value set(BUILD_SHARED_LIBS ${ORIGINAL_BUILD_SHARED_LIBS}) diff --git a/google/cloud/odbc/bq_driver/internal/data_translation.cc b/google/cloud/odbc/bq_driver/internal/data_translation.cc index eace909d6a..e0b0ef8568 100644 --- a/google/cloud/odbc/bq_driver/internal/data_translation.cc +++ b/google/cloud/odbc/bq_driver/internal/data_translation.cc @@ -104,7 +104,7 @@ odbc_internal::StatusRecord ConvertFromNumericDSValue(DSValue const& src_dsval, "DSValueToWchar Conversion Failed"}; break; } - SQLLEN wchar_capacity = dest_data.buflen / sizeof(SQLWCHAR); + SQLLEN wchar_capacity = dest_data.buflen / WireWcharSize(); auto src_len = static_cast(wstr->length()); SQLINTEGER required_chars = src_len + 1; WStrToOutputBufferResponse(wstr.GetValue(), dest_data.buf, wchar_capacity, @@ -326,7 +326,7 @@ odbc_internal::StatusRecord ConvertFromStringDSValue(DSValue const& src_dsval, } auto src_len = static_cast(wide_str.length()); - SQLLEN wchar_capacity = dest_data.buflen / sizeof(SQLWCHAR); + SQLLEN wchar_capacity = dest_data.buflen / WireWcharSize(); SQLINTEGER required_chars = src_len + 1; return WStrToOutputBufferResponse(wide_str, dest_data.buf, wchar_capacity, src_len, required_chars, @@ -907,7 +907,7 @@ odbc_internal::StatusRecord ConvertFromTimeDSValue(DSValue const& src_dsval, "DSValueToWchar Conversion Failed"}; break; } - SQLLEN wchar_capacity = buffer_length / sizeof(SQLWCHAR); + SQLLEN wchar_capacity = buffer_length / WireWcharSize(); SQLLEN required_chars = static_cast(wstr->length()) + 1; return WStrToOutputBufferResponse( wstr.GetValue(), dest_buf, wchar_capacity, k_time_src_len, @@ -1007,30 +1007,23 @@ odbc_internal::StatusRecord ConvertFromTimestampDSValue( "DSValueToWchar Conversion Failed"}; break; } - std::wstring wstr_val = wstr.GetValue(); - if (!wstr_val.empty() && wstr_val.back() == L'\0') { - wstr_val.pop_back(); - } - std::vector wstr_data(wstr_val.begin(), wstr_val.end()); - wstr_data.emplace_back(L'\0'); - - auto* dest = reinterpret_cast(dest_buf); - SQLLEN wchar_capacity = buffer_length / sizeof(SQLWCHAR); + size_t const wire_sz = WireWcharSize(); + SQLLEN wchar_capacity = buffer_length / static_cast(wire_sz); if (wchar_capacity > k_timestamp_src_len) { if (res_len) { - *res_len = k_timestamp_src_len * sizeof(SQLWCHAR); + *res_len = static_cast(wstr.GetValue().size() * wire_sz); } - std::memcpy(dest, wstr_data.data(), - (k_timestamp_src_len) * sizeof(SQLWCHAR)); - dest[k_timestamp_src_len] = L'\0'; + WriteWideToWireBuffer(wstr.GetValue(), dest_buf, + wstr.GetValue().size()); + WriteWireNul(dest_buf, wstr.GetValue().size()); } else if (20 <= wchar_capacity && wchar_capacity <= k_timestamp_src_len) { if (res_len) { - *res_len = wchar_capacity * sizeof(SQLWCHAR); + *res_len = wchar_capacity * static_cast(wire_sz); } - std::memcpy(dest, wstr_data.data(), - (wchar_capacity) * sizeof(SQLWCHAR)); - dest[wchar_capacity - 1] = L'\0'; + WriteWideToWireBuffer(wstr.GetValue(), dest_buf, + static_cast(wchar_capacity - 1)); + WriteWireNul(dest_buf, static_cast(wchar_capacity - 1)); LOG(WARNING) << "ConvertFromTimestampDSValue:: Data truncated for SQL_C_WCHAR."; status_record = StatusRecord{SQLStates::k_01004(), "Data truncated"}; @@ -1189,29 +1182,22 @@ odbc_internal::StatusRecord ConvertFromDatetimeDSValue(DSValue const& src_dsval, "DSValueToWchar Conversion Failed"}; break; } - std::wstring wstr_val = wstr.GetValue(); - if (!wstr_val.empty() && wstr_val.back() == L'\0') { - wstr_val.pop_back(); - } - std::vector wstr_data(wstr_val.begin(), wstr_val.end()); - wstr_data.emplace_back(L'\0'); - - auto* dest = reinterpret_cast(dest_buf); - SQLLEN wchar_capacity = buffer_length / sizeof(SQLWCHAR); + size_t const wire_sz = WireWcharSize(); + SQLLEN wchar_capacity = buffer_length / static_cast(wire_sz); if (wchar_capacity > k_datetime_src_len) { if (res_len) { - *res_len = k_datetime_src_len * sizeof(SQLWCHAR); + *res_len = static_cast(wstr.GetValue().size() * wire_sz); } - std::memcpy(dest, wstr_data.data(), - (k_datetime_src_len) * sizeof(SQLWCHAR)); - dest[k_datetime_src_len] = L'\0'; + WriteWideToWireBuffer(wstr.GetValue(), dest_buf, + wstr.GetValue().size()); + WriteWireNul(dest_buf, wstr.GetValue().size()); } else if (20 <= wchar_capacity && wchar_capacity <= k_datetime_src_len) { if (res_len) { - *res_len = wchar_capacity * sizeof(SQLWCHAR); + *res_len = wchar_capacity * static_cast(wire_sz); } - std::memcpy(dest, wstr_data.data(), - (wchar_capacity) * sizeof(SQLWCHAR)); - dest[wchar_capacity - 1] = L'\0'; + WriteWideToWireBuffer(wstr.GetValue(), dest_buf, + static_cast(wchar_capacity - 1)); + WriteWireNul(dest_buf, static_cast(wchar_capacity - 1)); LOG(WARNING) << "ConvertFromDatetimeDSValue:: Data truncated for SQL_C_WCHAR."; status_record = StatusRecord{SQLStates::k_01004(), "Data truncated"}; @@ -1398,7 +1384,7 @@ odbc_internal::StatusRecord ConvertFromDateDSValue(DSValue const& src_dsval, return StatusRecord{SQLStates::k_HY000(), "DSValueToWchar Conversion Failed"}; } - SQLLEN wchar_capacity = buffer_length / sizeof(SQLWCHAR); + SQLLEN wchar_capacity = buffer_length / WireWcharSize(); auto src_len = static_cast(wstr->length()); SQLINTEGER required_chars = src_len + 1; return WStrToOutputBufferResponse( @@ -1434,7 +1420,7 @@ StatusRecord ConvertStringToJsonOutputBuffer(std::string const& src_str, return StatusRecord{SQLStates::k_HY000(), "Conversion to UTF-16 failed"}; } - SQLLEN wchar_capacity = buffer_length / sizeof(SQLWCHAR); + SQLLEN wchar_capacity = buffer_length / WireWcharSize(); auto src_len = static_cast(wide_string->length()); SQLINTEGER required_chars = src_len + 1; return WStrToOutputBufferResponse(wide_string.GetValue(), dest_buf, @@ -1498,15 +1484,11 @@ StatusRecord ConvertFromArrayDSValue(DSValue const& src_dsval, if (!wide_string.Ok()) { return StatusRecord{SQLStates::k_HY000(), "Conversion Failed"}; } - std::wstring wide_val = wide_string.GetValue(); - if (!wide_val.empty() && wide_val.back() == L'\0') { - wide_val.pop_back(); - } - SQLLEN wchar_capacity = dest_data.buflen / sizeof(SQLWCHAR); - auto src_len = static_cast(wide_val.length()); + SQLLEN wchar_capacity = dest_data.buflen / WireWcharSize(); + auto src_len = static_cast(wide_string->length()); SQLINTEGER required_chars = src_len + 1; return WStrToOutputBufferResponse( - wide_val, dest_data.buf, wchar_capacity, src_len, required_chars, + *wide_string, dest_data.buf, wchar_capacity, src_len, required_chars, reinterpret_cast(dest_data.result_len)); } case SQL_C_BINARY: { @@ -1621,7 +1603,7 @@ odbc_internal::StatusRecord ConvertFromIntervalDSValue(DSValue const& src_dsval, StatusRecord{SQLStates::k_HY000(), wstr.GetStatusRecord().message}; break; } - SQLLEN wchar_capacity = buffer_length / sizeof(SQLWCHAR); + SQLLEN wchar_capacity = buffer_length / WireWcharSize(); auto interval_char_length = static_cast(wstr.GetValue().length()); return WStrIntervalBufferResponse( @@ -1907,7 +1889,7 @@ StatusRecord ConvertFromGeographyDSValue(DSValue const& src_dsval, } std::memset(dest_data.buf, 0, buffer_length); std::wstring const& wide_str = wstr.GetValue(); - SQLLEN wchar_capacity = buffer_length / sizeof(SQLWCHAR); + SQLLEN wchar_capacity = buffer_length / WireWcharSize(); SQLLEN src_len = static_cast(wide_str.length()); SQLLEN required_chars = src_len + 1; status_record = WStrToOutputBufferResponse( @@ -2054,38 +2036,35 @@ StatusRecord ConvertBytesToWChar(DSValue const& conn_val, "UTF-8 to UTF-16 conversion failed."}; } - std::wstring utf16_value = utf16_str.GetValue(); - if (!utf16_value.empty() && utf16_value.back() == L'\0') { - utf16_value.pop_back(); - } - size_t const required_size = utf16_value.length() * sizeof(SQLWCHAR); - - auto* buffer = reinterpret_cast(dest_data.buf); + std::wstring const& utf16_value = utf16_str.GetValue(); - // Handle truncation if buffer is insufficient - if (dest_data.buflen < required_size) { - size_t num_chars_to_copy = (dest_data.buflen / sizeof(SQLWCHAR)) - 1; - std::memcpy(buffer, utf16_value.data(), - num_chars_to_copy * sizeof(SQLWCHAR)); - buffer[num_chars_to_copy] = L'\0'; + // Narrow wchar_t -> wire encoding directly into the caller's buffer. + // No intermediate vector; WriteWideToWireBuffer is a memcpy when the wire + // SQLWCHAR width matches sizeof(wchar_t) and a per-element narrowing loop + // only on the iODBC-built / unixODBC-loaded path. + size_t const wire_sz = WireWcharSize(); + size_t const src_chars = utf16_value.size(); + size_t const required_size = src_chars * wire_sz; + if (static_cast(dest_data.buflen) < required_size) { + size_t num_chars_to_copy = dest_data.buflen / wire_sz; + if (num_chars_to_copy > 0) { + num_chars_to_copy--; // leave one slot for the null terminator + WriteWideToWireBuffer(utf16_value, dest_data.buf, num_chars_to_copy); + WriteWireNul(dest_data.buf, num_chars_to_copy); + } if (dest_data.result_len) { - *dest_data.result_len = dest_data.buflen; + *dest_data.result_len = required_size; } LOG(WARNING) << "ConvertBytesToWChar:: String data, right truncated."; return StatusRecord{SQLStates::k_01004(), "String data, right truncated"}; } - for (size_t i = 0; i < utf16_str.GetValue().size(); ++i) { - buffer[i] = static_cast(utf16_str.GetValue()[i]); - } - size_t buffer_chars = dest_data.buflen / sizeof(SQLWCHAR); - if (utf16_str.GetValue().size() < buffer_chars) { - buffer[utf16_str.GetValue().size()] = L'\0'; + WriteWideToWireBuffer(utf16_value, dest_data.buf, src_chars); + if (static_cast(dest_data.buflen) >= required_size + wire_sz) { + WriteWireNul(dest_data.buf, src_chars); } - - // Set output length if (dest_data.result_len) { - *dest_data.result_len = utf16_str.GetValue().size() * sizeof(SQLWCHAR); + *dest_data.result_len = required_size; } return status_record; } @@ -2253,7 +2232,7 @@ StatusRecord ConvertFromRangeDSValue(DSValue const& src_dsval, return StatusRecord{SQLStates::k_HY000(), "Conversion to SQL_C_WCHAR failed."}; } - SQLLEN wchar_capacity = buffer_length / sizeof(SQLWCHAR); + SQLLEN wchar_capacity = buffer_length / WireWcharSize(); SQLLEN src_len = static_cast(wstr->length()); SQLLEN required_chars = src_len + 1; return WStrToOutputBufferResponse( diff --git a/google/cloud/odbc/bq_driver/internal/data_translation_inv.cc b/google/cloud/odbc/bq_driver/internal/data_translation_inv.cc index 76832234b3..c20ae95e00 100644 --- a/google/cloud/odbc/bq_driver/internal/data_translation_inv.cc +++ b/google/cloud/odbc/bq_driver/internal/data_translation_inv.cc @@ -53,7 +53,7 @@ StatusRecordOr ConvertFromCharBuffer(DataBuffer& src_data, auto* wchar_buf = static_cast(src_buf); if ((result_len > 0) || (result_len == SQL_NTS)) { if (result_len > 0) { - result_len /= sizeof(SQLWCHAR); + result_len /= WireWcharSize(); } auto utf8_res = BqConvertSQLWCHARToString( wchar_buf, static_cast(result_len)); diff --git a/google/cloud/odbc/bq_driver/internal/odbc_desc_attr.cc b/google/cloud/odbc/bq_driver/internal/odbc_desc_attr.cc index 8052921952..1a28dcb25d 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_desc_attr.cc +++ b/google/cloud/odbc/bq_driver/internal/odbc_desc_attr.cc @@ -14,6 +14,7 @@ #include "google/cloud/odbc/bq_driver/internal/odbc_desc_attr.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/sql_state_constants.h" #include "google/cloud/odbc/internal/status_record_or.h" #include @@ -388,7 +389,7 @@ StatusRecord DescriptorRecord::SetOctetLength(SQLSMALLINT type, case SQL_WCHAR: case SQL_WVARCHAR: case SQL_WLONGVARCHAR: - octet_length = value * sizeof(SQLWCHAR); + octet_length = value * WireWcharSize(); break; case SQL_DECIMAL: case SQL_NUMERIC: diff --git a/google/cloud/odbc/bq_driver/internal/odbc_sql_columns.cc b/google/cloud/odbc/bq_driver/internal/odbc_sql_columns.cc index fc445675b5..18a4c39254 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_sql_columns.cc +++ b/google/cloud/odbc/bq_driver/internal/odbc_sql_columns.cc @@ -356,7 +356,8 @@ StatusRecordOr ProcessTableResults( for (TableFieldSchema const& table_field_schema : bq_table.schema.fields) { // bq_table_column could contain a search pattern character so do a regex // match. - auto column_pattern = BuildRegex(bq_table_column, metadata_id); + std::unique_ptr column_pattern = + BuildRegex(bq_table_column, metadata_id); if (re2::RE2::FullMatch(table_field_schema.name, *column_pattern)) { auto ds_row_status = CreateResultSetDSRow( conn_handle, bq_table.table_reference.project_id, diff --git a/google/cloud/odbc/bq_driver/internal/odbc_sql_columns_test.cc b/google/cloud/odbc/bq_driver/internal/odbc_sql_columns_test.cc index 70974bf7f2..09b572964f 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_sql_columns_test.cc +++ b/google/cloud/odbc/bq_driver/internal/odbc_sql_columns_test.cc @@ -299,7 +299,7 @@ void ProcessTableResultsHelper(std::string const& column, expected_sql_int_row.ord_pos = (column == "%" || column.empty()) ? 2 : 1; expected_sql_int_row.is_nullable = "NO"; - auto column_pattern = BuildRegex(column, metadata_id); + std::unique_ptr column_pattern = BuildRegex(column, metadata_id); if (!metadata_id && (column.empty() || column == "%")) { ASSERT_EQ(result_set.rows.size(), 2); diff --git a/google/cloud/odbc/bq_driver/internal/odbc_sql_tables.cc b/google/cloud/odbc/bq_driver/internal/odbc_sql_tables.cc index 88e7f6db84..adb278ae54 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_sql_tables.cc +++ b/google/cloud/odbc/bq_driver/internal/odbc_sql_tables.cc @@ -83,7 +83,8 @@ StatusRecordOr> GetFilteredProjectIds( ODBCBQClient& bq_client, std::string const& projects_filter, SQLULEN metadata_id) { std::vector project_ids; - auto filter_regex = BuildRegex(projects_filter, metadata_id); + std::unique_ptr filter_regex = + BuildRegex(projects_filter, metadata_id); // For now, we use default options. // We can set timeout here as needed later. Options options; diff --git a/google/cloud/odbc/bq_driver/internal/odbc_type_utils.cc b/google/cloud/odbc/bq_driver/internal/odbc_type_utils.cc index 993fa7374d..365c722aa4 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_type_utils.cc +++ b/google/cloud/odbc/bq_driver/internal/odbc_type_utils.cc @@ -42,25 +42,24 @@ SQLRETURN AddressToPointer(SQLPOINTER ptr, SQLPOINTER out_buf, } odbc_internal::StatusRecord WStrIntervalBufferResponse( - std::wstring wstr, SQLPOINTER dest_buf, SQLLEN buffer_length, + std::wstring const& wstr, SQLPOINTER dest_buf, SQLLEN buffer_length, SQLINTEGER char_len, SQLINTEGER whole_digits_count, SQLLEN* res_len) { auto status_record = odbc_internal::StatusRecord::Ok(); - std::vector wstr_data(wstr.begin(), wstr.end()); - wstr_data.emplace_back(L'\0'); + size_t const wire_sz = WireWcharSize(); - auto* dest = static_cast(dest_buf); if (buffer_length > char_len) { if (res_len) { - *res_len = char_len * sizeof(SQLWCHAR); + *res_len = static_cast(char_len) * static_cast(wire_sz); } - std::memcpy(dest, wstr_data.data(), (char_len) * sizeof(SQLWCHAR)); - dest[char_len] = L'\0'; + WriteWideToWireBuffer(wstr, dest_buf, static_cast(char_len)); + WriteWireNul(dest_buf, static_cast(char_len)); } else if (buffer_length > whole_digits_count) { if (res_len) { - *res_len = buffer_length * sizeof(SQLWCHAR); + *res_len = buffer_length * static_cast(wire_sz); } - std::memcpy(dest, wstr_data.data(), (buffer_length) * sizeof(SQLWCHAR)); - dest[buffer_length - 1] = L'\0'; + WriteWideToWireBuffer(wstr, dest_buf, + static_cast(buffer_length - 1)); + WriteWireNul(dest_buf, static_cast(buffer_length - 1)); status_record = odbc_internal::StatusRecord{ google::cloud::odbc_internal::SQLStates::k_01004(), "Data truncated"}; } else { diff --git a/google/cloud/odbc/bq_driver/internal/odbc_type_utils.h b/google/cloud/odbc/bq_driver/internal/odbc_type_utils.h index 6630ef929f..71f40e05a9 100644 --- a/google/cloud/odbc/bq_driver/internal/odbc_type_utils.h +++ b/google/cloud/odbc/bq_driver/internal/odbc_type_utils.h @@ -15,8 +15,10 @@ #ifndef CPP_BIGQUERY_ODBC_GOOGLE_CLOUD_ODBC_BQ_DRIVER_INTERNAL_ODBC_TYPE_UTILS_H #define CPP_BIGQUERY_ODBC_GOOGLE_CLOUD_ODBC_BQ_DRIVER_INTERNAL_ODBC_TYPE_UTILS_H +#include "google/cloud/odbc/bq_driver/internal/utils.h" #include "google/cloud/odbc/internal/diagnostic_records.h" #include "google/cloud/odbc/internal/sql_state_constants.h" +#include #include #include #include @@ -175,13 +177,49 @@ SQLRETURN IntValueToOutputBufferResponse(T val, SQLPOINTER buffer_ptr, return SQL_SUCCESS; } +// Writes `count` wide characters from `src` directly into `dest` using the +// current wire encoding. `dest` must point to caller-owned storage of at +// least `count * WireWcharSize()` bytes. The caller writes its own NUL +// terminator (`WriteWireNul` below) if it wants one. +// +// When the wire SQLWCHAR width matches `sizeof(wchar_t)` (Windows and the +// iODBC build — the common case), this is a single memcpy of the wstring's +// raw bytes +inline void WriteWideToWireBuffer(std::wstring const& src, void* dest, + size_t count) { + if (count > src.size()) count = src.size(); + +#if !defined(_WIN32) + if (IsRuntimeWireUtf16Le() || sizeof(SQLWCHAR) != sizeof(wchar_t)) { + auto* d = static_cast(dest); + for (size_t i = 0; i < count; ++i) { + // Cast to unsigned 32-bit first to avoid signed→unsigned misuse warning, + // then narrow to uint16_t (valid for BMP code points / UTF-16 units). + d[i] = static_cast(static_cast(src[i])); + } + return; + } +#endif + + std::memcpy(dest, src.data(), count * sizeof(SQLWCHAR)); +} + +// Writes a single wire-format NUL terminator (one code unit, 2 or 4 bytes) +// at byte offset `char_index * WireWcharSize()` from `dest`. +inline void WriteWireNul(void* dest, size_t char_index) { + size_t const wire_sz = WireWcharSize(); + std::memset(static_cast(dest) + (char_index * wire_sz), 0, wire_sz); +} + inline odbc_internal::StatusRecord WStrToOutputBufferResponse( - std::wstring wstr, SQLPOINTER dest_buf, SQLLEN buffer_length, + std::wstring const& wstr, SQLPOINTER dest_buf, SQLLEN buffer_length, SQLINTEGER src_len, SQLINTEGER supp_max_len, SQLLEN* res_len) { auto status_record = odbc_internal::StatusRecord::Ok(); + size_t const wire_sz = WireWcharSize(); + if (wstr.empty()) { if (dest_buf && buffer_length > 0) { - reinterpret_cast(dest_buf)[0] = L'\0'; + WriteWireNul(dest_buf, 0); } if (res_len) { *res_len = 0; @@ -189,21 +227,18 @@ inline odbc_internal::StatusRecord WStrToOutputBufferResponse( return status_record; } - std::vector wstr_data(wstr.begin(), wstr.end()); - - auto* dest = reinterpret_cast(dest_buf); if (buffer_length > src_len) { if (res_len) { - *res_len = src_len * sizeof(SQLWCHAR); + *res_len = src_len * static_cast(wire_sz); } - std::memcpy(dest, wstr_data.data(), (src_len) * sizeof(SQLWCHAR)); - dest[src_len] = L'\0'; + WriteWideToWireBuffer(wstr, dest_buf, src_len); + WriteWireNul(dest_buf, src_len); } else if (supp_max_len <= buffer_length && buffer_length <= src_len) { if (res_len) { - *res_len = buffer_length * sizeof(SQLWCHAR); + *res_len = buffer_length * static_cast(wire_sz); } - std::memcpy(dest, wstr_data.data(), (buffer_length) * sizeof(SQLWCHAR)); - dest[buffer_length - 1] = L'\0'; + WriteWideToWireBuffer(wstr, dest_buf, buffer_length - 1); + WriteWireNul(dest_buf, buffer_length - 1); status_record = odbc_internal::StatusRecord{ google::cloud::odbc_internal::SQLStates::k_01004(), "Data truncated"}; } else { @@ -221,7 +256,7 @@ SQLRETURN AddressToPointer(SQLPOINTER ptr, SQLPOINTER out_buf, SQLSMALLINT* str_len_ptr); odbc_internal::StatusRecord WStrIntervalBufferResponse( - std::wstring wstr, SQLPOINTER dest_buf, SQLLEN buffer_length, + std::wstring const& wstr, SQLPOINTER dest_buf, SQLLEN buffer_length, SQLINTEGER char_len, SQLINTEGER whole_digits_count, SQLLEN* res_len); } // namespace google::cloud::odbc_bq_driver_internal diff --git a/google/cloud/odbc/bq_driver/internal/trace_utils.cc b/google/cloud/odbc/bq_driver/internal/trace_utils.cc index 9d9ee8aded..c748c10f21 100644 --- a/google/cloud/odbc/bq_driver/internal/trace_utils.cc +++ b/google/cloud/odbc/bq_driver/internal/trace_utils.cc @@ -295,6 +295,8 @@ TraceOptions::CreateTraceOptionsFile( log_file_size = std::strtol(s.second.c_str(), nullptr, 10); } else if (s.first == kMaxThreadsParam) { max_threads = std::stoull(s.second); + } else if (s.first == kWcharEncoding) { + SetWcharEncodingFromConfig(s.second); } } diff --git a/google/cloud/odbc/bq_driver/internal/trace_utils.h b/google/cloud/odbc/bq_driver/internal/trace_utils.h index 6124b12495..4de27108d8 100644 --- a/google/cloud/odbc/bq_driver/internal/trace_utils.h +++ b/google/cloud/odbc/bq_driver/internal/trace_utils.h @@ -41,6 +41,9 @@ inline std::string const kLogPath = "LogPath"; inline std::string const kLogFileCount = "LogFileCount"; inline std::string const kLogFileSize = "LogFileSize"; inline std::string const kMaxThreadsParam = "MaxThreads"; +// Key controlling the wire encoding of SQLWCHAR buffers on Linux/macOS. +// Accepted values: "UTF-16LE", "UCS-4LE", or empty (auto-detect). +inline std::string const kWcharEncoding = "WcharEncoding"; inline std::uint32_t const kDefaultMaxThreads = 8; inline std::string const kDefaultMaxFiles = "50"; inline std::string const kDefaultMaxSize = "2000"; diff --git a/google/cloud/odbc/bq_driver/internal/utils.cc b/google/cloud/odbc/bq_driver/internal/utils.cc index b9fb984502..e11b6ec5c5 100644 --- a/google/cloud/odbc/bq_driver/internal/utils.cc +++ b/google/cloud/odbc/bq_driver/internal/utils.cc @@ -21,6 +21,7 @@ #include "google/cloud/odbc/bq_driver/internal/utils.h" #include "google/cloud/internal/getenv.h" #include +#include #include #include #include @@ -39,6 +40,52 @@ bool g_suppress_dropdown = false; using ::google::cloud::odbc_internal::SQLStates; using ::google::cloud::odbc_internal::StatusRecord; using ::google::cloud::odbc_internal::StatusRecordOr; + +namespace { +// Wire-encoding override read from the [Driver] WcharEncoding key in +// google.googlebigqueryodbc.ini +// Only meaningful when sizeof(SQLWCHAR) == 4 (iODBC build on Linux/macOS). +// kDefault use sizeof(SQLWCHAR) as the wire size (no adaptation) +// kUtf16Le 2-byte UTF-16LE wire format (unixODBC loaded driver) +// kUtf32Le 4-byte UTF-32LE wire format (native iODBC) +enum class WcharEncodingOverride { kDefault, kUtf16Le, kUtf32Le }; +std::atomic g_wchar_encoding_override{ + WcharEncodingOverride::kDefault}; +} // namespace + +bool IsRuntimeWireUtf16Le() { +#if defined(_WIN32) + return false; +#else + return g_wchar_encoding_override.load(std::memory_order_relaxed) == + WcharEncodingOverride::kUtf16Le; +#endif +} + +size_t WireWcharSize() { + return IsRuntimeWireUtf16Le() ? 2U : sizeof(SQLWCHAR); +} + +void SetWcharEncodingFromConfig(std::string const& value) { +#if !defined(_WIN32) + if (value == "UTF-16LE") { + g_wchar_encoding_override.store(WcharEncodingOverride::kUtf16Le, + std::memory_order_relaxed); + LOG(INFO) << "WcharEncoding: UTF-16LE wire format (2 bytes/char)"; + } else if (value == "UTF-32LE") { + g_wchar_encoding_override.store(WcharEncodingOverride::kUtf32Le, + std::memory_order_relaxed); + LOG(INFO) << "WcharEncoding: UTF-32LE wire format (4 bytes/char)"; + } else if (value.empty()) { + g_wchar_encoding_override.store(WcharEncodingOverride::kDefault, + std::memory_order_relaxed); + LOG(INFO) << "WcharEncoding: default (sizeof(SQLWCHAR) bytes/char)"; + } else { + LOG(WARNING) << "WcharEncoding: unrecognised value '" << value << "'"; + } +#endif +} + #ifdef _WIN32 using google::cloud::odbc_bigquery_client_interface::OauthMechanism; static std::string const kOAuthMechanism = "OAuthMechanism"; @@ -166,7 +213,7 @@ size_t BufferSizeForType(SQLSMALLINT type, size_t requested) { minimum_size = sizeof(SQL_TIMESTAMP_STRUCT); break; case SQL_C_WCHAR: - minimum_size = sizeof(SQLWCHAR); + minimum_size = WireWcharSize(); break; case SQL_C_SBIGINT: minimum_size = sizeof(SQLBIGINT); @@ -777,7 +824,10 @@ odbc_internal::StatusRecordOr Utf8ToUtf16( return StatusRecord{SQLStates::k_HY000(), "Error while converting string to wstring"}; } - utf16Str.push_back(L'\0'); + // MultiByteToWideChar was called with an explicit input length, so + // utf16Length is the character count without a null terminator and + // wstring::size() already reflects the actual character count, + // matching the Linux iconv path behaviour. return utf16Str; #else iconv_t cd = iconv_open(kFromCode.c_str(), "UTF-8"); @@ -809,10 +859,11 @@ odbc_internal::StatusRecordOr Utf8ToUtf16( iconv_close(cd); - // Resize the output string to the actual converted size + // Resize the output string to the actual converted size. No trailing NUL is + // appended: wstring::size() is the character count, matching the Windows + // MultiByteToWideChar path above. Callers own their own NUL termination. utf16str.resize((outbuf - reinterpret_cast(utf16str.data())) / sizeof(wchar_t)); - utf16str.push_back(L'\0'); return utf16str; #endif } @@ -822,7 +873,31 @@ odbc_internal::StatusRecordOr BqConvertSQLWCHARToString( if (in_str == nullptr) { return StatusRecord{SQLStates::k_HY000(), "in_str string is empty/Null"}; } - if (((in_str != nullptr) && (in_str[0] == '\0'))) { + +#if !defined(_WIN32) + // When WcharEncoding=UTF-16LE is set in google.googlebigqueryodbc.ini, + // the ODBC manager delivers 2-byte UTF-16LE code units packed into the + // buffer even though sizeof(SQLWCHAR)==4. Read them as uint16_t. + if (IsRuntimeWireUtf16Le()) { + auto const* utf16 = reinterpret_cast(in_str); + if (utf16[0] == 0) { + return std::string(); + } + SQLINTEGER count = in_str_len; + if (count == SQL_NTS || count == NULL) { + count = 0; + while (utf16[count] != 0) ++count; + } + std::wstring wstr; + wstr.reserve(count); + for (SQLINTEGER i = 0; i < count; ++i) { + wstr.push_back(static_cast(utf16[i])); + } + return Utf16ToUtf8(wstr); + } +#endif + + if (in_str[0] == '\0') { return std::string(); } if (in_str_len == SQL_NTS || in_str_len == NULL) { diff --git a/google/cloud/odbc/bq_driver/internal/utils.h b/google/cloud/odbc/bq_driver/internal/utils.h index a557ce6725..e4eab8398d 100644 --- a/google/cloud/odbc/bq_driver/internal/utils.h +++ b/google/cloud/odbc/bq_driver/internal/utils.h @@ -226,6 +226,26 @@ odbc_internal::StatusRecordOr Utf8ToUtf16( odbc_internal::StatusRecordOr BqConvertSQLWCHARToString( SQLWCHAR* in_str, SQLINTEGER in_str_len); +// Returns true when WcharEncoding=UTF-16LE is set in +// google.googlebigqueryodbc.ini. Always false on Windows. +bool IsRuntimeWireUtf16Le(); + +// Apply the WcharEncoding value read from google.googlebigqueryodbc.ini (or +// the Windows registry equivalent). Accepted values: +// "UTF-16LE" 2-byte wire format (unixODBC loaded under iODBC build) +// "UTF-32LE" 4-byte wire format (native iODBC / wchar_t) +// "" default: use sizeof(SQLWCHAR) as-is +// No-op on Windows. +void SetWcharEncodingFromConfig(std::string const& value); + +// Bytes per wide character on the wire between this driver and its loaded +// manager. Equals sizeof(SQLWCHAR) by default; equals 2 once the UTF-16LE +// wire format has been latched. Use this in *every* arithmetic expression +// that converts between byte counts and character counts on a buffer that +// crosses the driver/manager boundary — never use sizeof(SQLWCHAR) directly +// for that purpose. +size_t WireWcharSize(); + std::wstring SQLWcharToWstring(const SQLWCHAR* in_str); bool IsDiagIdentifierString(SQLSMALLINT DiagIdentifier); diff --git a/google/cloud/odbc/bq_driver/odbc_api.cc b/google/cloud/odbc/bq_driver/odbc_api.cc index 92c4bf0cc1..5a8ccebfe7 100644 --- a/google/cloud/odbc/bq_driver/odbc_api.cc +++ b/google/cloud/odbc/bq_driver/odbc_api.cc @@ -64,6 +64,9 @@ using google::cloud::odbc_bq_driver_internal::IsInfoTypeString; using google::cloud::odbc_bq_driver_internal::StatementHandle; using ::google::cloud::odbc_bq_driver_internal::TraceOptions; using google::cloud::odbc_bq_driver_internal::Utf8ToUtf16; +using google::cloud::odbc_bq_driver_internal::WireWcharSize; +using google::cloud::odbc_bq_driver_internal::WriteWideToWireBuffer; +using google::cloud::odbc_bq_driver_internal::WriteWireNul; using google::cloud::odbc_bq_driver_internal::WStrToOutputBufferResponse; using ::google::cloud::odbc_internal::SQLStates; using google::cloud::odbc_internal::StatusRecord; @@ -74,7 +77,6 @@ using ::google::cloud::odbc_bq_driver::HandleLockError; using google::cloud::odbc_bq_driver::ToCharStr; using google::cloud::odbc_bq_driver::ToSqlChar; -using google::cloud::odbc_bq_driver::ToSqlWChar; constexpr int kBufferLength = 4096; @@ -291,7 +293,9 @@ SQLRETURN SQL_API SQLDriverConnectW( if (!utf16_out_conn_str) { return utf16_out_conn_str.GetCalculatedReturnCode(); } - outConnectionString = ToSqlWChar(utf16_out_conn_str->data()); + + WriteWideToWireBuffer(*utf16_out_conn_str, outConnectionString, + out_conn_str_len); } if (outConnectionStringLen) *outConnectionStringLen = out_conn_str_len; @@ -392,11 +396,16 @@ SQLRETURN SQL_API SQLBrowseConnectW(SQLHDBC connectionHandle, if (!utf16_out_conn_str) { return utf16_out_conn_str.GetCalculatedReturnCode(); } - std::memset(outConnectionString, '\0', - outConnectionStringBufferLen * sizeof(SQLWCHAR)); - std::memcpy((SQLWCHAR*)outConnectionString, - ToSqlWChar(utf16_out_conn_str->data()), - utf16_out_conn_str->size() * sizeof(SQLWCHAR)); + { + size_t const dest_chars = + static_cast(outConnectionStringBufferLen); + size_t const to_copy = + std::min(utf16_out_conn_str->size(), dest_chars); + // memset zeros the entire dest, which leaves the trailing wire NUL in + // place after we write `to_copy` chars. + std::memset(outConnectionString, '\0', dest_chars * WireWcharSize()); + WriteWideToWireBuffer(*utf16_out_conn_str, outConnectionString, to_copy); + } } return rc; @@ -543,7 +552,7 @@ SQLRETURN SQL_API SQLConnectW(SQLHDBC connectionHandle, SQLWCHAR* serverName, return utf16_server_name.GetCalculatedReturnCode(); } serverNameLen = utf16_server_name->length(); - std::memcpy(serverName, ToSqlWChar(utf16_server_name->data()), serverNameLen); + WriteWideToWireBuffer(*utf16_server_name, serverName, serverNameLen); if (w_user_name_len > 0) { StatusRecordOr utf16_user_name = Utf8ToUtf16(*utf8_user_name); @@ -551,7 +560,7 @@ SQLRETURN SQL_API SQLConnectW(SQLHDBC connectionHandle, SQLWCHAR* serverName, return utf16_user_name.GetCalculatedReturnCode(); } userNameLen = utf16_user_name->length(); - std::memcpy(userName, ToSqlWChar(utf16_user_name->data()), userNameLen); + WriteWideToWireBuffer(*utf16_user_name, userName, userNameLen); } if (w_auth_str_len > 0) { @@ -560,7 +569,7 @@ SQLRETURN SQL_API SQLConnectW(SQLHDBC connectionHandle, SQLWCHAR* serverName, return utf16_auth_str.GetCalculatedReturnCode(); } authStringLen = utf16_auth_str->length(); - std::memcpy(authString, ToSqlWChar(utf16_auth_str->data()), authStringLen); + WriteWideToWireBuffer(*utf16_auth_str, authString, authStringLen); } return rc; @@ -643,15 +652,15 @@ SQLRETURN SQL_API SQLGetInfoW(SQLHDBC connectionHandle, SQLUSMALLINT infoType, return utf16_info_val.GetCalculatedReturnCode(); } - std::vector sql_w_str(utf16_info_val->begin(), - utf16_info_val->end()); - sql_w_str.emplace_back(L'\0'); - std::size_t bytes_available = - static_cast(infoValueBufferLen); - std::size_t bytes_to_copy = - std::min(sql_w_str.size() * sizeof(SQLWCHAR), bytes_available); - - std::memcpy(infoValue, sql_w_str.data(), bytes_to_copy); + size_t const wire_sz = WireWcharSize(); + size_t const dest_chars = + static_cast(infoValueBufferLen) / wire_sz; + size_t const to_copy = + std::min(utf16_info_val->size(), dest_chars); + WriteWideToWireBuffer(*utf16_info_val, infoValue, to_copy); + if (to_copy < dest_chars) { + WriteWireNul(infoValue, to_copy); + } } } else { if (info_val_buffer_len > 0) { @@ -662,7 +671,7 @@ SQLRETURN SQL_API SQLGetInfoW(SQLHDBC connectionHandle, SQLUSMALLINT infoType, } } if (infoValueStringLen) - *infoValueStringLen = info_val_buffer_len * sizeof(SQLWCHAR); + *infoValueStringLen = info_val_buffer_len * WireWcharSize(); return rc; } @@ -806,7 +815,7 @@ SQLRETURN SQL_API SQLSetConnectAttrW(SQLHDBC connectionHandle, ConnectionValueType::kSqlChr) { if (valueStringLen && valueStringLen > 0) { updated_attrib_status = - ConvertSQLPointerToSQLChar(value, valueStringLen / sizeof(SQLWCHAR)); + ConvertSQLPointerToSQLChar(value, valueStringLen / WireWcharSize()); } else { updated_attrib_status = ConvertSQLPointerToSQLChar(value, valueStringLen); } @@ -900,9 +909,10 @@ SQLRETURN SQL_API SQLGetConnectAttrW(SQLHDBC connectionHandle, // Handle Unicode conversion of input parameters. // Call to internal common function for SQLGetConnectAttr and // SQLGetConnectAttrW in odbc_connection.h. + SQLINTEGER internal_str_len = 0; rc = ::google::cloud::odbc_bq_driver::SQLGetConnectAttrInternal( - connectionHandle, attribute, updated_attrib_val, valueBufferLen, - valueStringLen); + connectionHandle, attribute, updated_attrib_val, + static_cast(kBufferLength), &internal_str_len); // Handle unicode conversion for attribute string values for output // parameters. if (SQL_SUCCEEDED(rc) && conn_attr.GetAttributeValueType(attribute) == @@ -912,14 +922,18 @@ SQLRETURN SQL_API SQLGetConnectAttrW(SQLHDBC connectionHandle, if (!updated_out_attr_status) { return updated_out_attr_status.GetCalculatedReturnCode(); } - *valueStringLen = - wcslen(updated_out_attr_status->data()) * sizeof(SQLWCHAR); - std::vector sql_w_str( - updated_out_attr_status->c_str(), - updated_out_attr_status->c_str() + *valueStringLen); - sql_w_str.emplace_back(L'\0'); - std::memset(value, '\0', valueBufferLen); - std::memcpy(value, sql_w_str.data(), sql_w_str.size()); + { + size_t const wire_sz = WireWcharSize(); + size_t const dest_chars = static_cast(valueBufferLen) / wire_sz; + size_t const to_copy = + std::min(updated_out_attr_status->size(), dest_chars); + if (valueStringLen) { + *valueStringLen = + static_cast(updated_out_attr_status->size() * wire_sz); + } + std::memset(value, '\0', valueBufferLen); + WriteWideToWireBuffer(*updated_out_attr_status, value, to_copy); + } } return rc; @@ -1151,13 +1165,17 @@ SQLRETURN SQL_API SQLGetDescFieldW(SQLHDESC descriptorHandle, if (!utf16_out_desc_val) { return utf16_out_desc_val.GetCalculatedReturnCode(); } - out_desc_val_string_len = - wcslen(utf16_out_desc_val->data()) * sizeof(SQLWCHAR); - std::vector sql_w_str(utf16_out_desc_val->begin(), - utf16_out_desc_val->end()); - sql_w_str.emplace_back(L'\0'); - std::memset(outDescValue, '\0', outDescValueBufferLen); - std::memcpy(outDescValue, sql_w_str.data(), out_desc_val_string_len); + { + size_t const wire_sz = WireWcharSize(); + size_t const dest_chars = + static_cast(outDescValueBufferLen) / wire_sz; + size_t const to_copy = + std::min(utf16_out_desc_val->size(), dest_chars); + out_desc_val_string_len = + static_cast(utf16_out_desc_val->size() * wire_sz); + std::memset(outDescValue, '\0', outDescValueBufferLen); + WriteWideToWireBuffer(*utf16_out_desc_val, outDescValue, to_copy); + } } else { std::memcpy(outDescValue, (SQLPOINTER)out_desc_val, out_desc_val_string_len); @@ -1238,9 +1256,10 @@ SQLRETURN SQL_API SQLGetDescRecW( if (!utf16_name) { return utf16_name.GetCalculatedReturnCode(); } - std::memset(name, '\0', nameBufferLen * sizeof(SQLWCHAR)); - std::memcpy(name, ToSqlWChar(utf16_name->data()), - name_string_len * sizeof(SQLWCHAR)); + size_t const dest_chars = static_cast(nameBufferLen); + size_t const to_copy = std::min(utf16_name->size(), dest_chars); + std::memset(name, '\0', dest_chars * WireWcharSize()); + WriteWideToWireBuffer(*utf16_name, name, to_copy); } if (nameStringLen) *nameStringLen = name_string_len; @@ -1526,11 +1545,11 @@ SQLRETURN SQL_API SQLGetCursorNameW(SQLHSTMT statementHandle, if (!utf16_cur_name) { return utf16_cur_name.GetCalculatedReturnCode(); } - std::vector sql_w_str(utf16_cur_name->begin(), - utf16_cur_name->end()); - sql_w_str.emplace_back(L'\0'); - std::memcpy(cursorName, sql_w_str.data(), - (sql_w_str.size() + 1) * sizeof(SQLWCHAR)); + { + WriteWideToWireBuffer(*utf16_cur_name, cursorName, + utf16_cur_name->size()); + WriteWireNul(cursorName, utf16_cur_name->size()); + } } if (cursorNameStringLen) *cursorNameStringLen = cursor_name_len; @@ -2103,14 +2122,13 @@ SQLRETURN SQL_API SQLColAttributeW(SQLHSTMT statementHandle, return updated_out_character_attr_status.GetCalculatedReturnCode(); } std::wstring const& wstr = *updated_out_character_attr_status; - size_t const bytes_to_copy = - std::min(static_cast(characterAttributeBufferLen), - wstr.size() * sizeof(SQLWCHAR)); - - std::memcpy(characterAttribute, wstr.data(), bytes_to_copy); - if (characterAttributeBufferLen >= sizeof(SQLWCHAR)) { - SQLWCHAR* wchar_buf = static_cast(characterAttribute); - wchar_buf[bytes_to_copy / sizeof(SQLWCHAR)] = 0; + size_t const wire_sz = WireWcharSize(); + size_t const dest_chars = + static_cast(characterAttributeBufferLen) / wire_sz; + size_t const to_copy = std::min(wstr.size(), dest_chars); + WriteWideToWireBuffer(wstr, characterAttribute, to_copy); + if (static_cast(characterAttributeBufferLen) >= wire_sz) { + WriteWireNul(characterAttribute, to_copy); } character_attribute_string_len = static_cast(wstr.size()); @@ -2123,7 +2141,7 @@ SQLRETURN SQL_API SQLColAttributeW(SQLHSTMT statementHandle, *characterAttributeStringLen = character_attribute_string_len; #ifdef WIN32 *characterAttributeStringLen = - character_attribute_string_len * sizeof(SQLWCHAR); + character_attribute_string_len * WireWcharSize(); #endif // WIN32 } @@ -2202,9 +2220,13 @@ SQLRETURN SQL_API SQLColAttributesW(SQLHSTMT statementHandle, if (!utf16_character_attribute) { return utf16_character_attribute.GetCalculatedReturnCode(); } - std::memcpy(characterAttribute, - (SQLPOINTER)ToSqlWChar(utf16_character_attribute->data()), - character_attribute_buffer_len); + size_t const wire_sz = WireWcharSize(); + size_t const dest_chars = + static_cast(character_attribute_buffer_len) / wire_sz; + size_t const to_copy = + std::min(utf16_character_attribute->size(), dest_chars); + WriteWideToWireBuffer(*utf16_character_attribute, characterAttribute, + to_copy); } if (characterAttributeStringLen) *characterAttributeStringLen = character_attribute_buffer_len; @@ -2285,12 +2307,14 @@ SQLRETURN SQL_API SQLDescribeColW( if (!utf16_col_name) { return utf16_col_name.GetCalculatedReturnCode(); } - std::vector sql_w_str(utf16_col_name->begin(), - utf16_col_name->end()); - sql_w_str.emplace_back(L'\0'); - std::memset(columnName, '\0', columnNameBufferLen); - std::memcpy(columnName, sql_w_str.data(), - column_name_string_len * sizeof(SQLWCHAR)); + { + // columnNameBufferLen is in SQLWCHAR characters per ODBC spec. + size_t const dest_chars = static_cast(columnNameBufferLen); + size_t const to_copy = + std::min(utf16_col_name->size(), dest_chars); + std::memset(columnName, '\0', dest_chars * WireWcharSize()); + WriteWideToWireBuffer(*utf16_col_name, columnName, to_copy); + } } if (columnNameLen) { @@ -2484,7 +2508,7 @@ SQLRETURN SQL_API SQLGetDiagFieldW(SQLSMALLINT handleType, SQLHANDLE handle, // in odbc_diagnostics.h. rc = google::cloud::odbc_bq_driver::SQLGetDiagFieldInternal( handleType, handle, recNumber, diagIdentifier, updated_diag_info, - diagInfoBufferLen, &diag_info_str_len); + static_cast(kBufferLength), &diag_info_str_len); // Handle Unicode conversion of output parameters. if (SQL_SUCCEEDED(rc) && diag_info_str_len > 0) { @@ -2495,13 +2519,16 @@ SQLRETURN SQL_API SQLGetDiagFieldW(SQLSMALLINT handleType, SQLHANDLE handle, if (!updated_out_diag_info_status) { return updated_out_diag_info_status.GetCalculatedReturnCode(); } - diag_info_str_len = - wcslen(updated_out_diag_info_status->data()) * sizeof(SQLWCHAR); - std::vector sql_w_str( - updated_out_diag_info_status->c_str(), - updated_out_diag_info_status->c_str() + diag_info_str_len); - sql_w_str.emplace_back(L'\0'); - std::memcpy(diagInfo, sql_w_str.data(), sql_w_str.size()); + { + size_t const wire_sz = WireWcharSize(); + size_t const dest_chars = + static_cast(diagInfoBufferLen) / wire_sz; + size_t const to_copy = + std::min(updated_out_diag_info_status->size(), dest_chars); + diag_info_str_len = + static_cast(updated_out_diag_info_status->size() * wire_sz); + WriteWideToWireBuffer(*updated_out_diag_info_status, diagInfo, to_copy); + } } else { std::memcpy(diagInfo, updated_diag_info, diagInfoBufferLen); @@ -2587,8 +2614,11 @@ SQLRETURN SQL_API SQLGetDiagRecW(SQLSMALLINT handleType, SQLHANDLE handle, if (!utf16_sql_state) { return utf16_sql_state.GetCalculatedReturnCode(); } - std::memcpy(sqlState, ToSqlWChar(utf16_sql_state->data()), - utf16_sql_state->size() * sizeof(SQLWCHAR)); + { + WriteWideToWireBuffer(*utf16_sql_state, sqlState, + utf16_sql_state->size()); + WriteWireNul(sqlState, utf16_sql_state->size()); + } } if (messageText && message_text_buffer_len > 0) { @@ -2597,9 +2627,14 @@ SQLRETURN SQL_API SQLGetDiagRecW(SQLSMALLINT handleType, SQLHANDLE handle, if (!utf16_msg_txt) { return utf16_msg_txt.GetCalculatedReturnCode(); } - std::memset(messageText, '\0', messageTextBufferLen); - std::memcpy(messageText, ToSqlWChar(utf16_msg_txt->data()), - utf16_msg_txt->size() * sizeof(SQLWCHAR)); + { + // messageTextBufferLen is in SQLWCHAR characters per ODBC spec. + size_t const dest_chars = static_cast(messageTextBufferLen); + size_t const to_copy = + std::min(utf16_msg_txt->size(), dest_chars); + std::memset(messageText, '\0', dest_chars * WireWcharSize()); + WriteWideToWireBuffer(*utf16_msg_txt, messageText, to_copy); + } } if (messageTextLen) *messageTextLen = message_text_buffer_len; diff --git a/google/cloud/odbc/bq_driver/odbc_sql_results.cc b/google/cloud/odbc/bq_driver/odbc_sql_results.cc index c9f1bd4e4f..dc800a116c 100644 --- a/google/cloud/odbc/bq_driver/odbc_sql_results.cc +++ b/google/cloud/odbc/bq_driver/odbc_sql_results.cc @@ -50,6 +50,7 @@ using google::cloud::odbc_bq_driver_internal::StatementHandle; using google::cloud::odbc_bq_driver_internal::StmtStates; using google::cloud::odbc_bq_driver_internal::StringValueToOutputBufferResponse; using google::cloud::odbc_bq_driver_internal::ToSqlPointer; +using google::cloud::odbc_bq_driver_internal::WireWcharSize; using google::cloud::odbc_bq_driver_internal::WriteRowset; using google::cloud::odbc_internal::SQLStates; using google::cloud::odbc_internal::StatusRecord; @@ -776,7 +777,7 @@ SQLRETURN SQLGetDataInternal(SQLHSTMT statement_handle, // 3. If the data fits or is not a variable-length type, return it directly in // the caller’s buffer. SQLLEN target_buff_len = (target_c_type == SQL_C_WCHAR) - ? (target_value_buffer_len / sizeof(SQLWCHAR)) + ? (target_value_buffer_len / WireWcharSize()) : target_value_buffer_len; if (offset == 0) { if ((ds_val.size() > target_buff_len) && @@ -789,7 +790,7 @@ SQLRETURN SQLGetDataInternal(SQLHSTMT statement_handle, size_t buffer_size = 0; if (target_c_type == SQL_C_WCHAR) { - buffer_size = (ds_val.size() + 1) * sizeof(SQLWCHAR); + buffer_size = (ds_val.size() + 1) * WireWcharSize(); } else { buffer_size = ds_val.size() + 1; } @@ -845,18 +846,18 @@ SQLRETURN SQLGetDataInternal(SQLHSTMT statement_handle, result_set.translated_data.row_offset = offset + target_value_buffer_len; } else if (target_c_type == SQL_C_WCHAR) { auto data_size = result_set.translated_data.data.size(); - auto max_buff_chars = target_value_buffer_len / sizeof(SQLWCHAR); - auto offset_chars = offset / sizeof(SQLWCHAR); + auto max_buff_chars = target_value_buffer_len / WireWcharSize(); + auto offset_chars = offset / WireWcharSize(); auto remain_chars = (data_size > offset_chars) ? (data_size - offset_chars) : 0; auto copy_chars = (remain_chars >= max_buff_chars) ? (max_buff_chars - 1) : remain_chars; std::memcpy(target_value, result_set.translated_data.data.data() + offset, - copy_chars * sizeof(SQLWCHAR)); + copy_chars * WireWcharSize()); reinterpret_cast(target_value)[copy_chars] = 0; result_set.translated_data.row_offset = - offset + (copy_chars * sizeof(SQLWCHAR)); + offset + (copy_chars * WireWcharSize()); } else { std::memcpy(target_value, result_set.translated_data.data.data() + offset, target_value_buffer_len - 1); diff --git a/google/cloud/odbc/integration_tests/odbc_driver_tests/catalog_test.cc b/google/cloud/odbc/integration_tests/odbc_driver_tests/catalog_test.cc index 3a32ffe9e3..3052712907 100644 --- a/google/cloud/odbc/integration_tests/odbc_driver_tests/catalog_test.cc +++ b/google/cloud/odbc/integration_tests/odbc_driver_tests/catalog_test.cc @@ -1804,7 +1804,6 @@ TEST(CatalogTest, SQLTables_Filter_DefaultDataset_SchemaNull) { EXPECT_EQ(Disconnect(conn), SQL_SUCCESS); } - #ifdef BQ_DRIVER_INTEGRATION_TESTS // This test case currently crashes with the existing ODBC Driver for BigQuery // v3.1.6.1026. The crash occurs in SQLColumns when schema_name is NULL, 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 eadd203c8a..6f02daee0d 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 @@ -616,6 +616,7 @@ TEST(StatementTest, SQLExecDirect_htapi_bytes_type) { SQLRETURN status; auto conn = std::make_shared(); EXPECT_EQ( + Connect(kDefaultConnectionString + ";AllowHtapiForLargeResults=1;HTAPI_ActivationThreshold=0", conn),