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
65 changes: 55 additions & 10 deletions driver/statement.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -72,7 +72,38 @@ void Statement::executeQuery(std::unique_ptr<ResultMutator> && mutator) {
*param_set_processed_ptr = 0;

next_param_set_idx = 0;

const auto param_set_array_size =
getEffectiveDescriptor(SQL_ATTR_APP_PARAM_DESC).getAttrAs<SQLULEN>(SQL_DESC_ARRAY_SIZE, 1);

// A set the driver never reaches is SQL_PARAM_UNUSED; every set that is sent overwrites its
// own entry in requestNextPackOfResultSets.
if (auto * param_status_ptr = getEffectiveDescriptor(SQL_ATTR_IMP_PARAM_DESC)
.getAttrAs<SQLUSMALLINT *>(SQL_DESC_ARRAY_STATUS_PTR, 0)) {
for (SQLULEN i = 0; i < param_set_array_size; ++i)
param_status_ptr[i] = SQL_PARAM_UNUSED;
}

requestNextPackOfResultSets(std::move(mutator));

// ODBC executes the statement once per parameter set during SQLExecute; SQLMoreResults then
// walks the result sets those executions produced. A statement that returns nothing produces
// none, so its remaining sets have to be sent here rather than left to SQLMoreResults. Stop
// as soon as a set produces a result set, so a result-returning statement keeps handing them
// back one per SQLMoreResults.
if (!hasResultSet()) {
while (next_param_set_idx < param_set_array_size) {
std::unique_ptr<ResultMutator> carried_mutator;
if (result_reader)
carried_mutator = result_reader->releaseMutator();

requestNextPackOfResultSets(std::move(carried_mutator));

if (hasResultSet())
break;
}
}

is_executed = true;
}

Expand Down Expand Up @@ -122,6 +153,30 @@ void Statement::requestNextPackOfResultSets(std::unique_ptr<ResultMutator> && mu
if (next_param_set_idx >= param_set_array_size)
return;

const auto param_set_idx = next_param_set_idx;
auto & ipd_desc = getEffectiveDescriptor(SQL_ATTR_IMP_PARAM_DESC);
auto * param_status_ptr = ipd_desc.getAttrAs<SQLUSMALLINT *>(SQL_DESC_ARRAY_STATUS_PTR, 0);

// SQL_ATTR_PARAMS_PROCESSED_PTR counts the sets processed including one that fails, so it is
// written before the attempt rather than after it.
// TODO: set this only after this single query is fully fetched (when output parameter support is added)
if (auto * processed_ptr = ipd_desc.getAttrAs<SQLULEN *>(SQL_DESC_ROWS_PROCESSED_PTR, 0))
*processed_ptr = param_set_idx + 1;

try {
sendParamSet(std::move(mutator));
}
catch (...) {
if (param_status_ptr)
param_status_ptr[param_set_idx] = SQL_PARAM_ERROR;
throw;
}

if (param_status_ptr)
param_status_ptr[param_set_idx] = SQL_PARAM_SUCCESS;
}

void Statement::sendParamSet(std::unique_ptr<ResultMutator> && mutator) {
getDiagHeader().setAttr(SQL_DIAG_ROW_COUNT, -1);

auto & connection = getParent();
Expand All @@ -135,11 +190,6 @@ void Statement::requestNextPackOfResultSets(std::unique_ptr<ResultMutator> && mu
uri.addQueryParameter(key, value);
}

// TODO: set this only after this single query is fully fetched (when output parameter support is added)
auto * param_set_processed_ptr = getEffectiveDescriptor(SQL_ATTR_IMP_PARAM_DESC).getAttrAs<SQLULEN *>(SQL_DESC_ROWS_PROCESSED_PTR, 0);
if (param_set_processed_ptr)
*param_set_processed_ptr = next_param_set_idx;

Poco::Net::HTTPRequest request;
request.setMethod(Poco::Net::HTTPRequest::HTTP_POST);
request.setVersion(Poco::Net::HTTPRequest::HTTP_1_1);
Expand Down Expand Up @@ -489,8 +539,6 @@ std::vector<ParamBindingInfo> Statement::getParamsBindingInfo(std::size_t param_
if (fully_bound_param_count > 0)
param_bindings.reserve(fully_bound_param_count);

auto * array_status_ptr = ipd_desc.getAttrAs<SQLUSMALLINT *>(SQL_DESC_ARRAY_STATUS_PTR, 0);

const auto bind_type = apd_desc.getAttrAs<SQLULEN>(SQL_DESC_BIND_TYPE, SQL_PARAM_BIND_TYPE_DEFAULT);
const auto * bind_offset_ptr = apd_desc.getAttrAs<SQLULEN *>(SQL_DESC_BIND_OFFSET_PTR, 0);
const auto bind_offset = (bind_offset_ptr ? *bind_offset_ptr : 0);
Expand Down Expand Up @@ -531,9 +579,6 @@ std::vector<ParamBindingInfo> Statement::getParamsBindingInfo(std::size_t param_
param_bindings.emplace_back(binding_info);
}

if (array_status_ptr)
array_status_ptr[param_set_idx] = SQL_PARAM_SUCCESS; // TODO: elaborate?

return param_bindings;
}

Expand Down
1 change: 1 addition & 0 deletions driver/statement.h
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,7 @@ class Statement

private:
void requestNextPackOfResultSets(std::unique_ptr<ResultMutator> && mutator);
void sendParamSet(std::unique_ptr<ResultMutator> && mutator);

/// Drops the keep-alive HTTP connection unless the previous response body was
/// fully and cleanly consumed.
Expand Down
81 changes: 81 additions & 0 deletions driver/test/statement_parameter_bindings_it.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -562,3 +562,84 @@ TEST_P(StringParameterBindingTest, SpecialCharactersRoundTrip)

STMT_OK(SQLFreeStmt(hstmt, SQL_CLOSE));
}

