Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 47 additions & 9 deletions cpp/src/arrow/flight/sql/odbc/odbc_api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<ODBCConnection*>(parent);

return ODBCConnection::ExecuteWithDiagnostics(connection, SQL_ERROR, [=]() {
std::shared_ptr<ODBCDescriptor> descriptor = connection->createDescriptor();

if (descriptor) {
*result = reinterpret_cast<SQLHDESC>(descriptor.get());

return SQL_SUCCESS;
}

return SQL_ERROR;
});
}

default:
break;
Expand Down Expand Up @@ -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<ODBCDescriptor*>(handle);

if (!descriptor) {
return SQL_INVALID_HANDLE;
}

descriptor->ReleaseDescriptor();

return SQL_SUCCESS;
}

default:
break;
Expand Down Expand Up @@ -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;

Expand Down Expand Up @@ -277,7 +306,9 @@ SQLRETURN SQLGetDiagField(SQLSMALLINT handleType, SQLHANDLE handle, SQLSMALLINT
}

case SQL_HANDLE_DESC: {
return SQL_ERROR;
ODBCDescriptor* descriptor = reinterpret_cast<ODBCDescriptor*>(handle);
diagnostics = &descriptor->GetDiagnostics();
break;
}

case SQL_HANDLE_STMT: {
Expand Down Expand Up @@ -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<ODBCDescriptor*>(handle);
ODBCConnection* connection = &descriptor->GetConnection();
std::string dsn = connection->GetDSN();
return GetStringAttribute(isUnicode, dsn, true, diagInfoPtr, bufferLength,
stringLengthPtr, *diagnostics);
break;
}

case SQL_HANDLE_STMT: {
Expand Down Expand Up @@ -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;

Expand Down Expand Up @@ -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: {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,11 @@ class ODBCStatement : public ODBCHandle<ODBCStatement> {
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(); }
Expand Down
99 changes: 99 additions & 0 deletions cpp/src/arrow/flight/sql/odbc/tests/connection_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<SQLPOINTER>(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<SQLPOINTER>(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) {
Expand Down
127 changes: 124 additions & 3 deletions cpp/src/arrow/flight/sql/odbc/tests/errors_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ TYPED_TEST(FlightSQLODBCTestBase, TestSQLGetDiagFieldWForConnectFailure) {
static_cast<SQLSMALLINT>(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;
Expand Down Expand Up @@ -170,7 +170,7 @@ TYPED_TEST(FlightSQLODBCTestBase, TestSQLGetDiagFieldWForConnectFailureNTS) {
static_cast<SQLSMALLINT>(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;
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -233,7 +354,7 @@ TYPED_TEST(FlightSQLODBCTestBase, TestSQLGetDiagRecForConnectFailure) {
static_cast<SQLSMALLINT>(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;
Expand Down
Loading