diff --git a/google/cloud/odbc/bq_driver/internal/trace_utils.cc b/google/cloud/odbc/bq_driver/internal/trace_utils.cc index ceec5b5397..9efbcf6f0c 100644 --- a/google/cloud/odbc/bq_driver/internal/trace_utils.cc +++ b/google/cloud/odbc/bq_driver/internal/trace_utils.cc @@ -32,9 +32,13 @@ static std::once_flag absl_log_init_flag; std::shared_ptr TraceOptions::options_file_ = nullptr; std::mutex TraceOptions::mu_; -odbc_internal::StatusRecordOr> const - kTraceOptsFile = - TraceOptions::CreateTraceOptionsFile(GetOdbcTraceConfigPath()); +odbc_internal::StatusRecordOr>& +GetTraceOptsFile() { + static auto* trace_opts = + new odbc_internal::StatusRecordOr>( + TraceOptions::CreateTraceOptionsFile(GetOdbcTraceConfigPath())); + return *trace_opts; +} #ifdef _WIN32 constexpr char kPathSeparator = '\\'; @@ -67,6 +71,10 @@ FileLogSink::FileLogSink(std::shared_ptr opts) } FileLogSink::~FileLogSink() { + if (is_registered_) { + absl::log_internal::RemoveLogSink(this); + is_registered_ = false; + } // Close the file pointer if it was opened if (fp_ != nullptr) { fclose(fp_); @@ -165,11 +173,11 @@ void UpdateTraceOption(std::optional log_level, std::optional log_file_size, std::optional log_file_count, std::optional max_threads) { - if (!kTraceOptsFile.Ok() || !(log_level || log_path || log_file_size || - log_file_count || max_threads)) + if (!GetTraceOptsFile().Ok() || !(log_level || log_path || log_file_size || + log_file_count || max_threads)) return; - auto const& trace_options = kTraceOptsFile.GetValue(); + auto const& trace_options = GetTraceOptsFile().GetValue(); std::lock_guard lock(trace_options->m); if (log_level) { @@ -190,8 +198,8 @@ std::string GetLogFileWithIndex(std::string const& log_path) { std::string base_dir = log_path; int file_index = 0; - if (kTraceOptsFile.Ok()) { - auto const& trace_opts = kTraceOptsFile.GetValue(); + if (GetTraceOptsFile().Ok()) { + auto const& trace_opts = GetTraceOptsFile().GetValue(); file_index = trace_opts->current_file_index; } std::string separator = @@ -204,15 +212,16 @@ std::string GetLogFileWithIndex(std::string const& log_path) { void FileLogSink::InitializeFileLog( std::shared_ptr const& trace_opts) { - if (file_sink_ || !trace_opts) return; + if (!trace_opts) return; - if (file_sink_) { - absl::log_internal::RemoveLogSink(file_sink_.get()); - file_sink_ = nullptr; - } + file_sink_ = nullptr; - file_sink_ = std::make_unique(trace_opts); - absl::log_internal::AddLogSink(file_sink_.get()); + auto new_sink = std::make_unique(trace_opts); + if (new_sink->IsOpen()) { + absl::log_internal::AddLogSink(new_sink.get()); + new_sink->is_registered_ = true; + } + file_sink_ = std::move(new_sink); } bool TraceOptions::InitializeLogging(bool is_trace_override) { @@ -220,8 +229,8 @@ bool TraceOptions::InitializeLogging(bool is_trace_override) { std::call_once(absl_log_init_flag, []() { absl::InitializeLog(); }); absl::SetStderrThreshold(absl::LogSeverityAtLeast::kInfinity); - if (!kTraceOptsFile.Ok()) return false; - auto const& trace_opts = kTraceOptsFile.GetValue(); + if (!GetTraceOptsFile().Ok()) return false; + auto const& trace_opts = GetTraceOptsFile().GetValue(); // If logging is disabled, return false if (trace_opts->log_level <= 0) { diff --git a/google/cloud/odbc/bq_driver/internal/trace_utils.h b/google/cloud/odbc/bq_driver/internal/trace_utils.h index 97348a5a7a..f35f9188dd 100644 --- a/google/cloud/odbc/bq_driver/internal/trace_utils.h +++ b/google/cloud/odbc/bq_driver/internal/trace_utils.h @@ -147,6 +147,7 @@ class FileLogSink : public absl::LogSink { void Send(absl::LogEntry const& entry) override; [[nodiscard]] int GetLogLevel() const { return opts_->log_level; } + [[nodiscard]] bool IsOpen() const { return fp_ != nullptr; } static void InitializeFileLog( std::shared_ptr const& trace_opts); @@ -159,6 +160,7 @@ class FileLogSink : public absl::LogSink { std::size_t current_file_size_; std::mutex log_mutex_; FILE* fp_ = nullptr; + bool is_registered_ = false; }; // Get abseil severity as per internal driver log levels @@ -198,8 +200,8 @@ std::string GetFormattedMsg(absl::LogEntry const& entry); // Struct types. ///////////////////////////////////////////// -extern odbc_internal::StatusRecordOr> const - kTraceOptsFile; +odbc_internal::StatusRecordOr>& +GetTraceOptsFile(); } // namespace google::cloud::odbc_bq_driver_internal diff --git a/google/cloud/odbc/bq_driver/internal/utils.cc b/google/cloud/odbc/bq_driver/internal/utils.cc index 91388d1019..32bc9d1c73 100644 --- a/google/cloud/odbc/bq_driver/internal/utils.cc +++ b/google/cloud/odbc/bq_driver/internal/utils.cc @@ -79,19 +79,13 @@ size_t WireWcharSize() { void SetWcharEncodingFromConfig(std::string const& value) { if (value == "UTF-8" || value == "UTF8") { g_wire_encoding.store(WireEncoding::kUtf8, std::memory_order_relaxed); - LOG(INFO) << "WcharEncoding: UTF-8 wire format (1 byte/char)"; } else if (value == "UTF-16LE" || value == "UTF16LE" || value == "UTF-16") { g_wire_encoding.store(WireEncoding::kUtf16Le, std::memory_order_relaxed); - LOG(INFO) << "WcharEncoding: UTF-16LE wire format (2 bytes/char)"; } else if (value == "UTF-32LE" || value == "UTF32LE" || value == "UTF-32" || value == "UCS-4LE") { g_wire_encoding.store(WireEncoding::kUtf32Le, std::memory_order_relaxed); - LOG(INFO) << "WcharEncoding: UTF-32LE wire format (4 bytes/char)"; } else if (value.empty() || value == "default") { g_wire_encoding.store(WireEncoding::kDefault, std::memory_order_relaxed); - LOG(INFO) << "WcharEncoding: default (sizeof(SQLWCHAR) bytes/char)"; - } else { - LOG(WARNING) << "WcharEncoding: unrecognised value '" << value << "'"; } } #endif diff --git a/google/cloud/odbc/bq_driver/odbc_connection.cc b/google/cloud/odbc/bq_driver/odbc_connection.cc index ceddacd470..9baa70a973 100644 --- a/google/cloud/odbc/bq_driver/odbc_connection.cc +++ b/google/cloud/odbc/bq_driver/odbc_connection.cc @@ -37,9 +37,8 @@ using google::cloud::odbc_bq_driver_internal::DescriptorHandle; using google::cloud::odbc_bq_driver_internal::Dsn; using google::cloud::odbc_bq_driver_internal::EnvironmentHandle; using google::cloud::odbc_bq_driver_internal::GetDefaultPemFile; -using google::cloud::odbc_bq_driver_internal::GetMissingAttributesStr; +using google::cloud::odbc_bq_driver_internal::GetTraceOptsFile; using google::cloud::odbc_bq_driver_internal::GetUpperStr; -using google::cloud::odbc_bq_driver_internal::kTraceOptsFile; using google::cloud::odbc_bq_driver_internal::LogAndReturnCode; using google::cloud::odbc_bq_driver_internal::PopulateOutputConnectionString; using google::cloud::odbc_bq_driver_internal::Section; @@ -325,8 +324,8 @@ SQLRETURN SQLDriverConnectInternal(SQLHDBC conn_handle, SQLHWND window_handle, dsn_section[property] = it.second; } } - if (kTraceOptsFile.Ok()) { - auto const& trace_options = kTraceOptsFile.GetValue(); + if (GetTraceOptsFile().Ok()) { + auto const& trace_options = GetTraceOptsFile().GetValue(); if (!trace_options->logging_enabled) { config_res = ConfigTraceFromSection(dsn_section); } diff --git a/google/cloud/odbc/bq_driver/odbc_driver_metadata.cc b/google/cloud/odbc/bq_driver/odbc_driver_metadata.cc index 3b6ae05f64..a82bbceacc 100644 --- a/google/cloud/odbc/bq_driver/odbc_driver_metadata.cc +++ b/google/cloud/odbc/bq_driver/odbc_driver_metadata.cc @@ -443,11 +443,6 @@ SQLRETURN SQLTablesInternal(SQLHSTMT stmt_handle, SQLCHAR* catalog_name, return LogAndReturnCode(handle, input_param_status); } - std::string project_filter = ToCharStr(catalog_name, kMatchAll); - std::string dataset_filter = ToCharStr(schema_name, kMatchAll); - std::string table_filter = ToCharStr(table_name, kMatchAll); - std::string table_type_filter = ToCharStr(table_type, kMatchAll); - if (handle.GetConnectionHandle() == nullptr) { LOG(ERROR) << "SQLTables:: Internal connection handle is null"; return LogAndReturnCode(handle, @@ -455,6 +450,25 @@ SQLRETURN SQLTablesInternal(SQLHSTMT stmt_handle, SQLCHAR* catalog_name, "Internal connection handle is null"}); } ConnectionHandle& conn_handle = *(handle.GetConnectionHandle()); + std::string catalog_str; + if (catalog_name == nullptr || catalog_name_len == 0) { + SQLINTEGER catalog_len = 0; + SQLCHAR current_catalog[256] = {0}; + conn_handle.GetAttribute(SQL_ATTR_CURRENT_CATALOG, current_catalog, + sizeof(current_catalog), &catalog_len); + + if (catalog_len > 0) { + catalog_str.assign(reinterpret_cast(current_catalog), catalog_len); + catalog_name = reinterpret_cast(catalog_str.data()); + catalog_name_len = static_cast(catalog_str.size()); + } + } + + std::string project_filter = ToCharStr(catalog_name, kMatchAll); + std::string dataset_filter = ToCharStr(schema_name, kMatchAll); + std::string table_filter = ToCharStr(table_name, kMatchAll); + std::string table_type_filter = ToCharStr(table_type, kMatchAll); + if (!metadata_id && dataset_filter == kMatchAll) { auto const dsn = conn_handle.GetDsn(); if (dsn.filter_tables_on_default_dataset && !dsn.default_dataset.empty()) { 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 a272f2f362..80a8bccdf7 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 @@ -2164,4 +2164,55 @@ TEST(SQLTables, Check_SQLTablesDescriptors) { EXPECT_EQ(Disconnect(conn), SQL_SUCCESS); } +TEST(CatalogTest, SQLTables_NullCatalogFiltersToCurrentProject) { + auto conn = std::make_shared(); + ASSERT_EQ(Connect(kDefaultConnectionString, conn), SQL_SUCCESS); + + SQLRETURN status = SQLSetStmtAttr(conn->hstmt, SQL_ATTR_METADATA_ID, + (SQLPOINTER)SQL_FALSE, 0); + CheckError(status, "SQLSetStmtAttr", conn); + + SQLCHAR current_catalog[256] = {0}; + SQLINTEGER catalog_len = 0; + SQLRETURN attr_status = + SQLGetConnectAttr(conn->hdbc, SQL_ATTR_CURRENT_CATALOG, current_catalog, + sizeof(current_catalog), &catalog_len); + + ASSERT_TRUE(SQL_SUCCEEDED(attr_status)) + << "Failed to get SQL_ATTR_CURRENT_CATALOG"; + std::string expected_catalog(reinterpret_cast(current_catalog)); + + SQLCHAR table_type[] = "TABLE,VIEW"; + SQLRETURN rc = SQLTables(conn->hstmt, NULL, 0, // Catalog (NULL) + NULL, 0, // Schema + NULL, 0, // Table name + table_type, SQL_NTS); // Table type + + ASSERT_TRUE(SQL_SUCCEEDED(rc)) << "SQLTables call failed."; + + SQLCHAR out_catalog[256] = {0}; + SQLLEN out_len = 0; + SQLBindCol(conn->hstmt, 1, SQL_C_CHAR, out_catalog, sizeof(out_catalog), + &out_len); + + int row_count = 0; + bool foreign_catalog_found = false; + + while (SQLFetch(conn->hstmt) == SQL_SUCCESS) { + row_count++; + std::string fetched_catalog(reinterpret_cast(out_catalog)); + + if (fetched_catalog != expected_catalog) { + foreign_catalog_found = true; + } + } + + EXPECT_FALSE(foreign_catalog_found) + << "SQLTables returned data for projects outside the configured DSN."; + EXPECT_GT(row_count, 0) + << "Expected to find at least one table/view in the default project."; + + EXPECT_EQ(Disconnect(conn), SQL_SUCCESS); +} + } // namespace google::cloud::odbc_tests diff --git a/google/cloud/odbc/integration_tests/odbc_driver_tests/examples/catalog_performance_example.cc b/google/cloud/odbc/integration_tests/odbc_driver_tests/examples/catalog_performance_example.cc index 84283edcf6..f3139a51d1 100644 --- a/google/cloud/odbc/integration_tests/odbc_driver_tests/examples/catalog_performance_example.cc +++ b/google/cloud/odbc/integration_tests/odbc_driver_tests/examples/catalog_performance_example.cc @@ -373,19 +373,17 @@ TEST_P(DataFetchPerformanceParamTest, Benchmark) { DescribeCol(conn, col_ptr, i); SqlToCdataTypes(col_ptr); + + ret = SQLBindCol( + conn->hstmt, i, col_ptr->data_type, col_ptr->data_buf.target_value, + col_ptr->data_buf.buffer_length, &(col_ptr->data_buf.str_len)); + CheckError(ret, "SQLBindCol(" + std::to_string(i) + ")", conn); } int row_count = 0; auto fetch_start = std::chrono::high_resolution_clock::now(); while ((ret = SQLFetch(conn->hstmt)) == SQL_SUCCESS || ret == SQL_SUCCESS_WITH_INFO) { - for (int i = 1; i <= num_cols; i++) { - auto const& col_ptr = cols[i - 1]; - SQLRETURN get_data_ret = SQLGetData( - conn->hstmt, i, col_ptr->data_type, col_ptr->data_buf.target_value, - col_ptr->data_buf.buffer_length, &(col_ptr->data_buf.str_len)); - CheckError(get_data_ret, "SQLGetData(" + std::to_string(i) + ")", conn); - } row_count++; } auto fetch_end = std::chrono::high_resolution_clock::now(); @@ -472,10 +470,10 @@ INSTANTIATE_TEST_SUITE_P( return std::get<0>(info.param); }); -class DataFetchBindColPerformanceParamTest +class DataFetchPerformanceParamTest_WithSQLGetData : public ::testing::TestWithParam {}; -TEST_P(DataFetchBindColPerformanceParamTest, Benchmark) { +TEST_P(DataFetchPerformanceParamTest_WithSQLGetData, Benchmark) { auto conn = std::make_shared(); std::string connection_string = @@ -509,17 +507,19 @@ TEST_P(DataFetchBindColPerformanceParamTest, Benchmark) { DescribeCol(conn, col_ptr, i); SqlToCdataTypes(col_ptr); - - ret = SQLBindCol( - conn->hstmt, i, col_ptr->data_type, col_ptr->data_buf.target_value, - col_ptr->data_buf.buffer_length, &(col_ptr->data_buf.str_len)); - CheckError(ret, "SQLBindCol(" + std::to_string(i) + ")", conn); } int row_count = 0; auto fetch_start = std::chrono::high_resolution_clock::now(); while ((ret = SQLFetch(conn->hstmt)) == SQL_SUCCESS || ret == SQL_SUCCESS_WITH_INFO) { + for (int i = 1; i <= num_cols; i++) { + auto const& col_ptr = cols[i - 1]; + SQLRETURN get_data_ret = SQLGetData( + conn->hstmt, i, col_ptr->data_type, col_ptr->data_buf.target_value, + col_ptr->data_buf.buffer_length, &(col_ptr->data_buf.str_len)); + CheckError(get_data_ret, "SQLGetData(" + std::to_string(i) + ")", conn); + } row_count++; } auto fetch_end = std::chrono::high_resolution_clock::now(); @@ -542,9 +542,9 @@ TEST_P(DataFetchBindColPerformanceParamTest, Benchmark) { EXPECT_EQ(Disconnect(conn), SQL_SUCCESS); } -inline std::vector GetDataFetchBindColBenchmarkParams() { +inline std::vector GetDataFetchSQLGetDataBenchmarkParams() { std::vector const benchmark_configs = { - {"all_bq_types_2", + {"all_bq_types_2_SQLGetData", "SELECT * FROM " "`bigquery-devtools-drivers.INTEGRATION_TEST_FORMAT.all_bq_types_2`", {{"1M", 1000000}}}, @@ -563,8 +563,8 @@ inline std::vector GetDataFetchBindColBenchmarkParams() { } INSTANTIATE_TEST_SUITE_P( - , DataFetchBindColPerformanceParamTest, - ::testing::ValuesIn(GetDataFetchBindColBenchmarkParams()), + , DataFetchPerformanceParamTest_WithSQLGetData, + ::testing::ValuesIn(GetDataFetchSQLGetDataBenchmarkParams()), [](::testing::TestParamInfo const& info) { return std::get<0>(info.param); });