From 4263457531a9088244fd349e56c61d69ea9278b9 Mon Sep 17 00:00:00 2001 From: Oleg Ovcharuk Date: Wed, 30 Sep 2026 16:53:32 +0300 Subject: [PATCH 1/3] Add StrictSerializableRW query commit timestamps --- CHANGELOG.md | 2 + docs/apireference.rst | 9 ++ docs/query.rst | 24 ++++ examples/query-service/basic_example.py | 5 + .../query-service/basic_example_asyncio.py | 5 + ydb/_grpc/grpcwrapper/ydb_query.py | 14 ++- .../grpcwrapper/ydb_query_public_types.py | 61 ++++++++++ ydb/aio/query/base.py | 7 +- ydb/aio/query/pool_test.py | 37 +++++- ydb/aio/query/transaction.py | 21 +++- ydb/query/__init__.py | 4 + ydb/query/base.py | 31 ++++- ydb/query/pool_test.py | 115 +++++++++++++++++- ydb/query/transaction.py | 36 +++++- 14 files changed, 357 insertions(+), 14 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 0803ca6be..a728f2f9b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,5 @@ +* Add StrictSerializableRW Query transactions and optional commit timestamps for successful writes + ## 3.33.1 ## * Add `ydb.convert_floats_to_embedding_bytes` to encode numeric vectors for YDB KNN queries diff --git a/docs/apireference.rst b/docs/apireference.rst index 5dc3144ca..0bc49b066 100644 --- a/docs/apireference.rst +++ b/docs/apireference.rst @@ -247,6 +247,15 @@ Transaction Modes :undoc-members: :exclude-members: name, to_proto +.. autoclass:: ydb.QueryStrictSerializableReadWrite + :members: + :inherited-members: + :undoc-members: + :exclude-members: name, to_proto + +.. autoclass:: ydb.VirtualTimestamp + :members: + .. autoclass:: ydb.QuerySnapshotReadOnly :members: :inherited-members: diff --git a/docs/query.rst b/docs/query.rst index cb7a67437..a9fabf40a 100644 --- a/docs/query.rst +++ b/docs/query.rst @@ -264,6 +264,8 @@ Transaction Modes - Description * - :class:`~ydb.QuerySerializableReadWrite` - Full ACID serializable isolation. Default. Supports reads and writes. + * - :class:`~ydb.QueryStrictSerializableReadWrite` + - Strict serializable read-write mode. A successful write commit may report a virtual timestamp. * - :class:`~ydb.QuerySnapshotReadOnly` - Consistent read-only snapshot taken at transaction start. * - :class:`~ydb.QuerySnapshotReadWrite` @@ -278,6 +280,28 @@ Transaction Modes Manual Transaction Control ^^^^^^^^^^^^^^^^^^^^^^^^^^ +Commit timestamps +~~~~~~~~~~~~~~~~~ + +Choose ``ydb.QueryStrictSerializableReadWrite()`` to receive a commit timestamp +for a successful transaction with write effects. ``tx.commit_timestamp`` is a +``ydb.VirtualTimestamp`` with unsigned 64-bit ``plan_step`` and ``tx_id`` fields, +or ``None`` when the server did not send one. It is available after ``tx.commit()`` +or after fully consuming a ``tx.execute(..., commit_tx=True)`` result stream. +The same property is available on async query transactions. + +.. code-block:: python + + with session.transaction(ydb.QueryStrictSerializableReadWrite()) as tx: + with tx.execute("UPSERT INTO users (id, name) VALUES (1, 'Alice')", commit_tx=True): + pass + timestamp = tx.commit_timestamp + +Virtual timestamps are ordered lexicographically by ``plan_step`` and then +``tx_id``. Comparison requires the same configured driver endpoint and database +path; it raises ``ValueError`` when either is missing or differs. This check +cannot determine whether different endpoint aliases refer to the same database. + Use ``session.transaction()`` when you need fine-grained control: **Synchronous:** diff --git a/examples/query-service/basic_example.py b/examples/query-service/basic_example.py index 854c2dfe7..2e14166a9 100644 --- a/examples/query-service/basic_example.py +++ b/examples/query-service/basic_example.py @@ -51,6 +51,11 @@ def callee(session): tx.commit() + with session.transaction(ydb.QueryStrictSerializableReadWrite()) as tx: + with tx.execute("UPSERT INTO example (key, value) VALUES (3, 'strict')", commit_tx=True): + pass + print(f"StrictSerializableRW commit timestamp: {tx.commit_timestamp}") + print("=" * 50) print("AFTER COMMIT TX") diff --git a/examples/query-service/basic_example_asyncio.py b/examples/query-service/basic_example_asyncio.py index c26db5356..245f0a6b7 100644 --- a/examples/query-service/basic_example_asyncio.py +++ b/examples/query-service/basic_example_asyncio.py @@ -52,6 +52,11 @@ async def callee(session): await tx.commit() + async with session.transaction(ydb.QueryStrictSerializableReadWrite()) as tx: + async with await tx.execute("UPSERT INTO example (key, value) VALUES (3, 'strict')", commit_tx=True): + pass + print(f"StrictSerializableRW commit timestamp: {tx.commit_timestamp}") + print("=" * 50) print("AFTER COMMIT TX") diff --git a/ydb/_grpc/grpcwrapper/ydb_query.py b/ydb/_grpc/grpcwrapper/ydb_query.py index db3c925dd..79f42b457 100644 --- a/ydb/_grpc/grpcwrapper/ydb_query.py +++ b/ydb/_grpc/grpcwrapper/ydb_query.py @@ -102,11 +102,23 @@ def from_proto(msg: ydb_query_pb2.BeginTransactionResponse) -> "BeginTransaction @dataclass class CommitTransactionResponse(IFromProto["ydb_query_pb2.CommitTransactionResponse", "CommitTransactionResponse"]): status: Optional[ServerStatus] + commit_timestamp: Optional[public_types.VirtualTimestamp] = None @staticmethod - def from_proto(msg: ydb_query_pb2.CommitTransactionResponse) -> "CommitTransactionResponse": + def from_proto( + msg: ydb_query_pb2.CommitTransactionResponse, + database: Optional[str] = None, + endpoint: Optional[str] = None, + ) -> "CommitTransactionResponse": return CommitTransactionResponse( status=ServerStatus(msg.status, msg.issues), + commit_timestamp=( + public_types.VirtualTimestamp( + msg.commit_timestamp.plan_step, msg.commit_timestamp.tx_id, database, endpoint + ) + if msg.HasField("commit_timestamp") + else None + ), ) diff --git a/ydb/_grpc/grpcwrapper/ydb_query_public_types.py b/ydb/_grpc/grpcwrapper/ydb_query_public_types.py index 22ad786cc..1ad886c83 100644 --- a/ydb/_grpc/grpcwrapper/ydb_query_public_types.py +++ b/ydb/_grpc/grpcwrapper/ydb_query_public_types.py @@ -1,6 +1,8 @@ import abc import enum import typing +from dataclasses import dataclass, field +from functools import total_ordering from .common_utils import IFromProto, IToProto @@ -65,6 +67,65 @@ def to_proto(self) -> ydb_query_pb2.SerializableModeSettings: return ydb_query_pb2.SerializableModeSettings() +class QueryStrictSerializableReadWrite(BaseQueryTxMode): + """Serializable read-write mode that can report a write commit timestamp.""" + + @property + def name(self) -> str: + return "strict_serializable_read_write" + + def to_proto(self) -> ydb_query_pb2.StrictSerializableRWModeSettings: + return ydb_query_pb2.StrictSerializableRWModeSettings() + + +@total_ordering +@dataclass(frozen=True, eq=False) +class VirtualTimestamp: + """A database-local commit timestamp in unsigned protobuf uint64 coordinates. + + Comparisons require matching configured endpoint and database. Different + endpoints may still address the same database; the SDK cannot verify that. + """ + + plan_step: int + tx_id: int + database: typing.Optional[str] = field(default=None, repr=False) + endpoint: typing.Optional[str] = field(default=None, repr=False) + + def __post_init__(self) -> None: + for value in (self.plan_step, self.tx_id): + if isinstance(value, bool) or not isinstance(value, int) or not 0 <= value < 1 << 64: + raise ValueError("VirtualTimestamp coordinates must be uint64 values") + + def _check_identity(self, other: "VirtualTimestamp") -> None: + if ( + not isinstance(self.database, str) + or not self.database + or not isinstance(self.endpoint, str) + or not self.endpoint + or not isinstance(other.database, str) + or not other.database + or not isinstance(other.endpoint, str) + or not other.endpoint + or (self.endpoint, self.database) != (other.endpoint, other.database) + ): + raise ValueError("Cannot compare timestamps without the same configured database") + + def __eq__(self, other: object) -> bool: + if not isinstance(other, VirtualTimestamp): + return NotImplemented + self._check_identity(other) + return (self.plan_step, self.tx_id) == (other.plan_step, other.tx_id) + + def __lt__(self, other: object) -> bool: + if not isinstance(other, VirtualTimestamp): + return NotImplemented + self._check_identity(other) + return (self.plan_step, self.tx_id) < (other.plan_step, other.tx_id) + + __hash__ = None # type: ignore[assignment] + + class QueryOnlineReadOnly(BaseQueryTxMode): """Each read operation in the transaction is reading the data that is most recent at execution time. The consistency of retrieved data depends on the allow_inconsistent_reads setting: diff --git a/ydb/aio/query/base.py b/ydb/aio/query/base.py index 7aa2f229d..318ee849f 100644 --- a/ydb/aio/query/base.py +++ b/ydb/aio/query/base.py @@ -5,10 +5,11 @@ class AsyncResponseContextIterator(_utilities.AsyncResponseIterator): """Async ExecuteQuery result stream.""" - def __init__(self, it, wrapper, on_error=None, on_finish=None): + def __init__(self, it, wrapper, on_error=None, on_finish=None, on_complete=None): super().__init__(it, wrapper) self._on_error = on_error self._on_finish = on_finish + self._on_complete = on_complete async def __aenter__(self) -> "AsyncResponseContextIterator": return self @@ -26,6 +27,9 @@ async def _next(self): except StopAsyncIteration: # Normal stream termination is not an error and must not invalidate # the session. + if self._on_complete is not None: + self._on_complete() + self._on_complete = None self._call_on_finish() raise except BaseException as e: @@ -45,6 +49,7 @@ def _call_on_finish(self, exception=None): self._on_finish(exception) self._on_finish = None self._on_error = None + self._on_complete = None def __del__(self): self._call_on_finish() diff --git a/ydb/aio/query/pool_test.py b/ydb/aio/query/pool_test.py index 64682994c..3ae00eb0f 100644 --- a/ydb/aio/query/pool_test.py +++ b/ydb/aio/query/pool_test.py @@ -4,7 +4,7 @@ import unittest from unittest.mock import AsyncMock, MagicMock, patch -from ydb import issues +from ydb import _apis, issues, QueryStrictSerializableReadWrite from ydb.aio import _utilities as aio_utilities from ydb.aio.query.pool import QuerySessionPool from ydb.aio.query.session import QuerySession @@ -376,3 +376,38 @@ async def fake_execute_call(**kwargs): await tx.execute("SELECT 1") self.assertIsNone(captured.get("pool_id")) + + +class TestStrictSerializableReadWriteAsync(unittest.IsolatedAsyncioTestCase): + async def test_execute_commit_timestamp_from_trailing_part(self): + driver = MagicMock() + driver._driver_config.endpoint = "localhost:2135" + driver._driver_config.database = "/Root/test" + session = MagicMock() + session._driver_config = driver._driver_config + session.session_id = "session" + session.node_id = None + session._endpoint_key = None + session._settings = None + tx = QueryTxContext(driver, session, QueryStrictSerializableReadWrite()) + + early = _apis.ydb_query.ExecuteQueryResponsePart(status=_apis.StatusIds.SUCCESS) + early.commit_timestamp.plan_step = 1 + trailing = _apis.ydb_query.ExecuteQueryResponsePart(status=_apis.StatusIds.SUCCESS) + trailing.commit_timestamp.plan_step = 2 + trailing.commit_timestamp.tx_id = 3 + + async def responses(): + yield early + yield trailing + + async def execute_call(**kwargs): + return responses() + + with patch.object(type(tx), "_execute_call", side_effect=execute_call): + stream = await tx.execute("UPSERT INTO t (id) VALUES (1)", commit_tx=True) + self.assertIsNone(tx.commit_timestamp) + async for _ in stream: + pass + + self.assertEqual((tx.commit_timestamp.plan_step, tx.commit_timestamp.tx_id), (2, 3)) diff --git a/ydb/aio/query/transaction.py b/ydb/aio/query/transaction.py index 33a00dea4..4ea95993c 100644 --- a/ydb/aio/query/transaction.py +++ b/ydb/aio/query/transaction.py @@ -233,17 +233,30 @@ async def execute( settings=settings, pool_id=pool_id, ) - self._prev_stream = AsyncResponseContextIterator( - it=stream_it, - wrapper=lambda resp: base.wrap_execute_query_response( + timestamp_tracker = base.CommitTimestampTracker(self.session) if commit_tx else None + + def wrap_response(resp): + result = base.wrap_execute_query_response( rpc_state=None, response_pb=resp, session=self.session, tx=self, commit_tx=commit_tx, settings=self.session._settings, - ), + ) + if timestamp_tracker is not None: + timestamp_tracker.observe(resp) + return result + + def finish_commit_timestamp(): + if timestamp_tracker is not None: + self._commit_timestamp = timestamp_tracker.commit_timestamp + + self._prev_stream = AsyncResponseContextIterator( + it=stream_it, + wrapper=wrap_response, on_error=self.session._on_execute_stream_error, on_finish=span_finish_callback(span), + on_complete=finish_commit_timestamp if timestamp_tracker is not None else None, ) return self._prev_stream diff --git a/ydb/query/__init__.py b/ydb/query/__init__.py index 9325709d1..d7f8c5683 100644 --- a/ydb/query/__init__.py +++ b/ydb/query/__init__.py @@ -3,6 +3,8 @@ "QueryExplainResultFormat", "QueryOnlineReadOnly", "QuerySerializableReadWrite", + "QueryStrictSerializableReadWrite", + "VirtualTimestamp", "QuerySnapshotReadOnly", "QuerySnapshotReadWrite", "QueryStaleReadOnly", @@ -37,6 +39,8 @@ BaseQueryTxMode, QueryOnlineReadOnly, QuerySerializableReadWrite, + QueryStrictSerializableReadWrite, + VirtualTimestamp, QuerySnapshotReadOnly, QuerySnapshotReadWrite, QueryStaleReadOnly, diff --git a/ydb/query/base.py b/ydb/query/base.py index 2a592aade..9ab80abbe 100644 --- a/ydb/query/base.py +++ b/ydb/query/base.py @@ -19,6 +19,7 @@ from .._grpc.grpcwrapper.ydb_query_public_types import ( BaseQueryTxMode, ArrowFormatSettings, + VirtualTimestamp, ) from ..connection import _RpcState as RpcState from .. import convert @@ -78,10 +79,11 @@ class QueryResultSetFormat(enum.IntEnum): class SyncResponseContextIterator(_utilities.SyncResponseIterator): """Streams ExecuteQuery results.""" - def __init__(self, it, wrapper, on_error=None, on_finish=None): + def __init__(self, it, wrapper, on_error=None, on_finish=None, on_complete=None): super().__init__(it, wrapper) self._on_error = on_error self._on_finish = on_finish + self._on_complete = on_complete def __enter__(self) -> "SyncResponseContextIterator": return self @@ -99,6 +101,9 @@ def _next(self): except StopIteration: # Normal stream termination is not an error and must not invalidate # the session. + if self._on_complete is not None: + self._on_complete() + self._on_complete = None self._call_on_finish() raise except BaseException as e: @@ -117,6 +122,7 @@ def _call_on_finish(self, exception=None): self._on_finish(exception) self._on_finish = None self._on_error = None + self._on_complete = None def __del__(self): self._call_on_finish() @@ -269,6 +275,29 @@ def wrap_execute_query_response( return None +class CommitTimestampTracker: + """Publish only a timestamp carried by the final, successfully drained part.""" + + def __init__(self, session: "BaseQuerySession"): + self._session = session + self._last_timestamp: Optional[VirtualTimestamp] = None + + def observe(self, response_pb: _apis.ydb_query.ExecuteQueryResponsePart) -> None: + self._last_timestamp = None + if response_pb.HasField("commit_timestamp"): + config = self._session._driver_config + self._last_timestamp = VirtualTimestamp( + response_pb.commit_timestamp.plan_step, + response_pb.commit_timestamp.tx_id, + getattr(config, "database", None), + getattr(config, "endpoint", None), + ) + + @property + def commit_timestamp(self) -> Optional[VirtualTimestamp]: + return self._last_timestamp + + class TxEvent(enum.Enum): BEFORE_COMMIT = "BEFORE_COMMIT" AFTER_COMMIT = "AFTER_COMMIT" diff --git a/ydb/query/pool_test.py b/ydb/query/pool_test.py index 0745b3aab..a372b61b1 100644 --- a/ydb/query/pool_test.py +++ b/ydb/query/pool_test.py @@ -8,12 +8,13 @@ from unittest.mock import patch -from ydb import _utilities, issues +from ydb import _apis, _utilities, issues, QueryStrictSerializableReadWrite, VirtualTimestamp from ydb.convert import _ResultSet, aggregate_result_sets_by_index, aggregate_result_sets_by_index_async from ydb.query.base import create_execute_query_request from ydb.query.pool import QuerySessionPool from ydb.query.session import QuerySession from ydb.query.transaction import QueryTxContext +from ydb.query.transaction import QueryTxStateEnum, wrap_tx_commit_response from ydb._grpc.grpcwrapper import ydb_query_public_types as _ydb_query_public @@ -373,3 +374,115 @@ def fake_execute_call(**kwargs): tx.execute("SELECT 1") self.assertIsNone(captured.get("pool_id")) + + +class TestStrictSerializableReadWrite(unittest.TestCase): + def _make_tx(self): + driver = MagicMock() + driver._driver_config.endpoint = "localhost:2135" + driver._driver_config.database = "/Root/test" + session = MagicMock() + session._driver_config = driver._driver_config + session.session_id = "session" + session.node_id = None + session._endpoint_key = None + session._settings = None + tx = QueryTxContext(driver, session, QueryStrictSerializableReadWrite()) + return tx, session + + def test_mode_uses_field_seven(self): + tx, _ = self._make_tx() + request = create_execute_query_request( + query="UPSERT INTO t (id) VALUES (1)", + session_id="session", + tx_id=None, + commit_tx=True, + tx_mode=tx._tx_state.tx_mode, + syntax=None, + exec_mode=None, + stats_mode=None, + schema_inclusion_mode=None, + result_set_format=None, + arrow_format_settings=None, + parameters=None, + concurrent_result_sets=None, + pool_id=None, + ).to_proto() + settings = request.tx_control.begin_tx + self.assertEqual(settings.WhichOneof("tx_mode"), "strict_serializable_read_write") + self.assertEqual(settings.DESCRIPTOR.fields_by_name["strict_serializable_read_write"].number, 7) + + def test_explicit_commit_timestamp_and_absence(self): + for timestamp in (_apis.ydb_common_pb2.VirtualTimestamp(plan_step=2**64 - 1, tx_id=2**63), None): + with self.subTest(timestamp=timestamp): + tx, session = self._make_tx() + tx._tx_state._change_state(QueryTxStateEnum.BEGINED) + response = _apis.ydb_query.CommitTransactionResponse(status=_apis.StatusIds.SUCCESS) + if timestamp is not None: + response.commit_timestamp.CopyFrom(timestamp) + wrap_tx_commit_response(None, response, session, tx._tx_state, tx) + if timestamp is None: + self.assertIsNone(tx.commit_timestamp) + else: + self.assertEqual(tx.commit_timestamp.plan_step, 2**64 - 1) + self.assertEqual(tx.commit_timestamp.tx_id, 2**63) + self.assertEqual(tx.commit_timestamp.database, "/Root/test") + + def test_failed_explicit_commit_does_not_publish_timestamp(self): + tx, session = self._make_tx() + tx._tx_state._change_state(QueryTxStateEnum.BEGINED) + response = _apis.ydb_query.CommitTransactionResponse(status=_apis.StatusIds.BAD_REQUEST) + response.commit_timestamp.plan_step = 5 + with self.assertRaises(issues.Error): + wrap_tx_commit_response(None, response, session, tx._tx_state, tx) + self.assertIsNone(tx.commit_timestamp) + + def test_execute_timestamp_only_from_trailing_part(self): + tx, _ = self._make_tx() + early = _apis.ydb_query.ExecuteQueryResponsePart(status=_apis.StatusIds.SUCCESS) + early.commit_timestamp.plan_step = 9 + trailing = _apis.ydb_query.ExecuteQueryResponsePart(status=_apis.StatusIds.SUCCESS) + trailing.commit_timestamp.plan_step = 10 + trailing.commit_timestamp.tx_id = 11 + with patch.object(type(tx), "_execute_call", return_value=iter((early, trailing))): + stream = tx.execute("UPSERT INTO t (id) VALUES (1)", commit_tx=True) + self.assertIsNone(tx.commit_timestamp) + list(stream) + self.assertEqual((tx.commit_timestamp.plan_step, tx.commit_timestamp.tx_id), (10, 11)) + + tx, _ = self._make_tx() + with patch.object( + type(tx), "_execute_call", return_value=iter((early, trailing.__class__(status=_apis.StatusIds.SUCCESS))) + ): + list(tx.execute("UPSERT INTO t (id) VALUES (1)", commit_tx=True)) + self.assertIsNone(tx.commit_timestamp) + + def test_interrupted_stream_does_not_publish_timestamp(self): + tx, _ = self._make_tx() + part = _apis.ydb_query.ExecuteQueryResponsePart(status=_apis.StatusIds.SUCCESS) + part.commit_timestamp.plan_step = 4 + + def responses(): + yield part + raise RuntimeError("stream interrupted") + + with patch.object(type(tx), "_execute_call", return_value=responses()): + with self.assertRaises(RuntimeError): + list(tx.execute("UPSERT INTO t (id) VALUES (1)", commit_tx=True)) + self.assertIsNone(tx.commit_timestamp) + + def test_unsigned_order_and_database_scope(self): + def ts(step, tx_id, database="/Root/test", endpoint="localhost:2135"): + return VirtualTimestamp(step, tx_id, database, endpoint) + + self.assertLess(ts(1, 2**64 - 1), ts(2, 0)) + self.assertLess(ts(2, 1), ts(2, 2**63)) + self.assertEqual(ts(2, 1), ts(2, 1)) + for other in (ts(2, 1, database="/Root/other"), ts(2, 1, endpoint="other:2135"), ts(2, 1, endpoint=None)): + with self.assertRaises(ValueError): + _ = ts(2, 1) < other + with self.assertRaises(ValueError): + _ = ts(2, 1) == other + for coordinates in ((-1, 0), (2**64, 0), (0, -1), (0, 2**64)): + with self.assertRaises(ValueError): + ts(*coordinates) diff --git a/ydb/query/transaction.py b/ydb/query/transaction.py index 0008ac71c..46af528d5 100644 --- a/ydb/query/transaction.py +++ b/ydb/query/transaction.py @@ -23,6 +23,7 @@ from ..observability.tracing import SpanName, create_ydb_span, span_finish_callback from .._grpc.grpcwrapper import ydb_topic as _ydb_topic from .._grpc.grpcwrapper import ydb_query as _ydb_query +from .._grpc.grpcwrapper.ydb_query_public_types import VirtualTimestamp from ..connection import _RpcState as RpcState from .._typing import DriverT @@ -192,10 +193,16 @@ def wrap_tx_commit_response( tx_state: QueryTxState, tx: "BaseQueryTxContext", ) -> "BaseQueryTxContext": - message = _ydb_query.CommitTransactionResponse.from_proto(response_pb) + config = session._driver_config + message = _ydb_query.CommitTransactionResponse.from_proto( + response_pb, + database=getattr(config, "database", None), + endpoint=getattr(config, "endpoint", None), + ) if message.status is not None: issues._process_response(message.status) tx_state._change_state(QueryTxStateEnum.COMMITTED) + tx._commit_timestamp = message.commit_timestamp return tx @@ -248,6 +255,7 @@ def __init__(self, driver: DriverT, session: "BaseQuerySession", tx_mode: base.B self._prev_stream = None self._external_error = None self._last_query_stats = None + self._commit_timestamp: Optional[VirtualTimestamp] = None @property def _driver_config(self): @@ -275,6 +283,11 @@ def tx_id(self) -> Optional[str]: def last_query_stats(self): return self._last_query_stats + @property + def commit_timestamp(self) -> Optional[VirtualTimestamp]: + """Commit timestamp, if a StrictSerializableRW write committed successfully.""" + return self._commit_timestamp + def _tx_identity(self) -> _ydb_topic.TransactionIdentity: if not self.tx_id: raise RuntimeError("Unable to get tx identity without started tx.") @@ -681,17 +694,30 @@ def execute( settings=settings, pool_id=pool_id, ) - self._prev_stream = base.SyncResponseContextIterator( - stream_it, - lambda resp: base.wrap_execute_query_response( + timestamp_tracker = base.CommitTimestampTracker(self.session) if commit_tx else None + + def wrap_response(resp): + result = base.wrap_execute_query_response( rpc_state=None, response_pb=resp, session=self.session, tx=self, commit_tx=commit_tx, settings=self.session._settings, - ), + ) + if timestamp_tracker is not None: + timestamp_tracker.observe(resp) + return result + + def finish_commit_timestamp(): + if timestamp_tracker is not None: + self._commit_timestamp = timestamp_tracker.commit_timestamp + + self._prev_stream = base.SyncResponseContextIterator( + stream_it, + wrap_response, on_error=self.session._on_execute_stream_error, on_finish=span_finish_callback(span), + on_complete=finish_commit_timestamp if timestamp_tracker is not None else None, ) return self._prev_stream From 08a290f0533eb1f182466914b218d7730d512a1c Mon Sep 17 00:00:00 2001 From: Oleg Ovcharuk Date: Wed, 30 Sep 2026 18:15:45 +0300 Subject: [PATCH 2/3] Cover strict query timestamp comparison and stream paths --- ydb/aio/query/pool_test.py | 20 ++++++++++++++++++++ ydb/aio/query/transaction.py | 14 ++++++++++---- ydb/query/pool_test.py | 11 +++++++++++ ydb/query/transaction.py | 13 +++++++++---- 4 files changed, 50 insertions(+), 8 deletions(-) diff --git a/ydb/aio/query/pool_test.py b/ydb/aio/query/pool_test.py index 3ae00eb0f..06453c938 100644 --- a/ydb/aio/query/pool_test.py +++ b/ydb/aio/query/pool_test.py @@ -411,3 +411,23 @@ async def execute_call(**kwargs): pass self.assertEqual((tx.commit_timestamp.plan_step, tx.commit_timestamp.tx_id), (2, 3)) + + async def test_execute_without_commit_does_not_publish_timestamp(self): + driver = MagicMock() + session = MagicMock() + session._settings = None + tx = QueryTxContext(driver, session, QueryStrictSerializableReadWrite()) + part = _apis.ydb_query.ExecuteQueryResponsePart(status=_apis.StatusIds.SUCCESS) + part.commit_timestamp.plan_step = 4 + + async def responses(): + yield part + + async def execute_call(**kwargs): + return responses() + + with patch.object(type(tx), "_execute_call", side_effect=execute_call): + stream = await tx.execute("SELECT 1") + async for _ in stream: + pass + self.assertIsNone(tx.commit_timestamp) diff --git a/ydb/aio/query/transaction.py b/ydb/aio/query/transaction.py index 4ea95993c..112c8b9ec 100644 --- a/ydb/aio/query/transaction.py +++ b/ydb/aio/query/transaction.py @@ -1,5 +1,6 @@ import logging from typing import ( + Callable, Optional, TYPE_CHECKING, ) @@ -248,15 +249,20 @@ def wrap_response(resp): timestamp_tracker.observe(resp) return result - def finish_commit_timestamp(): - if timestamp_tracker is not None: - self._commit_timestamp = timestamp_tracker.commit_timestamp + on_complete: Optional[Callable[[], None]] = None + if timestamp_tracker is not None: + tracker = timestamp_tracker + + def finish_commit_timestamp(): + self._commit_timestamp = tracker.commit_timestamp + + on_complete = finish_commit_timestamp self._prev_stream = AsyncResponseContextIterator( it=stream_it, wrapper=wrap_response, on_error=self.session._on_execute_stream_error, on_finish=span_finish_callback(span), - on_complete=finish_commit_timestamp if timestamp_tracker is not None else None, + on_complete=on_complete, ) return self._prev_stream diff --git a/ydb/query/pool_test.py b/ydb/query/pool_test.py index a372b61b1..07be79be0 100644 --- a/ydb/query/pool_test.py +++ b/ydb/query/pool_test.py @@ -471,6 +471,14 @@ def responses(): list(tx.execute("UPSERT INTO t (id) VALUES (1)", commit_tx=True)) self.assertIsNone(tx.commit_timestamp) + def test_execute_without_commit_does_not_publish_timestamp(self): + tx, _ = self._make_tx() + part = _apis.ydb_query.ExecuteQueryResponsePart(status=_apis.StatusIds.SUCCESS) + part.commit_timestamp.plan_step = 4 + with patch.object(type(tx), "_execute_call", return_value=iter((part,))): + list(tx.execute("SELECT 1")) + self.assertIsNone(tx.commit_timestamp) + def test_unsigned_order_and_database_scope(self): def ts(step, tx_id, database="/Root/test", endpoint="localhost:2135"): return VirtualTimestamp(step, tx_id, database, endpoint) @@ -478,6 +486,9 @@ def ts(step, tx_id, database="/Root/test", endpoint="localhost:2135"): self.assertLess(ts(1, 2**64 - 1), ts(2, 0)) self.assertLess(ts(2, 1), ts(2, 2**63)) self.assertEqual(ts(2, 1), ts(2, 1)) + self.assertNotEqual(ts(2, 1), object()) + with self.assertRaises(TypeError): + _ = ts(2, 1) < object() for other in (ts(2, 1, database="/Root/other"), ts(2, 1, endpoint="other:2135"), ts(2, 1, endpoint=None)): with self.assertRaises(ValueError): _ = ts(2, 1) < other diff --git a/ydb/query/transaction.py b/ydb/query/transaction.py index 46af528d5..247c0118e 100644 --- a/ydb/query/transaction.py +++ b/ydb/query/transaction.py @@ -709,15 +709,20 @@ def wrap_response(resp): timestamp_tracker.observe(resp) return result - def finish_commit_timestamp(): - if timestamp_tracker is not None: - self._commit_timestamp = timestamp_tracker.commit_timestamp + on_complete: Optional[Callable[[], None]] = None + if timestamp_tracker is not None: + tracker = timestamp_tracker + + def finish_commit_timestamp(): + self._commit_timestamp = tracker.commit_timestamp + + on_complete = finish_commit_timestamp self._prev_stream = base.SyncResponseContextIterator( stream_it, wrap_response, on_error=self.session._on_execute_stream_error, on_finish=span_finish_callback(span), - on_complete=finish_commit_timestamp if timestamp_tracker is not None else None, + on_complete=on_complete, ) return self._prev_stream From b2e72b7b5edb1fed4f3c639f028fc51d1d714cbb Mon Sep 17 00:00:00 2001 From: Oleg Ovcharuk Date: Thu, 1 Oct 2026 09:39:06 +0300 Subject: [PATCH 3/3] Cover interrupted async query commit stream --- ydb/aio/query/pool_test.py | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/ydb/aio/query/pool_test.py b/ydb/aio/query/pool_test.py index 06453c938..1e243e628 100644 --- a/ydb/aio/query/pool_test.py +++ b/ydb/aio/query/pool_test.py @@ -412,6 +412,28 @@ async def execute_call(**kwargs): self.assertEqual((tx.commit_timestamp.plan_step, tx.commit_timestamp.tx_id), (2, 3)) + async def test_interrupted_stream_does_not_publish_timestamp(self): + driver = MagicMock() + session = MagicMock() + session._settings = None + tx = QueryTxContext(driver, session, QueryStrictSerializableReadWrite()) + part = _apis.ydb_query.ExecuteQueryResponsePart(status=_apis.StatusIds.SUCCESS) + part.commit_timestamp.plan_step = 4 + + async def responses(): + yield part + raise RuntimeError("stream interrupted") + + async def execute_call(**kwargs): + return responses() + + with patch.object(type(tx), "_execute_call", side_effect=execute_call): + stream = await tx.execute("UPSERT INTO t (id) VALUES (1)", commit_tx=True) + with self.assertRaises(RuntimeError): + async for _ in stream: + pass + self.assertIsNone(tx.commit_timestamp) + async def test_execute_without_commit_does_not_publish_timestamp(self): driver = MagicMock() session = MagicMock()