diff --git a/cpp/src/arrow/flight/sql/odbc/odbc_impl/odbc_statement.cc b/cpp/src/arrow/flight/sql/odbc/odbc_impl/odbc_statement.cc index 43a3a95b8f10..7b0d868afe4c 100644 --- a/cpp/src/arrow/flight/sql/odbc/odbc_impl/odbc_statement.cc +++ b/cpp/src/arrow/flight/sql/odbc/odbc_impl/odbc_statement.cc @@ -241,7 +241,7 @@ ODBCStatement::ODBCStatement(ODBCConnection& connection, ird_(std::make_shared(spi_statement_->GetDiagnostics(), nullptr, this, false, false, connection.IsOdbc2Connection())), - current_ard_(built_in_apd_.get()), + current_ard_(built_in_ard_.get()), current_apd_(built_in_apd_.get()), row_number_(0), max_rows_(0), @@ -323,6 +323,10 @@ void ODBCStatement::ExecuteDirect(const std::string& query) { bool ODBCStatement::Fetch(size_t rows, SQLULEN* row_count_ptr, SQLUSMALLINT* row_status_array) { + if (!current_result_) { + throw DriverException("Invalid cursor state", "24000"); + } + if (has_reached_end_of_result_) { ird_->SetRowsProcessed(0); return false; @@ -558,8 +562,9 @@ void ODBCStatement::SetStmtAttr(SQLINTEGER statement_attribute, SQLPOINTER value switch (statement_attribute) { case SQL_ATTR_APP_PARAM_DESC: { - ODBCDescriptor* desc = static_cast(value); - if (desc && current_apd_ != desc) { + ODBCDescriptor* desc = + value ? static_cast(value) : built_in_apd_.get(); + if (current_apd_ != desc) { if (current_apd_ != built_in_apd_.get()) { current_apd_->DetachFromStatement(this, true); } @@ -571,8 +576,9 @@ void ODBCStatement::SetStmtAttr(SQLINTEGER statement_attribute, SQLPOINTER value return; } case SQL_ATTR_APP_ROW_DESC: { - ODBCDescriptor* desc = static_cast(value); - if (desc && current_ard_ != desc) { + ODBCDescriptor* desc = + value ? static_cast(value) : built_in_ard_.get(); + if (current_ard_ != desc) { if (current_ard_ != built_in_ard_.get()) { current_ard_->DetachFromStatement(this, false); } @@ -740,6 +746,10 @@ void ODBCStatement::CloseCursor(bool suppress_errors) { SQLRETURN ODBCStatement::GetData(SQLSMALLINT record_number, SQLSMALLINT c_type, SQLPOINTER data_ptr, SQLLEN buffer_length, SQLLEN* indicator_ptr) { + if (!current_result_) { + throw DriverException("Invalid cursor state", "24000"); + } + if (record_number == 0) { throw DriverException("Bookmarks are not supported", "07009"); } else if (static_cast(record_number) > ird_->GetRecords().size()) { @@ -785,6 +795,7 @@ SQLRETURN ODBCStatement::GetData(SQLSMALLINT record_number, SQLSMALLINT c_type, SQLRETURN ODBCStatement::GetMoreResults() { // Multiple result sets are not supported by Arrow protocol. + CloseCursor(/*suppress_errors=*/true); return SQL_NO_DATA; } @@ -811,6 +822,12 @@ void ODBCStatement::GetRowCount(SQLLEN* row_count_ptr) { void ODBCStatement::ReleaseStatement() { CloseCursor(true); + if (current_apd_ != built_in_apd_.get()) { + current_apd_->DetachFromStatement(this, true); + } + if (current_ard_ != built_in_ard_.get()) { + current_ard_->DetachFromStatement(this, false); + } connection_.DropStatement(this); } 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 837c18c9cbdf..121cdf3b5d33 100644 --- a/cpp/src/arrow/flight/sql/odbc/tests/connection_test.cc +++ b/cpp/src/arrow/flight/sql/odbc/tests/connection_test.cc @@ -528,6 +528,7 @@ TYPED_TEST(ConnectionTest, TestSQLSetStmtAttrDescriptor) { EXPECT_EQ(SQL_SUCCESS, SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, &internal_ard, sizeof(internal_ard), 0)); + EXPECT_NE(internal_apd, internal_ard); // Set APD descriptor to explicitly allocated handle EXPECT_EQ(SQL_SUCCESS, SQLSetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, @@ -549,12 +550,41 @@ TYPED_TEST(ConnectionTest, TestSQLSetStmtAttrDescriptor) { EXPECT_EQ(ard_descriptor, value); - // Free explicitly allocated APD and ARD descriptor handles - ASSERT_EQ(SQL_SUCCESS, SQLFreeHandle(SQL_HANDLE_DESC, apd_descriptor)); + // A null descriptor handle restores the corresponding implicit descriptor. + EXPECT_EQ(SQL_SUCCESS, + SQLSetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, SQL_NULL_HANDLE, 0)); + EXPECT_EQ(SQL_SUCCESS, + SQLSetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, SQL_NULL_HANDLE, 0)); + EXPECT_EQ(SQL_SUCCESS, SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, &value, + sizeof(value), 0)); + EXPECT_EQ(internal_apd, value); + EXPECT_EQ(SQL_SUCCESS, + SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, &value, sizeof(value), 0)); + EXPECT_EQ(internal_ard, value); + // Assign replacement descriptors after the reset. Freeing the old descriptors must + // not revert these replacements, which proves the reset detached the old handles. + SQLHDESC replacement_apd, replacement_ard; + ASSERT_EQ(SQL_SUCCESS, SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &replacement_apd)); + ASSERT_EQ(SQL_SUCCESS, SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &replacement_ard)); + ASSERT_EQ(SQL_SUCCESS, + SQLSetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, replacement_apd, 0)); + ASSERT_EQ(SQL_SUCCESS, + SQLSetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, replacement_ard, 0)); + + ASSERT_EQ(SQL_SUCCESS, SQLFreeHandle(SQL_HANDLE_DESC, apd_descriptor)); ASSERT_EQ(SQL_SUCCESS, SQLFreeHandle(SQL_HANDLE_DESC, ard_descriptor)); + EXPECT_EQ(SQL_SUCCESS, SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, &value, + sizeof(value), 0)); + EXPECT_EQ(replacement_apd, value); + EXPECT_EQ(SQL_SUCCESS, + SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_ROW_DESC, &value, sizeof(value), 0)); + EXPECT_EQ(replacement_ard, value); + + ASSERT_EQ(SQL_SUCCESS, SQLFreeHandle(SQL_HANDLE_DESC, replacement_apd)); + ASSERT_EQ(SQL_SUCCESS, SQLFreeHandle(SQL_HANDLE_DESC, replacement_ard)); - // Verify APD and ARD descriptors has been reverted to implicit descriptors + // Freeing the active explicit descriptors restores the implicit descriptors. value = nullptr; EXPECT_EQ(SQL_SUCCESS, SQLGetStmtAttr(this->stmt, SQL_ATTR_APP_PARAM_DESC, &value, @@ -568,4 +598,21 @@ TYPED_TEST(ConnectionTest, TestSQLSetStmtAttrDescriptor) { EXPECT_EQ(internal_ard, value); } +TYPED_TEST(ConnectionTest, TestExplicitDescriptorsOutliveStatement) { + SQLHSTMT statement; + SQLHDESC apd_descriptor, ard_descriptor; + ASSERT_EQ(SQL_SUCCESS, SQLAllocHandle(SQL_HANDLE_STMT, this->conn, &statement)); + ASSERT_EQ(SQL_SUCCESS, SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &apd_descriptor)); + ASSERT_EQ(SQL_SUCCESS, SQLAllocHandle(SQL_HANDLE_DESC, this->conn, &ard_descriptor)); + + ASSERT_EQ(SQL_SUCCESS, + SQLSetStmtAttr(statement, SQL_ATTR_APP_PARAM_DESC, apd_descriptor, 0)); + ASSERT_EQ(SQL_SUCCESS, + SQLSetStmtAttr(statement, SQL_ATTR_APP_ROW_DESC, ard_descriptor, 0)); + + ASSERT_EQ(SQL_SUCCESS, SQLFreeHandle(SQL_HANDLE_STMT, statement)); + EXPECT_EQ(SQL_SUCCESS, SQLFreeHandle(SQL_HANDLE_DESC, apd_descriptor)); + EXPECT_EQ(SQL_SUCCESS, SQLFreeHandle(SQL_HANDLE_DESC, ard_descriptor)); +} + } // namespace arrow::flight::sql::odbc diff --git a/cpp/src/arrow/flight/sql/odbc/tests/statement_test.cc b/cpp/src/arrow/flight/sql/odbc/tests/statement_test.cc index ba8b883aac47..1f49d4d7dc07 100644 --- a/cpp/src/arrow/flight/sql/odbc/tests/statement_test.cc +++ b/cpp/src/arrow/flight/sql/odbc/tests/statement_test.cc @@ -1981,12 +1981,29 @@ TYPED_TEST(StatementTest, TestSQLMoreResultsNoData) { ASSERT_EQ(SQL_SUCCESS, SQLExecDirect(this->stmt, wsql, wsql_len)); ASSERT_EQ(SQL_NO_DATA, SQLMoreResults(this->stmt)); + + // SQLMoreResults closes the current cursor when there is no next result. + ASSERT_EQ(SQL_ERROR, SQLCloseCursor(this->stmt)); + VerifyOdbcErrorState(SQL_HANDLE_STMT, this->stmt, kErrorState24000); } TYPED_TEST(StatementTest, TestSQLMoreResultsWithoutQuery) { ASSERT_EQ(SQL_NO_DATA, SQLMoreResults(this->stmt)); } +TYPED_TEST(StatementTest, TestSQLFetchWithoutCursor) { + ASSERT_EQ(SQL_ERROR, SQLFetch(this->stmt)); + VerifyOdbcErrorState(SQL_HANDLE_STMT, this->stmt, kErrorState24000); +} + +TYPED_TEST(StatementTest, TestSQLGetDataWithoutCursor) { + SQLINTEGER value; + SQLLEN indicator; + ASSERT_EQ(SQL_ERROR, + SQLGetData(this->stmt, 1, SQL_C_LONG, &value, sizeof(value), &indicator)); + VerifyOdbcErrorState(SQL_HANDLE_STMT, this->stmt, kErrorState24000); +} + TYPED_TEST(StatementTest, TestSQLNativeSqlReturnsInputString) { SQLWCHAR buf[1024]; SQLINTEGER buf_char_len = sizeof(buf) / GetSqlWCharSize();