TEST_F(StatementParameterBindingsTest, IntArrayInsertProcessesEverySet) {
auto create_query = fromUTF8<PTChar>(
"CREATE OR REPLACE TABLE param_array_insert (i Int32) engine MergeTree order by i");
STMT_OK(SQLExecDirect(hstmt, ptcharCast(create_query.data()), SQL_NTS));
STMT_OK(SQLFreeStmt(hstmt, SQL_CLOSE));

SQLINTEGER param[] = { 10, 20, 30, 40, 50 };
SQLLEN param_ind[] = { 0, 0, 0, 0, 0 };
SQLULEN params_processed = 0;

auto insert_query = fromUTF8<PTChar>("INSERT INTO param_array_insert (i) VALUES (?)");
STMT_OK(SQLSetStmtAttr(hstmt, SQL_ATTR_PARAMSET_SIZE, (SQLPOINTER)lengthof(param), 0));
STMT_OK(SQLSetStmtAttr(hstmt, SQL_ATTR_PARAMS_PROCESSED_PTR, &params_processed, 0));
STMT_OK(SQLPrepare(hstmt, ptcharCast(insert_query.data()), SQL_NTS));
STMT_OK(SQLBindParameter(
hstmt,
1,
SQL_PARAM_INPUT,
getCTypeFor<std::decay_t<decltype(param[0])>>(),
SQL_INTEGER,
0,
0,
param,
0,
param_ind));

// An INSERT returns no result sets, so nothing would prompt a caller to call SQLMoreResults:
// this one call has to process the whole array.
STMT_OK(SQLExecute(hstmt));
ASSERT_EQ(params_processed, lengthof(param));
STMT_OK(SQLFreeStmt(hstmt, SQL_CLOSE));

auto select_query = fromUTF8<PTChar>("SELECT i FROM param_array_insert ORDER BY i");
STMT_OK(SQLExecDirect(hstmt, ptcharCast(select_query.data()), SQL_NTS));
for (std::size_t i = 0; i < lengthof(param); ++i) {
SQLINTEGER value = 0;
STMT_OK(SQLFetch(hstmt));
STMT_OK(SQLGetData(hstmt, 1, SQL_C_SLONG, &value, sizeof(value), nullptr));
ASSERT_EQ(value, param[i]) << "row " << i;
}
ASSERT_EQ(SQLFetch(hstmt), SQL_NO_DATA);
STMT_OK(SQLFreeStmt(hstmt, SQL_CLOSE));
}

TEST_F(StatementParameterBindingsTest, IntArrayInsertReportsTheFailingSet) {
// copy_id is derived from id with accurateCast, so an id that does not fit Int16 makes the
// server reject that parameter set and nothing after it is sent.
auto create_query = fromUTF8<PTChar>(
"CREATE OR REPLACE TABLE param_array_error "
"(id Int32, value String, copy_id Int16 DEFAULT accurateCast(id, 'Int16')) "
"engine MergeTree order by id");
STMT_OK(SQLExecDirect(hstmt, ptcharCast(create_query.data()), SQL_NTS));
STMT_OK(SQLFreeStmt(hstmt, SQL_CLOSE));

SQLINTEGER ids[] = { 1, 2, 40000, 4, 5 };
SQLLEN id_ind[] = { 0, 0, 0, 0, 0 };
char values[][2] = { "a", "b", "c", "d", "e" };
SQLLEN value_ind[] = { SQL_NTS, SQL_NTS, SQL_NTS, SQL_NTS, SQL_NTS };
SQLULEN params_processed = 0;
SQLUSMALLINT param_status[lengthof(ids)] = {};

auto insert_query = fromUTF8<PTChar>("INSERT INTO param_array_error (id, value) VALUES (?, ?)");
STMT_OK(SQLSetStmtAttr(hstmt, SQL_ATTR_PARAMSET_SIZE, (SQLPOINTER)lengthof(ids), 0));
STMT_OK(SQLSetStmtAttr(hstmt, SQL_ATTR_PARAMS_PROCESSED_PTR, &params_processed, 0));
STMT_OK(SQLSetStmtAttr(hstmt, SQL_ATTR_PARAM_STATUS_PTR, param_status, 0));
STMT_OK(SQLPrepare(hstmt, ptcharCast(insert_query.data()), SQL_NTS));
STMT_OK(SQLBindParameter(hstmt, 1, SQL_PARAM_INPUT, SQL_C_SLONG, SQL_INTEGER, 0, 0, ids, 0, id_ind));
STMT_OK(SQLBindParameter(hstmt, 2, SQL_PARAM_INPUT, SQL_C_CHAR, SQL_VARCHAR, 255, 0, values, sizeof(values[0]), value_ind));

ASSERT_EQ(SQLExecute(hstmt), SQL_ERROR);

// The count includes the set that failed, and the status array says which one.
ASSERT_EQ(params_processed, 3u);
ASSERT_EQ(param_status[0], SQL_PARAM_SUCCESS);
ASSERT_EQ(param_status[1], SQL_PARAM_SUCCESS);
ASSERT_EQ(param_status[2], SQL_PARAM_ERROR);
ASSERT_EQ(param_status[3], SQL_PARAM_UNUSED);
ASSERT_EQ(param_status[4], SQL_PARAM_UNUSED);
STMT_OK(SQLFreeStmt(hstmt, SQL_CLOSE));
}
Loading