diff --git a/CHANGELOG.md b/CHANGELOG.md index 38c05c665..b933b4ab3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,5 @@ +* Add `pool_id` parameter to `QuerySession.execute`, `QueryTxContext.execute`, and `QuerySessionPool.execute_with_retries` to route queries to a specific resource pool + ## 3.31.1 ## * Support the new `SECRET` scheme entry type: `SchemeEntryType.SECRET` and `SchemeEntry.is_secret()` now recognise secret entries returned by `list_directory`/`describe_path` diff --git a/ydb/_grpc/grpcwrapper/ydb_query.py b/ydb/_grpc/grpcwrapper/ydb_query.py index 812f90c64..db3c925dd 100644 --- a/ydb/_grpc/grpcwrapper/ydb_query.py +++ b/ydb/_grpc/grpcwrapper/ydb_query.py @@ -168,13 +168,14 @@ class ExecuteQueryRequest(IToProto): schema_inclusion_mode: int result_set_format: int arrow_format_settings: Optional[public_types.ArrowFormatSettings] + pool_id: Optional[str] def to_proto(self) -> ydb_query_pb2.ExecuteQueryRequest: tx_control = self.tx_control.to_proto() if self.tx_control is not None else self.tx_control arrow_format_settings = ( self.arrow_format_settings.to_proto() if self.arrow_format_settings is not None else None ) - return ydb_query_pb2.ExecuteQueryRequest( + req = ydb_query_pb2.ExecuteQueryRequest( session_id=self.session_id, tx_control=tx_control, query_content=self.query_content.to_proto(), @@ -186,3 +187,6 @@ def to_proto(self) -> ydb_query_pb2.ExecuteQueryRequest: concurrent_result_sets=self.concurrent_result_sets, parameters=convert.query_parameters_to_pb(self.parameters), ) + if self.pool_id is not None: + req.pool_id = self.pool_id + return req diff --git a/ydb/aio/query/pool.py b/ydb/aio/query/pool.py index 7a71c5f73..25f92703f 100644 --- a/ydb/aio/query/pool.py +++ b/ydb/aio/query/pool.py @@ -219,6 +219,7 @@ async def execute_with_retries( parameters: Optional[dict] = None, retry_settings: Optional[RetrySettings] = None, *args, + pool_id: Optional[str] = None, **kwargs, ) -> List[convert.ResultSet]: """Special interface to execute a one-shot queries in a safe, retriable way. @@ -228,6 +229,7 @@ async def execute_with_retries( :param query: A query, yql or sql text. :param parameters: dict with parameters and YDB types; :param retry_settings: RetrySettings object. + :param pool_id: Optional resource pool ID for routing the query to a specific resource pool. :return: Result sets or exception in case of execution errors. """ @@ -236,7 +238,7 @@ async def execute_with_retries( async def wrapped_callee(): async with self.checkout(timeout=retry_settings.max_session_acquire_timeout) as session: - it = await session.execute(query, parameters, *args, **kwargs) + it = await session.execute(query, parameters, *args, pool_id=pool_id, **kwargs) return await convert.aggregate_result_sets_by_index_async(it) return await retry_operation_async(wrapped_callee, retry_settings) diff --git a/ydb/aio/query/pool_test.py b/ydb/aio/query/pool_test.py index 9b2e1dce4..71d083da1 100644 --- a/ydb/aio/query/pool_test.py +++ b/ydb/aio/query/pool_test.py @@ -2,12 +2,14 @@ import asyncio import unittest -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch from ydb import issues from ydb.aio.query.pool import QuerySessionPool from ydb.aio.query.session import QuerySession +from ydb.aio.query.transaction import QueryTxContext from ydb.observability.metrics import QuerySessionPoolMetrics +from ydb._grpc.grpcwrapper import ydb_query_public_types as _ydb_query_public def _make_pool(size=1): @@ -117,3 +119,131 @@ async def test_retry_reacquires_invalidated_session_before_first_use(self): self.assertEqual(result, "ok") live_session.explain.assert_awaited_once_with("SELECT 1") + + +async def _async_empty_iter(): + """Async-iterable that yields nothing; usable as a stub for session.execute return value.""" + if False: + yield + + +class TestQuerySessionExecutePoolId(unittest.IsolatedAsyncioTestCase): + """Test that pool_id flows from async session.execute() → _execute_call() → driver.""" + + def _make_session(self): + driver = MagicMock() + driver._driver_config.query_client_settings = None + session = QuerySession(driver) + session._session_id = "fake-session-id" + return session + + async def test_execute_passes_pool_id_to_execute_call(self): + session = self._make_session() + + captured = {} + + async def fake_execute_call(**kwargs): + captured.update(kwargs) + return _async_empty_iter() + + with patch.object(type(session), "_execute_call", side_effect=fake_execute_call): + await session.execute("SELECT 1", pool_id="my-pool") + + self.assertEqual(captured.get("pool_id"), "my-pool") + + async def test_execute_without_pool_id_passes_none(self): + session = self._make_session() + + captured = {} + + async def fake_execute_call(**kwargs): + captured.update(kwargs) + return _async_empty_iter() + + with patch.object(type(session), "_execute_call", side_effect=fake_execute_call): + await session.execute("SELECT 1") + + self.assertIsNone(captured.get("pool_id")) + + +class TestPoolIdParameter(unittest.IsolatedAsyncioTestCase): + async def test_execute_with_retries_passes_pool_id_to_session(self): + pool = _make_pool(size=1) + + session = MagicMock() + session.is_active = True + session.execute = AsyncMock(return_value=_async_empty_iter()) + + async def mock_acquire(timeout=None): + return session + + pool.acquire = mock_acquire + pool.release = AsyncMock() + + await pool.execute_with_retries("SELECT 1", pool_id="my-pool") + + session.execute.assert_awaited_once() + call_kwargs = session.execute.call_args[1] + self.assertEqual(call_kwargs.get("pool_id"), "my-pool") + + async def test_execute_with_retries_without_pool_id(self): + pool = _make_pool(size=1) + + session = MagicMock() + session.is_active = True + session.execute = AsyncMock(return_value=_async_empty_iter()) + + async def mock_acquire(timeout=None): + return session + + pool.acquire = mock_acquire + pool.release = AsyncMock() + + await pool.execute_with_retries("SELECT 1") + + session.execute.assert_awaited_once() + call_kwargs = session.execute.call_args[1] + self.assertIsNone(call_kwargs.get("pool_id")) + + +class TestQueryTxContextExecutePoolId(unittest.IsolatedAsyncioTestCase): + """Test that pool_id flows from async QueryTxContext.execute() → _execute_call().""" + + def _make_tx(self): + driver = MagicMock() + driver._driver_config.query_client_settings = None + session = MagicMock() + session.session_id = "fake-session-id" + session.node_id = None + session._endpoint_key = None + tx_mode = _ydb_query_public.QuerySerializableReadWrite() + tx = QueryTxContext(driver, session, tx_mode) + return tx + + async def test_execute_passes_pool_id_to_execute_call(self): + tx = self._make_tx() + + captured = {} + + async def fake_execute_call(**kwargs): + captured.update(kwargs) + return _async_empty_iter() + + with patch.object(type(tx), "_execute_call", side_effect=fake_execute_call): + await tx.execute("SELECT 1", pool_id="my-pool") + + self.assertEqual(captured.get("pool_id"), "my-pool") + + async def test_execute_without_pool_id_passes_none(self): + tx = self._make_tx() + + captured = {} + + async def fake_execute_call(**kwargs): + captured.update(kwargs) + return _async_empty_iter() + + with patch.object(type(tx), "_execute_call", side_effect=fake_execute_call): + await tx.execute("SELECT 1") + + self.assertIsNone(captured.get("pool_id")) diff --git a/ydb/aio/query/session.py b/ydb/aio/query/session.py index 8d9657cad..7f7654580 100644 --- a/ydb/aio/query/session.py +++ b/ydb/aio/query/session.py @@ -136,6 +136,7 @@ async def execute( schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode] = None, result_set_format: Optional[base.QueryResultSetFormat] = None, arrow_format_settings: Optional[base.ArrowFormatSettings] = None, + pool_id: Optional[str] = None, ) -> AsyncResponseContextIterator: """Sends a query to Query Service @@ -157,6 +158,7 @@ async def execute( 1) QueryResultSetFormat.VALUE, which is default; 2) QueryResultSetFormat.ARROW. :param arrow_format_settings: Settings for Arrow format when result_set_format is ARROW. + :param pool_id: Optional resource pool ID for routing the query to a specific resource pool. :return: Iterator with result sets """ @@ -182,6 +184,7 @@ async def execute( arrow_format_settings=arrow_format_settings, concurrent_result_sets=concurrent_result_sets, settings=settings, + pool_id=pool_id, ) return AsyncResponseContextIterator( it=stream_it, diff --git a/ydb/aio/query/transaction.py b/ydb/aio/query/transaction.py index 4592f5892..33a00dea4 100644 --- a/ydb/aio/query/transaction.py +++ b/ydb/aio/query/transaction.py @@ -177,6 +177,7 @@ async def execute( schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode] = None, result_set_format: Optional[base.QueryResultSetFormat] = None, arrow_format_settings: Optional[base.ArrowFormatSettings] = None, + pool_id: Optional[str] = None, ) -> AsyncResponseContextIterator: """Sends a query to Query Service @@ -204,6 +205,7 @@ async def execute( 1) QueryResultSetFormat.VALUE, which is default; 2) QueryResultSetFormat.ARROW. :param arrow_format_settings: Settings for Arrow format when result_set_format is ARROW. + :param pool_id: Optional resource pool ID for routing the query to a specific compute pool. :return: Iterator with result sets """ @@ -229,6 +231,7 @@ async def execute( arrow_format_settings=arrow_format_settings, concurrent_result_sets=concurrent_result_sets, settings=settings, + pool_id=pool_id, ) self._prev_stream = AsyncResponseContextIterator( it=stream_it, diff --git a/ydb/query/base.py b/ydb/query/base.py index 093c7d554..12752fe23 100644 --- a/ydb/query/base.py +++ b/ydb/query/base.py @@ -165,6 +165,7 @@ def create_execute_query_request( arrow_format_settings: Optional[ArrowFormatSettings], parameters: Optional[dict], concurrent_result_sets: Optional[bool], + pool_id: Optional[str], ) -> ydb_query.ExecuteQueryRequest: try: syntax = QuerySyntax.YQL_V1 if not syntax else syntax @@ -207,6 +208,7 @@ def create_execute_query_request( schema_inclusion_mode=schema_inclusion_mode, result_set_format=result_set_format, arrow_format_settings=arrow_format_settings, + pool_id=pool_id, ) except BaseException as e: raise issues.ClientInternalError("Unable to prepare execute request") from e diff --git a/ydb/query/pool.py b/ydb/query/pool.py index 869884102..79f5051a7 100644 --- a/ydb/query/pool.py +++ b/ydb/query/pool.py @@ -243,6 +243,7 @@ def execute_with_retries( parameters: Optional[dict] = None, retry_settings: Optional[RetrySettings] = None, *args, + pool_id: Optional[str] = None, **kwargs, ) -> List[convert.ResultSet]: """Special interface to execute a one-shot queries in a safe, retriable way. @@ -252,6 +253,7 @@ def execute_with_retries( :param query: A query, yql or sql text. :param parameters: dict with parameters and YDB types; :param retry_settings: RetrySettings object. + :param pool_id: Optional resource pool ID for routing the query to a specific compute pool. :return: Result sets or exception in case of execution errors. """ @@ -263,7 +265,7 @@ def execute_with_retries( def wrapped_callee(): with self.checkout(timeout=retry_settings.max_session_acquire_timeout) as session: - it = session.execute(query, parameters, *args, **kwargs) + it = session.execute(query, parameters, *args, pool_id=pool_id, **kwargs) return convert.aggregate_result_sets_by_index(it) return retry_operation_sync(wrapped_callee, retry_settings) @@ -274,6 +276,7 @@ def execute_with_retries_async( parameters: Optional[dict] = None, retry_settings: Optional[RetrySettings] = None, *args, + pool_id: Optional[str] = None, **kwargs, ) -> futures.Future: """Asynchronously execute a query with retries.""" @@ -287,6 +290,7 @@ def execute_with_retries_async( parameters, retry_settings, *args, + pool_id=pool_id, **kwargs, ) diff --git a/ydb/query/pool_test.py b/ydb/query/pool_test.py index 73653a731..33041ccf4 100644 --- a/ydb/query/pool_test.py +++ b/ydb/query/pool_test.py @@ -6,10 +6,15 @@ import unittest from unittest.mock import MagicMock +from unittest.mock import patch + from ydb import issues 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._grpc.grpcwrapper import ydb_query_public_types as _ydb_query_public def _make_pool(size=1): @@ -149,3 +154,169 @@ def test_retry_reacquires_invalidated_session_before_first_use(self): self.assertEqual(result, "ok") live_session.explain.assert_called_once_with("SELECT 1") + + +class TestQuerySessionExecutePoolId(unittest.TestCase): + """Test that pool_id flows from session.execute() → _execute_call() → driver.""" + + def _make_session(self): + driver = MagicMock() + driver._driver_config.query_client_settings = None + session = QuerySession(driver) + session._session_id = "fake-session-id" + return session + + def test_execute_passes_pool_id_to_driver(self): + session = self._make_session() + + captured = {} + + def fake_execute_call(**kwargs): + captured.update(kwargs) + return iter([]) + + with patch.object(type(session), "_execute_call", side_effect=fake_execute_call): + session.execute("SELECT 1", pool_id="my-pool") + + self.assertEqual(captured.get("pool_id"), "my-pool") + + def test_execute_without_pool_id_passes_none(self): + session = self._make_session() + + captured = {} + + def fake_execute_call(**kwargs): + captured.update(kwargs) + return iter([]) + + with patch.object(type(session), "_execute_call", side_effect=fake_execute_call): + session.execute("SELECT 1") + + self.assertIsNone(captured.get("pool_id")) + + +class TestCreateExecuteQueryRequest(unittest.TestCase): + def test_pool_id_is_set_in_request(self): + req = create_execute_query_request( + query="SELECT 1", + session_id="sess-1", + tx_id=None, + commit_tx=None, + tx_mode=None, + 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="my-pool", + ) + self.assertEqual(req.pool_id, "my-pool") + proto = req.to_proto() + self.assertEqual(proto.pool_id, "my-pool") + + def test_pool_id_defaults_to_none_and_is_absent_from_proto(self): + req = create_execute_query_request( + query="SELECT 1", + session_id="sess-1", + tx_id=None, + commit_tx=None, + tx_mode=None, + 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, + ) + self.assertIsNone(req.pool_id) + proto = req.to_proto() + self.assertEqual(proto.pool_id, "") + + +class TestPoolIdParameter(unittest.TestCase): + def test_execute_with_retries_passes_pool_id_to_session(self): + pool = _make_pool(size=1) + + session = MagicMock() + session.is_active = True + session.execute = MagicMock(return_value=[]) + + def mock_acquire(timeout=None): + return session + + pool.acquire = mock_acquire + pool.release = MagicMock() + + pool.execute_with_retries("SELECT 1", pool_id="my-pool") + + session.execute.assert_called_once() + call_kwargs = session.execute.call_args[1] + self.assertEqual(call_kwargs.get("pool_id"), "my-pool") + + def test_execute_with_retries_without_pool_id(self): + pool = _make_pool(size=1) + + session = MagicMock() + session.is_active = True + session.execute = MagicMock(return_value=[]) + + def mock_acquire(timeout=None): + return session + + pool.acquire = mock_acquire + pool.release = MagicMock() + + pool.execute_with_retries("SELECT 1") + + session.execute.assert_called_once() + call_kwargs = session.execute.call_args[1] + self.assertIsNone(call_kwargs.get("pool_id")) + + +class TestQueryTxContextExecutePoolId(unittest.TestCase): + """Test that pool_id flows from QueryTxContext.execute() → _execute_call().""" + + def _make_tx(self): + driver = MagicMock() + driver._driver_config.query_client_settings = None + session = MagicMock() + session.session_id = "fake-session-id" + session.node_id = None + session._endpoint_key = None + tx_mode = _ydb_query_public.QuerySerializableReadWrite() + tx = QueryTxContext(driver, session, tx_mode) + return tx + + def test_execute_passes_pool_id_to_execute_call(self): + tx = self._make_tx() + + captured = {} + + def fake_execute_call(**kwargs): + captured.update(kwargs) + return iter([]) + + with patch.object(type(tx), "_execute_call", side_effect=fake_execute_call): + tx.execute("SELECT 1", pool_id="my-pool") + + self.assertEqual(captured.get("pool_id"), "my-pool") + + def test_execute_without_pool_id_passes_none(self): + tx = self._make_tx() + + captured = {} + + def fake_execute_call(**kwargs): + captured.update(kwargs) + return iter([]) + + with patch.object(type(tx), "_execute_call", side_effect=fake_execute_call): + tx.execute("SELECT 1") + + self.assertIsNone(captured.get("pool_id")) diff --git a/ydb/query/session.py b/ydb/query/session.py index 7e2e8dafe..b8a7d82e1 100644 --- a/ydb/query/session.py +++ b/ydb/query/session.py @@ -307,6 +307,7 @@ def _execute_call( arrow_format_settings: Optional[base.ArrowFormatSettings] = None, concurrent_result_sets: bool = False, settings: Optional[BaseRequestSettings] = None, + pool_id: Optional[str] = None, ) -> Iterable[_apis.ydb_query.ExecuteQueryResponsePart]: ... @overload @@ -323,6 +324,7 @@ def _execute_call( arrow_format_settings: Optional[base.ArrowFormatSettings] = None, concurrent_result_sets: bool = False, settings: Optional[BaseRequestSettings] = None, + pool_id: Optional[str] = None, ) -> Awaitable[Iterable[_apis.ydb_query.ExecuteQueryResponsePart]]: ... def _execute_call( @@ -338,6 +340,7 @@ def _execute_call( arrow_format_settings: Optional[base.ArrowFormatSettings] = None, concurrent_result_sets: bool = False, settings: Optional[BaseRequestSettings] = None, + pool_id: Optional[str] = None, ) -> Union[ Iterable[_apis.ydb_query.ExecuteQueryResponsePart], Awaitable[Iterable[_apis.ydb_query.ExecuteQueryResponsePart]], @@ -361,6 +364,7 @@ def _execute_call( result_set_format=result_set_format, arrow_format_settings=arrow_format_settings, concurrent_result_sets=concurrent_result_sets, + pool_id=pool_id, ) return self._driver( @@ -484,6 +488,7 @@ def execute( schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode] = None, result_set_format: Optional[base.QueryResultSetFormat] = None, arrow_format_settings: Optional[base.ArrowFormatSettings] = None, + pool_id: Optional[str] = None, ) -> base.SyncResponseContextIterator: """Sends a query to Query Service @@ -505,6 +510,7 @@ def execute( 1) QueryResultSetFormat.VALUE, which is default; 2) QueryResultSetFormat.ARROW. :param arrow_format_settings: Settings for Arrow format when result_set_format is ARROW. + :param pool_id: Optional resource pool ID for routing the query to a specific compute pool. :return: Iterator with result sets """ @@ -530,6 +536,7 @@ def execute( arrow_format_settings=arrow_format_settings, concurrent_result_sets=concurrent_result_sets, settings=settings, + pool_id=pool_id, ) return base.SyncResponseContextIterator( stream_it, diff --git a/ydb/query/transaction.py b/ydb/query/transaction.py index e2cbb34ba..692b2a4c4 100644 --- a/ydb/query/transaction.py +++ b/ydb/query/transaction.py @@ -384,6 +384,7 @@ def _execute_call( arrow_format_settings: Optional[base.ArrowFormatSettings], concurrent_result_sets: Optional[bool], settings: Optional[BaseRequestSettings], + pool_id: Optional[str], ) -> Iterable[_apis.ydb_query.ExecuteQueryResponsePart]: ... @overload @@ -400,6 +401,7 @@ def _execute_call( arrow_format_settings: Optional[base.ArrowFormatSettings], concurrent_result_sets: Optional[bool], settings: Optional[BaseRequestSettings], + pool_id: Optional[str], ) -> Awaitable[Iterable[_apis.ydb_query.ExecuteQueryResponsePart]]: ... def _execute_call( @@ -415,6 +417,7 @@ def _execute_call( arrow_format_settings: Optional[base.ArrowFormatSettings], concurrent_result_sets: Optional[bool], settings: Optional[BaseRequestSettings], + pool_id: Optional[str], ) -> Union[ Iterable[_apis.ydb_query.ExecuteQueryResponsePart], Awaitable[Iterable[_apis.ydb_query.ExecuteQueryResponsePart]], @@ -441,6 +444,7 @@ def _execute_call( result_set_format=result_set_format, arrow_format_settings=arrow_format_settings, concurrent_result_sets=concurrent_result_sets, + pool_id=pool_id, ) return self._driver( @@ -616,6 +620,7 @@ def execute( schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode] = None, result_set_format: Optional[base.QueryResultSetFormat] = None, arrow_format_settings: Optional[base.ArrowFormatSettings] = None, + pool_id: Optional[str] = None, ) -> base.SyncResponseContextIterator: """Sends a query to Query Service @@ -644,6 +649,7 @@ def execute( 1) QueryResultSetFormat.VALUE, which is default; 2) QueryResultSetFormat.ARROW. :param arrow_format_settings: Settings for Arrow format when result_set_format is ARROW. + :param pool_id: Optional resource pool ID for routing the query to a specific compute pool. :return: Iterator with result sets """ @@ -669,6 +675,7 @@ def execute( parameters=parameters, concurrent_result_sets=concurrent_result_sets, settings=settings, + pool_id=pool_id, ) self._prev_stream = base.SyncResponseContextIterator( stream_it,