Skip to content

Commit 9d678e8

Browse files
SNOW-3264599: Python - implement cursor reset method
1 parent d1c0313 commit 9d678e8

4 files changed

Lines changed: 525 additions & 98 deletions

File tree

Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,79 @@
1+
from __future__ import annotations
2+
3+
from collections.abc import Generator
4+
from contextlib import contextmanager
5+
from typing import TYPE_CHECKING
6+
7+
from .protobuf_gen.database_driver_v1_pb2 import (
8+
ExecuteResult,
9+
StatementHandle,
10+
StatementNewRequest,
11+
StatementReleaseRequest,
12+
StatementSetSqlQueryRequest,
13+
)
14+
15+
16+
if TYPE_CHECKING:
17+
from ..connection import Connection
18+
19+
20+
@contextmanager
21+
def create_statement(connection: Connection, query: str) -> Generator[StatementHandle]:
22+
statement_request = StatementNewRequest(conn_handle=connection.conn_handle)
23+
statement = connection.db_api.statement_new(request=statement_request)
24+
stmt_handle = statement.stmt_handle
25+
sql_query_request = StatementSetSqlQueryRequest(stmt_handle=stmt_handle, query=query)
26+
try:
27+
connection.db_api.statement_set_sql_query(sql_query_request)
28+
yield stmt_handle
29+
finally:
30+
release_request = StatementReleaseRequest(stmt_handle=stmt_handle)
31+
connection.db_api.statement_release(release_request)
32+
33+
34+
def extract_rowcount(result: ExecuteResult) -> int:
35+
"""Extract rowcount from execute result."""
36+
if result and result.HasField("rows_affected"):
37+
return result.rows_affected
38+
return -1
39+
40+
41+
def extract_sqlstate(result: ExecuteResult | None) -> str | None:
42+
"""Extract sqlstate from execute result.
43+
44+
"00000" (successful completion) is treated as None for backwards compatibility with the old connector.
45+
"""
46+
sql_state = result.sql_state if result else None
47+
if sql_state and sql_state != "00000":
48+
return sql_state
49+
return None
50+
51+
52+
def get_stream_ptr(result: ExecuteResult | None) -> int:
53+
"""Get the ArrowArrayStream pointer from execute result.
54+
55+
Returns:
56+
int: The ArrowArrayStream pointer as an integer
57+
58+
Raises:
59+
RuntimeError: If result is invalid or stream pointer is null
60+
"""
61+
if result is None:
62+
raise RuntimeError("No query has been executed")
63+
64+
if not hasattr(result, "stream") or result.stream is None:
65+
raise RuntimeError("Execute result does not contain a valid stream")
66+
67+
if not hasattr(result.stream, "value") or result.stream.value is None:
68+
raise RuntimeError("Stream does not contain a valid pointer value")
69+
70+
stream_value = result.stream.value
71+
if len(stream_value) != 8:
72+
raise RuntimeError(f"Stream pointer value has wrong length: {len(stream_value)} (expected 8)")
73+
74+
stream_ptr = int.from_bytes(stream_value, byteorder="little", signed=False)
75+
76+
if stream_ptr == 0:
77+
raise RuntimeError("Stream pointer is null")
78+
79+
return stream_ptr

python/src/snowflake/connector/cursor.py

Lines changed: 78 additions & 91 deletions
Original file line numberDiff line numberDiff line change
@@ -35,10 +35,9 @@
3535
ExecuteResult,
3636
QueryBindings,
3737
StatementExecuteQueryRequest,
38-
StatementNewRequest,
39-
StatementReleaseRequest,
40-
StatementSetSqlQueryRequest,
38+
StatementHandle,
4139
)
40+
from ._internal.query_utils import create_statement, extract_rowcount, extract_sqlstate, get_stream_ptr
4241
from ._internal.type_codes import get_type_code
4342
from .errors import InterfaceError, NotSupportedError, ProgrammingError
4443

@@ -93,6 +92,13 @@ def from_column(cls, col: Any) -> ResultMetadata:
9392
is_nullable=col.nullable,
9493
)
9594

95+
@classmethod
96+
def create_description(cls, result: ExecuteResult | None) -> list[ResultMetadata] | None:
97+
"""Extract description from execute result column metadata."""
98+
if result and result.columns:
99+
return [cls.from_column(col) for col in result.columns]
100+
return None
101+
96102

