Skip to content
Draft
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
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -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

Expand Down
9 changes: 9 additions & 0 deletions docs/apireference.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
24 changes: 24 additions & 0 deletions docs/query.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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`
Expand All @@ -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:**
Expand Down
5 changes: 5 additions & 0 deletions examples/query-service/basic_example.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down
5 changes: 5 additions & 0 deletions examples/query-service/basic_example_asyncio.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down
14 changes: 13 additions & 1 deletion ydb/_grpc/grpcwrapper/ydb_query.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
),
)


Expand Down
61 changes: 61 additions & 0 deletions ydb/_grpc/grpcwrapper/ydb_query_public_types.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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:
Expand Down
7 changes: 6 additions & 1 deletion ydb/aio/query/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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()
Expand Down
79 changes: 78 additions & 1 deletion ydb/aio/query/pool_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -376,3 +376,80 @@ 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))

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()
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)
27 changes: 23 additions & 4 deletions ydb/aio/query/transaction.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import logging
from typing import (
Callable,
Optional,
TYPE_CHECKING,
)
Expand Down Expand Up @@ -233,17 +234,35 @@ 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

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=on_complete,
)
return self._prev_stream
4 changes: 4 additions & 0 deletions ydb/query/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
"QueryExplainResultFormat",
"QueryOnlineReadOnly",
"QuerySerializableReadWrite",
"QueryStrictSerializableReadWrite",
"VirtualTimestamp",
"QuerySnapshotReadOnly",
"QuerySnapshotReadWrite",
"QueryStaleReadOnly",
Expand Down Expand Up @@ -37,6 +39,8 @@
BaseQueryTxMode,
QueryOnlineReadOnly,
QuerySerializableReadWrite,
QueryStrictSerializableReadWrite,
VirtualTimestamp,
QuerySnapshotReadOnly,
QuerySnapshotReadWrite,
QueryStaleReadOnly,
Expand Down
Loading
Loading