Skip to content
Open
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
28 changes: 28 additions & 0 deletions python/adbc_driver_flightsql/tests/test_errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,34 @@ def _cancel():
cur.fetchone()


def test_read_partition_cancel(test_dbapi):
with test_dbapi.cursor() as cur:
partitions, _ = cur.adbc_execute_partitions("forever")
assert len(partitions) == 1

errors = []

def _read():
try:
cur.adbc_read_partition(partitions[0])
except Exception as e:
errors.append(e)

t = threading.Thread(target=_read, daemon=True)
t.start()
time.sleep(2)
cur.adbc_cancel()
t.join(timeout=30)
assert not t.is_alive(), "cancel did not interrupt ReadPartition"

assert len(errors) == 1
assert isinstance(errors[0], test_dbapi.OperationalError)
assert str(errors[0]) == (
"CANCELLED: [FlightSQL] context canceled"
" (Canceled; ReadPartition(DoGet)). Vendor code: 1"
)


def test_query_error_fetch(test_dbapi):
with test_dbapi.cursor() as cur:
cur.execute("error_do_get")
Expand Down
2 changes: 1 addition & 1 deletion python/adbc_driver_manager/adbc_driver_manager/_lib.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -1540,7 +1540,7 @@ cdef class AdbcStatement(_AdbcHandle):
check_error(status, &c_error)

def cancel(self) -> None:
"""Attempt to cancel any ongoing operations on the connection."""
"""Attempt to cancel any ongoing operations on the statement."""
cdef CAdbcError c_error = empty_error()
cdef CAdbcStatusCode status
with nogil:
Expand Down
39 changes: 26 additions & 13 deletions python/adbc_driver_manager/adbc_driver_manager/dbapi.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,18 @@
import typing
import warnings
import weakref
from typing import Any, Dict, List, Literal, Mapping, NoReturn, Optional, Tuple, Union
from typing import (
Any,
Callable,
Dict,
List,
Literal,
Mapping,
NoReturn,
Optional,
Tuple,
Union,
)

try:
import pyarrow
Expand Down Expand Up @@ -787,6 +798,7 @@ def __init__(

self._last_query: Optional[Union[str, bytes]] = None
self._results: Optional["_RowIterator"] = None
self._cancel: Callable[[], None] = self._stmt.cancel
self._arraysize = 1
self._rowcount = -1
self._bind_by_name = False
Expand All @@ -798,6 +810,7 @@ def _clear(self) -> None:
if self._results is not None:
self._results.close()
self._results = None
self._cancel = self._stmt.cancel

@property
def arraysize(self) -> int:
Expand Down Expand Up @@ -926,7 +939,7 @@ def execute(self, operation: Union[bytes, str], parameters=None) -> "Self":
handle, self._rowcount = _blocking_call(
self._stmt.execute_query, (), {}, self._stmt.cancel
)
self._results = _RowIterator(self._stmt, handle, self._backend)
self._results = _RowIterator(self._cancel, handle, self._backend)
return self

def executemany(self, operation: Union[bytes, str], seq_of_parameters) -> None:
Expand Down Expand Up @@ -1065,13 +1078,13 @@ def __next__(self) -> tuple:

def adbc_cancel(self) -> None:
"""
Cancel any ongoing operations on this statement.
Cancel any ongoing operations on this cursor.

Notes
-----
This is an extension and not part of the DBAPI standard.
"""
self._stmt.cancel()
self._cancel()

def adbc_ingest(
self,
Expand Down Expand Up @@ -1290,12 +1303,12 @@ def adbc_read_partition(self, partition: bytes) -> None:
"""
_requires_pyarrow()
self._clear()
self._results = None
self._cancel = self._conn._conn.cancel
handle = _blocking_call(
self._conn._conn.read_partition, (partition,), {}, self._stmt.cancel
self._conn._conn.read_partition, (partition,), {}, self._cancel
)
self._rowcount = -1
self._results = _RowIterator(self._stmt, handle, self._backend)
self._results = _RowIterator(self._cancel, handle, self._backend)

@property
def adbc_statement(self) -> _lib.AdbcStatement:
Expand Down Expand Up @@ -1440,11 +1453,11 @@ class _RowIterator(_Closeable):

def __init__(
self,
stmt: _lib.AdbcStatement,
cancel: Callable[[], None],
handle: _lib.ArrowArrayStreamHandle,
dbapi_backend: _dbapi_backend.DbapiBackend,
) -> None:
self._stmt = stmt
self._cancel = cancel
self._handle: Optional[_lib.ArrowArrayStreamHandle] = handle
self._backend = dbapi_backend
self._reader: Optional["AdbcRecordBatchReader"] = None
Expand Down Expand Up @@ -1495,7 +1508,7 @@ def fetchone(self) -> Optional[tuple]:
try:
while True:
self._current_batch = _blocking_call(
self.reader.read_next_batch, (), {}, self._stmt.cancel
self.reader.read_next_batch, (), {}, self._cancel
)
if self._current_batch.num_rows > 0:
break
Expand Down Expand Up @@ -1531,10 +1544,10 @@ def fetchall(self) -> List[tuple]:
return rows

def fetch_arrow_table(self) -> "pyarrow.Table":
return _blocking_call(self.reader.read_all, (), {}, self._stmt.cancel)
return _blocking_call(self.reader.read_all, (), {}, self._cancel)

def fetch_df(self) -> "pandas.DataFrame":
return _blocking_call(self.reader.read_pandas, (), {}, self._stmt.cancel)
return _blocking_call(self.reader.read_pandas, (), {}, self._cancel)

def fetch_polars(self) -> "polars.DataFrame":
import polars
Expand All @@ -1545,7 +1558,7 @@ def fetch_polars(self) -> "polars.DataFrame":
),
(),
{},
self._stmt.cancel,
self._cancel,
)

def fetch_arrow(self) -> _lib.ArrowArrayStreamHandle:
Expand Down
Loading