97103
# Backward compatibility alias
98104
ResultMetadataV2 = ResultMetadata
@@ -174,7 +180,7 @@ def __init__(self, connection: Connection) -> None:
174180
"""
175181
self._connection = connection
176182
self._description: list[ResultMetadata] | None = None
177-
self._rowcount: int | None = None
183+
self._rowcount: int = -1
178184
self._arraysize: int = 1
179185
self._sqlstate: str | None = None
180186
self._closed = False
@@ -304,9 +310,20 @@ def callproc(self, procname: str, parameters: Sequence[Any] | None = None) -> Se
304310
raise NotSupportedError("callproc is not implemented")
305311

306312
@pep249
307-
def close(self) -> None:
308-
"""Close the cursor now (rather than whenever __del__ is called)."""
309-
self._closed = True
313+
def close(self) -> bool | None:
314+
"""Close the cursor now (rather than whenever __del__ is called).
315+
316+
Returns whether the cursor was closed during this call.
317+
"""
318+
try:
319+
if self._closed:
320+
return False
321+
self.reset(closing=True)
322+
self._closed = True
323+
del self._messages[:]
324+
return True
325+
except Exception:
326+
return None
310327

311328
def _build_query_bindings(self, parameters: Sequence[Any]) -> QueryBindings | None:
312329
"""Serialize parameters and build a QueryBindings protobuf message.
@@ -397,6 +414,7 @@ def execute(
397414
) -> SnowflakeCursorBase:
398415
"""
399416
Execute a database operation (query or command).
417+
Resets the cursor state before the execution.
400418
401419
Args:
402420
operation (str): SQL statement to execute
@@ -405,52 +423,43 @@ def execute(
405423
For pyformat paramstyle: sequence (%s) or dict (%(name)s)
406424
For format paramstyle: sequence (%s)
407425
"""
426+
self.reset()
427+
return self._execute(operation, parameters, _is_put_get, **kwargs)
428+
429+
def _execute(
430+
self,
431+
operation: str,
432+
parameters: Sequence[Any] | dict[str, Any] | None = None,
433+
_is_put_get: bool | None = None,
434+
**kwargs: Any,
435+
) -> SnowflakeCursorBase:
436+
"""Execute query logic."""
408437
query, bindings = self._prepare_query(operation, parameters)
409-
stmt_handle = self._connection.db_api.statement_new(
410-
StatementNewRequest(conn_handle=self._connection.conn_handle)
411-
).stmt_handle
412-
self._connection.db_api.statement_set_sql_query(
413-
StatementSetSqlQueryRequest(stmt_handle=stmt_handle, query=query)
414-
)
415438

416-
request = StatementExecuteQueryRequest(stmt_handle=stmt_handle, bindings=bindings)
439+
result: ExecuteResult | None = None
440+
with create_statement(self.connection, query) as stmt_handle:
441+
result = self._execute_query(stmt_handle, bindings)
417442

443+
# populate description, rowcount, and sqlstate
444+
self._description = ResultMetadata.create_description(result)
445+
self._rowcount = extract_rowcount(result)
446+
self._sqlstate = extract_sqlstate(result)
447+
# save execute result
448+
self.execute_result = result
449+
450+
return self
451+
452+
def _execute_query(self, stmt_handle: StatementHandle, bindings: QueryBindings | None) -> ExecuteResult:
418453
try:
419-
self.execute_result = self._connection.db_api.statement_execute_query(request).result
454+
request = StatementExecuteQueryRequest(stmt_handle=stmt_handle, bindings=bindings)
455+
return self._connection.db_api.statement_execute_query(request).result
420456
except ProgrammingError as exc:
421457
self._sqlstate = exc.sqlstate or None
422458
raise
423-
finally:
424-
self._connection.db_api.statement_release(StatementReleaseRequest(stmt_handle=stmt_handle))
425-
426-
# Reset streaming state for a new result
427-
self._binding_data = None
428-
self._iterator = None
429-
self._fetch_mode = None
430-
self._rownumber = -1
431-
432-
# Populate description, rowcount, and sqlstate
433-
self._populate_description()
434-
self._populate_rowcount()
435-
self._populate_sqlstate()
436-
return self
437-
438-
def _populate_rowcount(self) -> None:
439-
if self.execute_result and self.execute_result.HasField("rows_affected"):
440-
self._rowcount = self.execute_result.rows_affected
441-
else:
442-
self._rowcount = None
443-
444-
def _populate_sqlstate(self) -> None:
445-
# "00000" (successful completion) is treated as None for
446-
# backwards compatibility with the old connector.
447-
sql_state = self.execute_result.sql_state if self.execute_result else None
448-
if sql_state and sql_state != "00000":
449-
self._sqlstate = sql_state
450-
else:
451-
self._sqlstate = None
452459

