diff --git a/cpp/src/arrow/flight/sql/odbc/odbc_api.cc b/cpp/src/arrow/flight/sql/odbc/odbc_api.cc index e19553160f07..75924f596c4a 100644 --- a/cpp/src/arrow/flight/sql/odbc/odbc_api.cc +++ b/cpp/src/arrow/flight/sql/odbc/odbc_api.cc @@ -107,9 +107,26 @@ SQLRETURN SQLAllocHandle(SQLSMALLINT type, SQLHANDLE parent, SQLHANDLE* result) }); } - // TODO Implement for case of descriptor - case SQL_HANDLE_DESC: - return SQL_INVALID_HANDLE; + case SQL_HANDLE_DESC: { + using ODBC::ODBCConnection; + using ODBC::ODBCDescriptor; + + *result = SQL_NULL_HDESC; + + ODBCConnection* connection = reinterpret_cast(parent); + + return ODBCConnection::ExecuteWithDiagnostics(connection, SQL_ERROR, [=]() { + std::shared_ptr descriptor = connection->createDescriptor(); + + if (descriptor) { + *result = reinterpret_cast(descriptor.get()); + + return SQL_SUCCESS; + } + + return SQL_ERROR; + }); + } default: break; @@ -164,8 +181,19 @@ SQLRETURN SQLFreeHandle(SQLSMALLINT type, SQLHANDLE handle) { return SQL_SUCCESS; } - case SQL_HANDLE_DESC: - return SQL_INVALID_HANDLE; + case SQL_HANDLE_DESC: { + using ODBC::ODBCDescriptor; + + ODBCDescriptor* descriptor = reinterpret_cast(handle); + + if (!descriptor) { + return SQL_INVALID_HANDLE; + } + + descriptor->ReleaseDescriptor(); + + return SQL_SUCCESS; + } default: break; @@ -242,6 +270,7 @@ SQLRETURN SQLGetDiagField(SQLSMALLINT handleType, SQLHANDLE handle, SQLSMALLINT using driver::odbcabstraction::Diagnostics; using ODBC::GetStringAttribute; using ODBC::ODBCConnection; + using ODBC::ODBCDescriptor; using ODBC::ODBCEnvironment; using ODBC::ODBCStatement; @@ -277,7 +306,9 @@ SQLRETURN SQLGetDiagField(SQLSMALLINT handleType, SQLHANDLE handle, SQLSMALLINT } case SQL_HANDLE_DESC: { - return SQL_ERROR; + ODBCDescriptor* descriptor = reinterpret_cast(handle); + diagnostics = &descriptor->GetDiagnostics(); + break; } case SQL_HANDLE_STMT: { @@ -405,8 +436,12 @@ SQLRETURN SQLGetDiagField(SQLSMALLINT handleType, SQLHANDLE handle, SQLSMALLINT } case SQL_HANDLE_DESC: { - // TODO Implement for case of descriptor - return SQL_ERROR; + ODBCDescriptor* descriptor = reinterpret_cast(handle); + ODBCConnection* connection = &descriptor->GetConnection(); + std::string dsn = connection->GetDSN(); + return GetStringAttribute(isUnicode, dsn, true, diagInfoPtr, bufferLength, + stringLengthPtr, *diagnostics); + break; } case SQL_HANDLE_STMT: { @@ -495,6 +530,7 @@ SQLRETURN SQLGetDiagRec(SQLSMALLINT handleType, SQLHANDLE handle, SQLSMALLINT re using driver::odbcabstraction::Diagnostics; using ODBC::GetStringAttribute; using ODBC::ODBCConnection; + using ODBC::ODBCDescriptor; using ODBC::ODBCEnvironment; using ODBC::ODBCStatement; @@ -525,7 +561,9 @@ SQLRETURN SQLGetDiagRec(SQLSMALLINT handleType, SQLHANDLE handle, SQLSMALLINT re } case SQL_HANDLE_DESC: { - return SQL_ERROR; + auto* descriptor = ODBCDescriptor::of(handle); + diagnostics = &descriptor->GetDiagnostics(); + break; } case SQL_HANDLE_STMT: { diff --git a/cpp/src/arrow/flight/sql/odbc/odbcabstraction/include/odbcabstraction/odbc_impl/odbc_statement.h b/cpp/src/arrow/flight/sql/odbc/odbcabstraction/include/odbcabstraction/odbc_impl/odbc_statement.h index 73cdc2448f8f..7fb8d5c57415 100644 --- a/cpp/src/arrow/flight/sql/odbc/odbcabstraction/include/odbcabstraction/odbc_impl/odbc_statement.h +++ b/cpp/src/arrow/flight/sql/odbc/odbcabstraction/include/odbcabstraction/odbc_impl/odbc_statement.h @@ -77,6 +77,11 @@ class ODBCStatement : public ODBCHandle { void SetStmtAttr(SQLINTEGER statementAttribute, SQLPOINTER value, SQLINTEGER bufferSize, bool isUnicode); + /** + * @brief Revert back to implicitly allocated internal descriptors. + * isApd as True indicates APD descritor is to be reverted. + * isApd as False indicates ARD descritor is to be reverted. + */ void RevertAppDescriptor(bool isApd); inline ODBCDescriptor* GetIRD() { return m_ird.get(); } diff --git a/cpp/src/arrow/flight/sql/odbc/tests/connection_test.cc b/cpp/src/arrow/flight/sql/odbc/tests/connection_test.cc index c935646de5b7..5bf737d55182 100644 --- a/cpp/src/arrow/flight/sql/odbc/tests/connection_test.cc +++ b/cpp/src/arrow/flight/sql/odbc/tests/connection_test.cc @@ -919,6 +919,105 @@ TYPED_TEST(FlightSQLODBCTestBase, TestCloseConnectionWithOpenStatement) { EXPECT_EQ(ret, SQL_SUCCESS); } +TYPED_TEST(FlightSQLODBCTestBase, TestSQLAllocFreeDesc) { + this->connect(); + SQLHDESC descriptor; + + // Allocate a descriptor using alloc handle + SQLRETURN ret = SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + // Free descriptor handle + ret = SQLFreeHandle(SQL_HANDLE_DESC, descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + this->disconnect(); +} + +TYPED_TEST(FlightSQLODBCTestBase, TestSQLSetStmtAttrDescriptor) { + this->connect(); + + SQLHDESC apd_descriptor, ard_descriptor; + + // Allocate an APD descriptor using alloc handle + SQLRETURN ret = SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &apd_descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + // Allocate an ARD descriptor using alloc handle + ret = SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &ard_descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + // Save implicitly allocated internal APD and ARD descriptor pointers + SQLPOINTER internal_apd, internal_ard = nullptr; + + ret = SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, &internal_apd, + sizeof(internal_apd), 0); + + EXPECT_EQ(ret, SQL_SUCCESS); + + ret = SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, &internal_ard, + sizeof(internal_ard), 0); + + EXPECT_EQ(ret, SQL_SUCCESS); + + // Set APD descriptor to explicitly allocated handle + ret = SQLSetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, + reinterpret_cast(apd_descriptor), 0); + + EXPECT_EQ(ret, SQL_SUCCESS); + + // Set ARD descriptor to explicitly allocated handle + ret = SQLSetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, + reinterpret_cast(ard_descriptor), 0); + + EXPECT_EQ(ret, SQL_SUCCESS); + + // Verify APD and ARD descriptors are set to explicitly allocated pointers + SQLPOINTER value = nullptr; + + ret = SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, &value, sizeof(value), 0); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_EQ(value, apd_descriptor); + + ret = SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, &value, sizeof(value), 0); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_EQ(value, ard_descriptor); + + // Free explicitly allocated APD and ARD descriptor handles + ret = SQLFreeHandle(SQL_HANDLE_DESC, apd_descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + ret = SQLFreeHandle(SQL_HANDLE_DESC, ard_descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + // Verify APD and ARD descriptors has been reverted to implicit descriptors + value = nullptr; + + ret = SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, &value, sizeof(value), 0); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_EQ(value, internal_apd); + + ret = SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, &value, sizeof(value), 0); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_EQ(value, internal_ard); + + this->disconnect(); +} + } // namespace arrow::flight::sql::odbc int main(int argc, char** argv) { diff --git a/cpp/src/arrow/flight/sql/odbc/tests/errors_test.cc b/cpp/src/arrow/flight/sql/odbc/tests/errors_test.cc index ee0a2846194d..276c16a113ff 100644 --- a/cpp/src/arrow/flight/sql/odbc/tests/errors_test.cc +++ b/cpp/src/arrow/flight/sql/odbc/tests/errors_test.cc @@ -62,7 +62,7 @@ TYPED_TEST(FlightSQLODBCTestBase, TestSQLGetDiagFieldWForConnectFailure) { static_cast(connect_str0.size()), outstr, ODBC_BUFFER_SIZE, &outstrlen, SQL_DRIVER_NOPROMPT); - EXPECT_TRUE(ret == SQL_ERROR); + EXPECT_EQ(ret, SQL_ERROR); // Retrieve all supported header level and record level data SQLSMALLINT HEADER_LEVEL = 0; @@ -170,7 +170,7 @@ TYPED_TEST(FlightSQLODBCTestBase, TestSQLGetDiagFieldWForConnectFailureNTS) { static_cast(connect_str0.size()), outstr, ODBC_BUFFER_SIZE, &outstrlen, SQL_DRIVER_NOPROMPT); - EXPECT_TRUE(ret == SQL_ERROR); + EXPECT_EQ(ret, SQL_ERROR); // Retrieve all supported header level and record level data SQLSMALLINT RECORD_1 = 1; @@ -199,6 +199,127 @@ TYPED_TEST(FlightSQLODBCTestBase, TestSQLGetDiagFieldWForConnectFailureNTS) { EXPECT_EQ(ret, SQL_SUCCESS); } +TYPED_TEST(FlightSQLODBCTestBase, + TestSQLGetDiagFieldWForDescriptorFailureFromDriverManager) { + this->connect(); + SQLHDESC descriptor; + + // Allocate a descriptor using alloc handle + SQLRETURN ret = SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + ret = SQLGetDescField(descriptor, 1, SQL_DESC_DATETIME_INTERVAL_CODE, 0, 0, 0); + + EXPECT_EQ(ret, SQL_ERROR); + + // Retrieve all supported header level and record level data + SQLSMALLINT HEADER_LEVEL = 0; + SQLSMALLINT RECORD_1 = 1; + + // SQL_DIAG_NUMBER + SQLINTEGER diag_number; + SQLSMALLINT diag_number_length; + + ret = SQLGetDiagField(SQL_HANDLE_DESC, descriptor, HEADER_LEVEL, SQL_DIAG_NUMBER, + &diag_number, sizeof(SQLINTEGER), &diag_number_length); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_EQ(diag_number, 1); + + // SQL_DIAG_SERVER_NAME + SQLWCHAR server_name[ODBC_BUFFER_SIZE]; + SQLSMALLINT server_name_length; + + ret = SQLGetDiagField(SQL_HANDLE_DESC, descriptor, RECORD_1, SQL_DIAG_SERVER_NAME, + server_name, ODBC_BUFFER_SIZE, &server_name_length); + + EXPECT_EQ(ret, SQL_SUCCESS); + + // SQL_DIAG_MESSAGE_TEXT + SQLWCHAR message_text[ODBC_BUFFER_SIZE]; + SQLSMALLINT message_text_length; + + ret = SQLGetDiagField(SQL_HANDLE_DESC, descriptor, RECORD_1, SQL_DIAG_MESSAGE_TEXT, + message_text, ODBC_BUFFER_SIZE, &message_text_length); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_GT(message_text_length, 100); + + // SQL_DIAG_NATIVE + SQLINTEGER diag_native; + SQLSMALLINT diag_native_length; + + ret = SQLGetDiagField(SQL_HANDLE_DESC, descriptor, RECORD_1, SQL_DIAG_NATIVE, + &diag_native, sizeof(diag_native), &diag_native_length); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_EQ(diag_native, 0); + + // SQL_DIAG_SQLSTATE + const SQLSMALLINT sql_state_size = 6; + SQLWCHAR sql_state[sql_state_size]; + SQLSMALLINT sql_state_length; + ret = SQLGetDiagField( + SQL_HANDLE_DESC, descriptor, RECORD_1, SQL_DIAG_SQLSTATE, sql_state, + sql_state_size * driver::odbcabstraction::GetSqlWCharSize(), &sql_state_length); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_EQ(std::wstring(sql_state), std::wstring(L"IM001")); + + // Free descriptor handle + ret = SQLFreeHandle(SQL_HANDLE_DESC, descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + this->disconnect(); +} + +TYPED_TEST(FlightSQLODBCTestBase, + TestSQLGetDiagRecForDescriptorFailureFromDriverManager) { + this->connect(); + SQLHDESC descriptor; + + // Allocate a descriptor using alloc handle + SQLRETURN ret = SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + ret = SQLGetDescField(descriptor, 1, SQL_DESC_DATETIME_INTERVAL_CODE, 0, 0, 0); + + EXPECT_EQ(ret, SQL_ERROR); + + SQLWCHAR sql_state[6]; + SQLINTEGER native_error; + SQLWCHAR message[ODBC_BUFFER_SIZE]; + SQLSMALLINT message_length; + + ret = SQLGetDiagRec(SQL_HANDLE_DESC, descriptor, 1, sql_state, &native_error, message, + ODBC_BUFFER_SIZE, &message_length); + + EXPECT_EQ(ret, SQL_SUCCESS); + + EXPECT_GT(message_length, 60); + + EXPECT_EQ(native_error, 0); + + // API not implemented error from driver manager + EXPECT_EQ(std::wstring(sql_state), std::wstring(L"IM001")); + + EXPECT_TRUE(!std::wstring(message).empty()); + + // Free descriptor handle + ret = SQLFreeHandle(SQL_HANDLE_DESC, descriptor); + + EXPECT_EQ(ret, SQL_SUCCESS); + + this->disconnect(); +} + TYPED_TEST(FlightSQLODBCTestBase, TestSQLGetDiagRecForConnectFailure) { // ODBC Environment SQLHENV env; @@ -233,7 +354,7 @@ TYPED_TEST(FlightSQLODBCTestBase, TestSQLGetDiagRecForConnectFailure) { static_cast(connect_str0.size()), outstr, ODBC_BUFFER_SIZE, &outstrlen, SQL_DRIVER_NOPROMPT); - EXPECT_TRUE(ret == SQL_ERROR); + EXPECT_EQ(ret, SQL_ERROR); SQLWCHAR sql_state[6]; SQLINTEGER native_error;