453460
@pep249
461+
@_requires_not_closed
462+
@_requires_open_connection
454463
def executemany(self, operation: str, seq_of_parameters: Sequence[Sequence[Any] | dict[str, Any]]) -> None:
455464
"""
456465
Execute a database operation repeatedly for each element in seq_of_parameters.
@@ -476,10 +485,11 @@ def executemany(self, operation: str, seq_of_parameters: Sequence[Sequence[Any]
476485
# - Client-side binding (pyformat/format)
477486
# - Dict parameters (server-side doesn't support named binding)
478487
if paramstyle.is_client_side() or isinstance(first_params, dict):
488+
self.reset()
479489
total_rowcount = 0
480490
unknown_rowcount = False
481491
for params in seq_of_parameters:
482-
self.execute(operation, params)
492+
self._execute(operation, params) # no reset between calls
483493
rc = self._rowcount
484494
if rc is None or rc == -1:
485495
unknown_rowcount = True
@@ -515,50 +525,8 @@ def executemany(self, operation: str, seq_of_parameters: Sequence[Sequence[Any]
515525
# Arrow stream helpers
516526
# ------------------------------------------------------------------
517527

518-
def _get_stream_ptr(self) -> int:
519-
"""Get the ArrowArrayStream pointer from execute result.
520-
521-
Returns:
522-
int: The ArrowArrayStream pointer as an integer
523-
524-
Raises:
525-
RuntimeError: If execute_result is invalid or stream pointer is null
526-
"""
527-
if self.execute_result is None:
528-
raise RuntimeError("No query has been executed")
529-
530-
if not hasattr(self.execute_result, "stream") or self.execute_result.stream is None:
531-
raise RuntimeError("Execute result does not contain a valid stream")
532-
533-
if not hasattr(self.execute_result.stream, "value") or self.execute_result.stream.value is None:
534-
raise RuntimeError("Stream does not contain a valid pointer value")
535-
536-
stream_value = self.execute_result.stream.value
537-
if len(stream_value) != 8:
538-
raise RuntimeError(f"Stream pointer value has wrong length: {len(stream_value)} (expected 8)")
539-
540-
stream_ptr = int.from_bytes(stream_value, byteorder="little", signed=False)
541-
542-
if stream_ptr == 0:
543-
raise RuntimeError("Stream pointer is null")
544-
545-
return stream_ptr
546-
547-
def _populate_description(self) -> None:
548-
"""Populate cursor description from execute result column metadata."""
549-
if self.execute_result is None:
550-
self._description = None
551-
return
552-
553-
columns = self.execute_result.columns
554-
if not columns:
555-
self._description = None
556-
return
557-
558-
self._description = [ResultMetadata.from_column(col) for col in columns]
559-
560528
def _get_iterator(self) -> ArrowStreamIterator:
561-
stream_ptr = self._get_stream_ptr()
529+
stream_ptr = get_stream_ptr(self.execute_result)
562530
arrow_context = ArrowConverterContext()
563531
return ArrowStreamIterator(
564532
stream_ptr,
@@ -572,7 +540,7 @@ def _get_table_iterator(
572540
self,
573541
force_microsecond_precision: bool = False,
574542
) -> ArrowStreamTableIterator:
575-
stream_ptr = self._get_stream_ptr()
543+
stream_ptr = get_stream_ptr(self.execute_result)
576544
arrow_context = ArrowConverterContext()
577545
return ArrowStreamTableIterator(
578546
stream_ptr,
@@ -849,8 +817,27 @@ def scroll(self, value: int, mode: str = "relative") -> None:
849817
raise NotSupportedError("scroll is not supported")
850818

851819
def reset(self, closing: bool = False) -> None:
852-
"""Reset the result set."""
853-
raise NotImplementedError("reset is not yet implemented")
820+
"""Reset the result set.
821+
822+
Clears all result-related state, preparing the cursor for a new execution.
823+
This is called automatically by execute() before each query, and by close().
824+
825+
Args:
826+
closing: If True, do not reset rowcount,
827+
see: SNOW-647539: Do not erase the rowcount information when closing the cursor.
828+
If False, reset rowcount to None.
829+
"""
830+
self._description = None
831+
if not closing:
832+
self._rowcount = -1
833+
self._sqlstate = None
834+
835+
self.execute_result = None
836+
self._iterator = None
837+
self._fetch_mode = None
838+
839+
self._binding_data = None
840+
self._rownumber = -1
854841

855842
def query_result(self, qid: str) -> SnowflakeCursorBase:
856843
"""Query the result of a previously executed query."""

python/tests/integ/test_cursor.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -427,16 +427,16 @@ def test_description_numeric_precision_and_scale(self, cursor):
427427
class TestCursorRowcount:
428428
"""Integration tests for Cursor.rowcount property."""
429429

430-
def test_rowcount_is_none_before_execute(self, connection):
431-
"""Test that rowcount returns None before any query is executed."""
430+
def test_rowcount_is_minus_one_before_execute(self, connection):
431+
"""Test that rowcount returns -1 before any query is executed."""
432432
# Given a new cursor
433433
cursor = connection.cursor()
434434

435435
# When accessing rowcount before execute
436436
result = cursor.rowcount
437437

438-
# Then it should be None (per PEP 249)
439-
assert result is None
438+
# Then it should be -1 (per PEP 249)
439+
assert result == -1
440440

441441
def test_rowcount_after_select_single_row(self, cursor):
442442
"""Test rowcount after a SELECT query returning single row."""

0 commit comments

Comments
 (0)