diff --git a/examples/opentelemetry/load_tank.py b/examples/opentelemetry/load_tank.py index 801ddadf5..b07715326 100644 --- a/examples/opentelemetry/load_tank.py +++ b/examples/opentelemetry/load_tank.py @@ -7,7 +7,7 @@ import random import time from dataclasses import dataclass -from typing import AsyncIterator, Tuple +from typing import AsyncIterator from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import OTLPMetricExporter from opentelemetry.sdk.metrics import MeterProvider @@ -84,7 +84,7 @@ def _load_config() -> LoadConfig: ) -async def _load_steps(config: LoadConfig) -> AsyncIterator[Tuple[int, str, int]]: +async def _load_steps(config: LoadConfig) -> AsyncIterator[tuple[int, str, int]]: pattern = ( (config.peak_rps, "Peak", config.peak_duration), (config.medium_rps, "Medium down", config.medium_duration), diff --git a/examples/reservations-bot-demo/cloud_function/controller.py b/examples/reservations-bot-demo/cloud_function/controller.py index 474721091..ff45363c4 100644 --- a/examples/reservations-bot-demo/cloud_function/controller.py +++ b/examples/reservations-bot-demo/cloud_function/controller.py @@ -1,4 +1,3 @@ -import typing import logging from storage import Storage from models import ( @@ -16,7 +15,7 @@ class Controller(object): def __init__(self, storage: Storage): self._storage = storage - def _find_available_table_id(self, request: ReservationCreateRequest) -> typing.Optional[int]: + def _find_available_table_id(self, request: ReservationCreateRequest) -> int | None: table_ids = set(self._storage.list_table_ids(cnt=request.cnt)) reserved_table_ids = set(self._storage.find_reserved_table_ids(cnt=request.cnt, dt=request.dt)) for table_id in table_ids.difference(reserved_table_ids): diff --git a/examples/reservations-bot-demo/cloud_function/models.py b/examples/reservations-bot-demo/cloud_function/models.py index bce3622a8..fa34fbd91 100644 --- a/examples/reservations-bot-demo/cloud_function/models.py +++ b/examples/reservations-bot-demo/cloud_function/models.py @@ -1,35 +1,34 @@ -import typing import datetime from pydantic import BaseModel class Reservation(BaseModel): - phone: typing.Optional[typing.Union[str, str]] = None - description: typing.Optional[typing.Union[bytes, str]] = None + phone: str | str | None = None + description: bytes | str | None = None table_id: int dt: datetime.datetime class Table(BaseModel): table_id: int = None - description: typing.Optional[typing.Union[bytes, str]] = None + description: bytes | str | None = None cnt: int class ReservationCreateRequest(BaseModel): dt: datetime.datetime cnt: int - description: typing.Optional[typing.Union[bytes, str]] = None - phone: typing.Optional[typing.Union[bytes, str]] + description: bytes | str | None = None + phone: bytes | str | None class ReservationCreateResponse(BaseModel): success: bool - table_id: typing.Optional[int] = None + table_id: int | None = None class ReservationCancelRequest(BaseModel): - phone: typing.Optional[typing.Union[bytes, str]] + phone: bytes | str | None dt: datetime.datetime diff --git a/examples/reservations-bot-demo/cloud_function/storage.py b/examples/reservations-bot-demo/cloud_function/storage.py index 2342db647..c1bda3979 100644 --- a/examples/reservations-bot-demo/cloud_function/storage.py +++ b/examples/reservations-bot-demo/cloud_function/storage.py @@ -1,5 +1,4 @@ import datetime -import typing import ydb from utils import session_pool_context, make_driver_config from config import Config @@ -10,7 +9,7 @@ def __init__(self, *, endpoint: str, database: str, path: str): self._database = database self._driver_config = make_driver_config(endpoint, database, path) - def list_table_ids(self, *, cnt: int = 0) -> typing.List[int]: + def list_table_ids(self, *, cnt: int = 0) -> list[int]: query = f"""PRAGMA TablePathPrefix("{self._database}"); DECLARE $cnt as Uint64; SELECT table_id FROM tables WHERE cnt >= $cnt; @@ -26,7 +25,7 @@ def transaction(session): tables = session_pool.retry_operation_sync(transaction) return list(map(lambda x: getattr(x, "table_id"), tables)) - def find_reserved_table_ids(self, *, cnt: int, dt: datetime.datetime) -> typing.List[int]: + def find_reserved_table_ids(self, *, cnt: int, dt: datetime.datetime) -> list[int]: query = f"""PRAGMA TablePathPrefix("{self._database}"); DECLARE $dt AS DateTime; DECLARE $reservation_period_minutes AS Int32; diff --git a/examples/time-series-serverless/database.py b/examples/time-series-serverless/database.py index 8e8d91815..cca3201dd 100644 --- a/examples/time-series-serverless/database.py +++ b/examples/time-series-serverless/database.py @@ -1,5 +1,4 @@ import ydb -from typing import List from config import ydb_configuration from exception import ConnectionFailure @@ -31,7 +30,7 @@ def create_driver(self) -> ydb.Driver: def table_client(self) -> ydb.TableClient: return self.driver.table_client - def bulk_upsert(self, rows: List, column_types: ydb.BulkUpsertColumns): + def bulk_upsert(self, rows: list, column_types: ydb.BulkUpsertColumns): self.table_client.bulk_upsert(self.config.full_path, rows, column_types) diff --git a/examples/time-series-serverless/time_series.py b/examples/time-series-serverless/time_series.py index b95ccbbe2..e2a77f668 100644 --- a/examples/time-series-serverless/time_series.py +++ b/examples/time-series-serverless/time_series.py @@ -1,4 +1,3 @@ -from typing import Dict import random import ydb @@ -23,7 +22,7 @@ def generate_time_series(parameters: Parameters): ydb_client.bulk_upsert(rows, column_types) -def do_handle(event: Dict, _) -> Response: +def do_handle(event: dict, _) -> Response: if "queryStringParameters" not in event: return BadRequest("Incorrect function call: non HTTP request") diff --git a/examples/topic/writer_async_example.py b/examples/topic/writer_async_example.py index 74b867d2c..b019bd154 100644 --- a/examples/topic/writer_async_example.py +++ b/examples/topic/writer_async_example.py @@ -1,6 +1,5 @@ import asyncio import datetime -from typing import Dict, List import ydb from ydb import TopicWriterMessage @@ -95,8 +94,8 @@ async def send_messages_and_wait_all_commit_with_results( await writer.flush() -async def switch_messages_with_many_producers(writers: Dict[str, ydb.TopicWriterAsyncIO], messages: List[str]): - futures = [] # type: List[asyncio.Future] +async def switch_messages_with_many_producers(writers: dict[str, ydb.TopicWriterAsyncIO], messages: list[str]): + futures = [] # type: list[asyncio.Future] for msg in messages: # select writer for the msg diff --git a/examples/topic/writer_example.py b/examples/topic/writer_example.py index 8bcacb7dd..d54bde329 100644 --- a/examples/topic/writer_example.py +++ b/examples/topic/writer_example.py @@ -1,6 +1,5 @@ import concurrent.futures import datetime -from typing import Dict, List from concurrent.futures import Future, wait # noqa: F401 import ydb @@ -123,7 +122,7 @@ def send_messages_and_wait_all_commit_with_flush(writer: ydb.TopicWriter): def send_messages_and_wait_all_commit_with_results(writer: ydb.TopicWriter): - futures = [] # type: List[concurrent.futures.Future] + futures = [] # type: list[concurrent.futures.Future] for i in range(10): future = writer.async_write_with_ack() futures.append(future) @@ -134,8 +133,8 @@ def send_messages_and_wait_all_commit_with_results(writer: ydb.TopicWriter): raise future.exception() -def switch_messages_with_many_producers(writers: Dict[str, ydb.TopicWriter], messages: List[str]): - futures = [] # type: List[Future] +def switch_messages_with_many_producers(writers: dict[str, ydb.TopicWriter], messages: list[str]): + futures = [] # type: list[Future] for msg in messages: # select writer for the msg diff --git a/generate_protoc.py b/generate_protoc.py index d69260285..f707d940c 100644 --- a/generate_protoc.py +++ b/generate_protoc.py @@ -2,13 +2,12 @@ import pathlib import shutil -from typing import List from argparse import ArgumentParser from grpc_tools import command -def files_filter(dir, items: List[str]) -> List[str]: +def files_filter(dir, items: list[str]) -> list[str]: ignored_names = ['.git'] ignore = [] diff --git a/tests/aio/query/test_query_session_pool.py b/tests/aio/query/test_query_session_pool.py index 7af2f8962..2d61c1b2e 100644 --- a/tests/aio/query/test_query_session_pool.py +++ b/tests/aio/query/test_query_session_pool.py @@ -4,7 +4,6 @@ import pytest import ydb -from typing import Optional from ydb import QueryExplainResultFormat from ydb.aio.query.pool import QuerySessionPool @@ -106,7 +105,7 @@ async def callee(session: QuerySession): ], ) @pytest.mark.asyncio - async def test_retry_tx_normal(self, pool: QuerySessionPool, tx_mode: Optional[ydb.BaseQueryTxMode]): + async def test_retry_tx_normal(self, pool: QuerySessionPool, tx_mode: ydb.BaseQueryTxMode | None): retry_no = 0 async def callee(tx: QueryTxContext): diff --git a/tests/observability/test_observability_enable.py b/tests/observability/test_observability_enable.py index eb251db4c..b38ec2dc5 100644 --- a/tests/observability/test_observability_enable.py +++ b/tests/observability/test_observability_enable.py @@ -11,7 +11,7 @@ * ``get_trace_metadata`` is empty until a provider is enabled. """ -from typing import Any, Dict, List, Optional, Tuple +from typing import Any import pytest @@ -44,12 +44,12 @@ class RecordingSpan: """Minimal Span implementation that records every call.""" - def __init__(self, name: str, attributes: Optional[dict], kind: Optional[str], sink: List[Dict[str, Any]]): + def __init__(self, name: str, attributes: dict | None, kind: str | None, sink: list[dict[str, Any]]): self.name = name - self.attributes: Dict[str, Any] = dict(attributes or {}) + self.attributes: dict[str, Any] = dict(attributes or {}) self.kind = kind self.ended = False - self.errors: List[BaseException] = [] + self.errors: list[BaseException] = [] self._sink = sink def set_error(self, exception): @@ -91,9 +91,9 @@ def __exit__(self_inner, exc_type, exc_val, exc_tb): class RecordingProvider: """Custom TracingProvider used across the tests.""" - def __init__(self, metadata: Optional[List[Tuple[str, str]]] = None): - self.spans: List[RecordingSpan] = [] - self.finished: List[Dict[str, Any]] = [] + def __init__(self, metadata: list[tuple[str, str]] | None = None): + self.spans: list[RecordingSpan] = [] + self.finished: list[dict[str, Any]] = [] self._metadata = metadata or [] def create_span(self, name, attributes=None, kind=None) -> Span: @@ -423,7 +423,7 @@ class TestSetPeerAttributes: """Direct unit tests for ``set_peer_attributes``.""" def _recording_span(self): - recorded: Dict[str, Any] = {} + recorded: dict[str, Any] = {} class _Span: def set_attribute(self, key, value): @@ -460,7 +460,7 @@ class TestSpanFinishCallback: """``span_finish_callback`` wires stream completion into span lifecycle.""" def test_finish_ends_span_on_success(self): - calls: List[str] = [] + calls: list[str] = [] class _Span: def set_error(self, exc): @@ -473,7 +473,7 @@ def end(self): assert calls == ["end"] def test_finish_records_error_then_ends_span(self): - calls: List[str] = [] + calls: list[str] = [] exc = RuntimeError("stream broke") class _Span: @@ -563,7 +563,7 @@ def test_set_attribute_forwards_to_underlying_otel_span(self): from ydb.opentelemetry.plugin import TracingSpan - recorded: Dict[str, Any] = {} + recorded: dict[str, Any] = {} class _FakeOtelSpan: def set_attribute(self, key, value): diff --git a/tests/query/test_query_session_pool.py b/tests/query/test_query_session_pool.py index 499febe6b..3c98897fd 100644 --- a/tests/query/test_query_session_pool.py +++ b/tests/query/test_query_session_pool.py @@ -5,7 +5,6 @@ import time from concurrent import futures -from typing import Optional from ydb import QueryExplainResultFormat from ydb.query.pool import QuerySessionPool @@ -102,7 +101,7 @@ def callee(session: QuerySession): (ydb.QueryStaleReadOnly()), ], ) - def test_retry_tx_normal(self, pool: QuerySessionPool, tx_mode: Optional[ydb.BaseQueryTxMode]): + def test_retry_tx_normal(self, pool: QuerySessionPool, tx_mode: ydb.BaseQueryTxMode | None): retry_no = 0 def callee(tx: QueryTxContext): diff --git a/tests/slo/src/core/metrics.py b/tests/slo/src/core/metrics.py index 313e1683d..edbdba29d 100644 --- a/tests/slo/src/core/metrics.py +++ b/tests/slo/src/core/metrics.py @@ -8,7 +8,7 @@ from contextlib import contextmanager from importlib.metadata import version from os import environ -from typing import Any, Optional, Tuple +from typing import Any OP_TYPE_READ, OP_TYPE_WRITE = "read", "write" OP_STATUS_SUCCESS, OP_STATUS_FAILURE = "success", "error" @@ -19,7 +19,7 @@ logger = logging.getLogger(__name__) -def _normalize_labels(labels: Any) -> Tuple[Any, ...]: +def _normalize_labels(labels: Any) -> tuple[Any, ...]: if labels is None: return tuple() if isinstance(labels, str): @@ -44,7 +44,7 @@ def stop( labels, start_time: float, attempts: int = 1, - error: Optional[Exception] = None, + error: Exception | None = None, ) -> None: pass @@ -95,7 +95,7 @@ def stop( labels, start_time: float, attempts: int = 1, - error: Optional[Exception] = None, + error: Exception | None = None, ) -> None: return None @@ -228,7 +228,7 @@ def stop( labels, start_time: float, attempts: int = 1, - error: Optional[Exception] = None, + error: Exception | None = None, ) -> None: labels_t = _normalize_labels(labels) duration = time.time() - start_time @@ -294,7 +294,7 @@ def inc_duplicated(self, n: int = 1) -> None: self._topic_duplicated.add(int(n), attributes={"ref": REF}) -def _resolve_metrics_endpoint(cli_endpoint: Optional[str]) -> str: +def _resolve_metrics_endpoint(cli_endpoint: str | None) -> str: """ Resolution order: 1. OTEL_EXPORTER_OTLP_METRICS_ENDPOINT (used as-is) @@ -315,7 +315,7 @@ def _resolve_metrics_endpoint(cli_endpoint: Optional[str]) -> str: return (cli_endpoint or "").strip() -def create_metrics(otlp_endpoint: Optional[str]) -> BaseMetrics: +def create_metrics(otlp_endpoint: str | None) -> BaseMetrics: """ Build a metrics exporter. diff --git a/tests/slo/src/root_runner.py b/tests/slo/src/root_runner.py index 2c8c2e75a..8bac4b69b 100644 --- a/tests/slo/src/root_runner.py +++ b/tests/slo/src/root_runner.py @@ -2,7 +2,6 @@ import ydb import ydb.aio import logging -from typing import Dict from core.metrics import WORKLOAD from runners.topic_runner import TopicRunner @@ -20,7 +19,7 @@ def _is_async_workload(args) -> bool: class SLORunner: def __init__(self): - self.runners: Dict[str, type(BaseRunner)] = {} + self.runners: dict[str, type(BaseRunner)] = {} def register_runner(self, prefix: str, runner_cls: type(BaseRunner)): self.runners[prefix] = runner_cls diff --git a/tests/slo/src/runners/base.py b/tests/slo/src/runners/base.py index 53b728eb0..bf114b4fb 100644 --- a/tests/slo/src/runners/base.py +++ b/tests/slo/src/runners/base.py @@ -1,6 +1,5 @@ import logging from abc import ABC, abstractmethod -from typing import Optional import ydb @@ -8,7 +7,7 @@ class BaseRunner(ABC): def __init__(self): self.logger = logging.getLogger(self.__class__.__module__) - self.driver: Optional[ydb.Driver] = None + self.driver: ydb.Driver | None = None @property @abstractmethod diff --git a/ydb/_apis.py b/ydb/_apis.py index f4b6dd537..f357aac92 100644 --- a/ydb/_apis.py +++ b/ydb/_apis.py @@ -63,23 +63,23 @@ ydb_coordination = ydb_coordination_pb2 -class CmsService(object): +class CmsService: Stub = ydb_cms_v1_pb2_grpc.CmsServiceStub -class DiscoveryService(object): +class DiscoveryService: Stub = ydb_discovery_v1_pb2_grpc.DiscoveryServiceStub ListEndpoints = "ListEndpoints" -class OperationService(object): +class OperationService: Stub = ydb_operation_v1_pb2_grpc.OperationServiceStub ForgetOperation = "ForgetOperation" GetOperation = "GetOperation" CancelOperation = "CancelOperation" -class SchemeService(object): +class SchemeService: Stub = ydb_scheme_v1_pb2_grpc.SchemeServiceStub MakeDirectory = "MakeDirectory" RemoveDirectory = "RemoveDirectory" @@ -88,7 +88,7 @@ class SchemeService(object): ModifyPermissions = "ModifyPermissions" -class TableService(object): +class TableService: Stub = ydb_table_v1_pb2_grpc.TableServiceStub StreamExecuteScanQuery = "StreamExecuteScanQuery" @@ -113,7 +113,7 @@ class TableService(object): BulkUpsert = "BulkUpsert" -class TopicService(object): +class TopicService: Stub = ydb_topic_v1_pb2_grpc.TopicServiceStub CreateTopic = "CreateTopic" @@ -127,7 +127,7 @@ class TopicService(object): CommitOffset = "CommitOffset" -class QueryService(object): +class QueryService: Stub = ydb_query_v1_pb2_grpc.QueryServiceStub CreateSession = "CreateSession" @@ -143,7 +143,7 @@ class QueryService(object): FetchScriptResults = "FetchScriptResults" -class CoordinationService(object): +class CoordinationService: Stub = ydb_coordination_v1_pb2_grpc.CoordinationServiceStub CreateNode = "CreateNode" AlterNode = "AlterNode" diff --git a/ydb/_errors.py b/ydb/_errors.py index 1e969c09c..6527364e1 100644 --- a/ydb/_errors.py +++ b/ydb/_errors.py @@ -1,5 +1,4 @@ from dataclasses import dataclass -from typing import Optional from . import issues @@ -57,4 +56,4 @@ def check_retriable_error(err, retry_settings, attempt): @dataclass class ErrorRetryInfo: is_retriable: bool - sleep_timeout_seconds: Optional[float] + sleep_timeout_seconds: float | None diff --git a/ydb/_session_impl.py b/ydb/_session_impl.py index 9d5f77703..ad84cf42b 100644 --- a/ydb/_session_impl.py +++ b/ydb/_session_impl.py @@ -73,7 +73,7 @@ def explain_data_query_request_factory(session_state, yql_text): return request -class _ExplainResponse(object): +class _ExplainResponse: def __init__(self, ast, plan): self.query_ast = ast self.query_plan = plan @@ -399,7 +399,7 @@ def wrap_read_table_response(response): return convert.ResultSet.from_message(response.result.result_set, snapshot=snapshot) -class SessionState(object): +class SessionState: def __init__(self, table_client_settings): self._session_id = None self._query_cache = _utilities.LRUCache(1000) diff --git a/ydb/_sp_impl.py b/ydb/_sp_impl.py index dfea89f9c..897318334 100644 --- a/ydb/_sp_impl.py +++ b/ydb/_sp_impl.py @@ -7,7 +7,7 @@ from . import settings, issues, _utilities, tracing -class SessionPoolImpl(object): +class SessionPoolImpl: def __init__( self, logger, diff --git a/ydb/_tx_ctx_impl.py b/ydb/_tx_ctx_impl.py index 3c7eec39e..89ed6a9b9 100644 --- a/ydb/_tx_ctx_impl.py +++ b/ydb/_tx_ctx_impl.py @@ -84,7 +84,7 @@ def commit_request_factory(session_state, tx_state): return request -class TxState(object): +class TxState: __slots__ = ("tx_id", "tx_mode", "dead", "initialized") def __init__(self, tx_mode): diff --git a/ydb/_typing.py b/ydb/_typing.py index 00530d57b..7c491a2d2 100644 --- a/ydb/_typing.py +++ b/ydb/_typing.py @@ -9,7 +9,6 @@ Any, Callable, Iterable, - Tuple, TypeVar, Union, TYPE_CHECKING, @@ -67,4 +66,4 @@ class GrpcStreamCall(grpc.Call, Iterable[_StreamItemT]): WrapResultFunc = Callable[..., Any] # Type for RPC call arguments tuple -WrapArgsType = Tuple[Any, ...] +WrapArgsType = tuple[Any, ...] diff --git a/ydb/_utilities.py b/ydb/_utilities.py index df152987e..92f2fdc7a 100644 --- a/ydb/_utilities.py +++ b/ydb/_utilities.py @@ -14,7 +14,7 @@ import random import time import urllib.parse -from typing import Dict, List, Optional, TYPE_CHECKING +from typing import Optional, TYPE_CHECKING from . import ydb_version import typing @@ -109,7 +109,7 @@ def get_query_hash(yql_text): return hashlib.sha256(str(yql_text).encode("utf-8")).hexdigest() -class LRUCache(object): +class LRUCache: def __init__(self, capacity=1000): self.items = collections.OrderedDict() self.capacity = capacity @@ -142,7 +142,7 @@ def from_bytes(val): return val -class AsyncResponseIterator(object): +class AsyncResponseIterator: def __init__(self, it, wrapper): self.it = it self.wrapper = wrapper @@ -164,7 +164,7 @@ def __next__(self): return self._next() -class SyncResponseIterator(object): +class SyncResponseIterator: def __init__(self, it, wrapper): self.it = it self.wrapper = wrapper @@ -229,7 +229,7 @@ def get_first_response(waiter): # Module-level thread pool for TCP race (reused across discovery cycles) _TCP_RACE_MAX_WORKERS = 30 -_TCP_RACE_EXECUTOR: Optional[concurrent.futures.ThreadPoolExecutor] = None +_TCP_RACE_EXECUTOR: concurrent.futures.ThreadPoolExecutor | None = None _EXECUTOR_LOCK = threading.Lock() _ATEXIT_REGISTERED = False @@ -268,7 +268,7 @@ def _shutdown_executor(): def _check_fastest_endpoint( - endpoints: List["resolver.EndpointInfo"], timeout: float = 5.0 + endpoints: list["resolver.EndpointInfo"], timeout: float = 5.0 ) -> Optional["resolver.EndpointInfo"]: """ Perform TCP race using a bounded thread pool and return the fastest endpoint. @@ -280,7 +280,7 @@ def _check_fastest_endpoint( per location to ensure fair representation of all locations in the race. If there are still too many locations, randomly samples them to stay within the limit. - :param endpoints: List of resolver.EndpointInfo objects + :param endpoints: list of resolver.EndpointInfo objects :param timeout: Maximum time to wait for any connection (seconds) :return: Fastest endpoint that connected successfully, or None if all failed """ @@ -329,7 +329,7 @@ def try_connect(endpoint: "resolver.EndpointInfo") -> Optional["resolver.Endpoin return None executor = _get_executor() - futures_list: List[concurrent.futures.Future] = [executor.submit(try_connect, ep) for ep in endpoints] + futures_list: list[concurrent.futures.Future] = [executor.submit(try_connect, ep) for ep in endpoints] try: for fut in concurrent.futures.as_completed(futures_list, timeout=timeout): @@ -346,14 +346,14 @@ def try_connect(endpoint: "resolver.EndpointInfo") -> Optional["resolver.Endpoin return None -def _split_endpoints_by_location(endpoints: List["resolver.EndpointInfo"]) -> Dict[str, List["resolver.EndpointInfo"]]: +def _split_endpoints_by_location(endpoints: list["resolver.EndpointInfo"]) -> dict[str, list["resolver.EndpointInfo"]]: """ Group endpoints by their location. - :param endpoints: List of resolver.EndpointInfo objects - :return: Dictionary mapping location -> list of resolver.EndpointInfo + :param endpoints: list of resolver.EndpointInfo objects + :return: dictionary mapping location -> list of resolver.EndpointInfo """ - result: Dict[str, List["resolver.EndpointInfo"]] = {} + result: dict[str, list["resolver.EndpointInfo"]] = {} for endpoint in endpoints: location = endpoint.location if location not in result: @@ -362,11 +362,11 @@ def _split_endpoints_by_location(endpoints: List["resolver.EndpointInfo"]) -> Di return result -def _get_random_endpoints(endpoints: List["resolver.EndpointInfo"], count: int) -> List["resolver.EndpointInfo"]: +def _get_random_endpoints(endpoints: list["resolver.EndpointInfo"], count: int) -> list["resolver.EndpointInfo"]: """ Get random sample of endpoints. - :param endpoints: List of resolver.EndpointInfo objects + :param endpoints: list of resolver.EndpointInfo objects :param count: Maximum number of endpoints to return :return: Random sample of resolver.EndpointInfo """ @@ -376,8 +376,8 @@ def _get_random_endpoints(endpoints: List["resolver.EndpointInfo"], count: int) def detect_local_dc( - endpoints: List["resolver.EndpointInfo"], max_per_location: int = 3, timeout: float = 5.0 -) -> Optional[str]: + endpoints: list["resolver.EndpointInfo"], max_per_location: int = 3, timeout: float = 5.0 +) -> str | None: """ Detect nearest datacenter by performing TCP race between endpoints. @@ -393,7 +393,7 @@ def detect_local_dc( 5. Return the location of the first endpoint that connects successfully 6. If all connections fail, return None - :param endpoints: List of resolver.EndpointInfo objects from discovery + :param endpoints: list of resolver.EndpointInfo objects from discovery :param max_per_location: Maximum number of endpoints to test per location (default: 3, must be >= 1) :param timeout: TCP connection timeout in seconds (default: 5.0, must be > 0) :return: Location string of the nearest datacenter, or None if detection failed diff --git a/ydb/aio/_utilities.py b/ydb/aio/_utilities.py index e3e6699eb..8733d21e5 100644 --- a/ydb/aio/_utilities.py +++ b/ydb/aio/_utilities.py @@ -2,7 +2,6 @@ import logging import random import time -from typing import Dict, List, Optional from .. import resolver @@ -52,8 +51,8 @@ async def get_first_response(): async def _check_fastest_endpoint( - endpoints: List[resolver.EndpointInfo], timeout: float = 5.0 -) -> Optional[resolver.EndpointInfo]: + endpoints: list[resolver.EndpointInfo], timeout: float = 5.0 +) -> resolver.EndpointInfo | None: """ Perform async TCP race: connect to all endpoints concurrently and return the fastest one. @@ -114,15 +113,15 @@ async def try_connect(endpoint): def _split_endpoints_by_location( - endpoints: List[resolver.EndpointInfo], -) -> Dict[str, List[resolver.EndpointInfo]]: + endpoints: list[resolver.EndpointInfo], +) -> dict[str, list[resolver.EndpointInfo]]: """ Group endpoints by their location. :param endpoints: List of resolver.EndpointInfo objects :return: Dictionary mapping location -> list of resolver.EndpointInfo """ - result: Dict[str, List[resolver.EndpointInfo]] = {} + result: dict[str, list[resolver.EndpointInfo]] = {} for endpoint in endpoints: location = endpoint.location if location not in result: @@ -131,7 +130,7 @@ def _split_endpoints_by_location( return result -def _get_random_endpoints(endpoints: List[resolver.EndpointInfo], count: int) -> List[resolver.EndpointInfo]: +def _get_random_endpoints(endpoints: list[resolver.EndpointInfo], count: int) -> list[resolver.EndpointInfo]: """ Get random sample of endpoints. @@ -145,8 +144,8 @@ def _get_random_endpoints(endpoints: List[resolver.EndpointInfo], count: int) -> async def detect_local_dc( - endpoints: List[resolver.EndpointInfo], max_per_location: int = 3, timeout: float = 5.0 -) -> Optional[str]: + endpoints: list[resolver.EndpointInfo], max_per_location: int = 3, timeout: float = 5.0 +) -> str | None: """ Detect nearest datacenter by performing async TCP race between endpoints. diff --git a/ydb/aio/connection.py b/ydb/aio/connection.py index 3bc72f920..edf6125aa 100644 --- a/ydb/aio/connection.py +++ b/ydb/aio/connection.py @@ -2,7 +2,7 @@ import logging import asyncio -from typing import Any, Callable, Dict, List, Optional, Tuple, TYPE_CHECKING +from typing import Any, Callable, TYPE_CHECKING import collections import grpc @@ -48,15 +48,15 @@ async def _construct_metadata( driver_config: DriverConfig, - settings: Optional[BaseRequestSettings], -) -> List[Tuple[str, str]]: + settings: BaseRequestSettings | None, +) -> list[tuple[str, str]]: """ Translates request settings into RPC metadata :param driver_config: A driver config :param settings: An instance of BaseRequestSettings :return: RPC metadata """ - metadata: List[Tuple[str, str]] = [] + metadata: list[tuple[str, str]] = [] if driver_config.database is not None: metadata.append((YDB_DATABASE_HEADER, driver_config.database)) @@ -163,8 +163,8 @@ class Connection: def __init__( self, endpoint: str, - driver_config: Optional[DriverConfig] = None, - endpoint_options: Optional[EndpointOptions] = None, + driver_config: DriverConfig | None = None, + endpoint_options: EndpointOptions | None = None, ) -> None: self.endpoint = endpoint self.endpoint_key = EndpointKey(self.endpoint, getattr(endpoint_options, "node_id", None)) @@ -175,12 +175,12 @@ def __init__( self._channel = channel_factory(self.endpoint, driver_config, grpc.aio, endpoint_options=endpoint_options) self._driver_config = driver_config - self._stub_instances: Dict[Any, Any] = {} - self._cleanup_callbacks: List[Callable[["Connection"], None]] = [] + self._stub_instances: dict[Any, Any] = {} + self._cleanup_callbacks: list[Callable[["Connection"], None]] = [] for stub in _stubs_list: self._stub_instances[stub] = stub(self._channel) - self.calls: Dict[Any, asyncio.Future[Any]] = {} + self.calls: dict[Any, asyncio.Future[Any]] = {} self.closing = False def _prepare_stub_instance(self, stub: Any) -> None: @@ -188,8 +188,8 @@ def _prepare_stub_instance(self, stub: Any) -> None: self._stub_instances[stub] = stub(self._channel) async def _prepare_call( - self, stub: Any, rpc_name: str, request: Any, settings: Optional[BaseRequestSettings] - ) -> Tuple[_RpcState, float, List[Tuple[str, str]]]: + self, stub: Any, rpc_name: str, request: Any, settings: BaseRequestSettings | None + ) -> tuple[_RpcState, float, list[tuple[str, str]]]: timeout, metadata = _get_request_timeout(settings), await _construct_metadata(self._driver_config, settings) # type: ignore[arg-type] _set_server_timeouts(request, settings, timeout) self._prepare_stub_instance(stub) @@ -208,10 +208,10 @@ async def __call__( request: Any, stub: Any, rpc_name: str, - wrap_result: Optional[Callable[..., Any]] = None, - settings: Optional[BaseRequestSettings] = None, - wrap_args: Tuple[Any, ...] = (), - on_disconnected: Optional[Callable[..., Any]] = None, + wrap_result: Callable[..., Any] | None = None, + settings: BaseRequestSettings | None = None, + wrap_args: tuple[Any, ...] = (), + on_disconnected: Callable[..., Any] | None = None, ) -> Any: """ Async method to execute request diff --git a/ydb/aio/coordination/client.py b/ydb/aio/coordination/client.py index 5983f8c8e..1b4716338 100644 --- a/ydb/aio/coordination/client.py +++ b/ydb/aio/coordination/client.py @@ -1,4 +1,4 @@ -from typing import Optional, TYPE_CHECKING +from typing import TYPE_CHECKING from ..._grpc.grpcwrapper.ydb_coordination import ( CreateNodeRequest, @@ -15,7 +15,7 @@ class CoordinationClient(BaseCoordinationClient["AsyncDriver"]): - async def create_node(self, path: str, config: Optional[NodeConfig] = None, settings=None): + async def create_node(self, path: str, config: NodeConfig | None = None, settings=None): self._log_experimental_api() return await self._call_create( diff --git a/ydb/aio/coordination/reconnector.py b/ydb/aio/coordination/reconnector.py index e5875f2db..9563ac541 100644 --- a/ydb/aio/coordination/reconnector.py +++ b/ydb/aio/coordination/reconnector.py @@ -2,7 +2,7 @@ import asyncio import logging -from typing import Any, Dict, Optional +from typing import Any from ... import issues from ..._grpc.grpcwrapper.common_utils import IToProto @@ -22,11 +22,11 @@ def __init__(self, driver, node_path: str, timeout_millis: int = 30000): self._stream = None self._session_id = None - self._pending_futures: Dict[int, asyncio.Future[Any]] = {} - self._pending_requests: Dict[int, IToProto] = {} + self._pending_futures: dict[int, asyncio.Future[Any]] = {} + self._pending_requests: dict[int, IToProto] = {} self._send_lock = asyncio.Lock() - self._connection_task: Optional[asyncio.Task[Any]] = None + self._connection_task: asyncio.Task[Any] | None = None self._closed = False async def stop(self): diff --git a/ydb/aio/coordination/stream.py b/ydb/aio/coordination/stream.py index a04280e6a..893849d7c 100644 --- a/ydb/aio/coordination/stream.py +++ b/ydb/aio/coordination/stream.py @@ -2,7 +2,6 @@ import asyncio import logging -from typing import Optional from ... import issues, _apis from ..._grpc.grpcwrapper.common_utils import IToProto, GrpcWrapperAsyncIO @@ -21,16 +20,16 @@ def __init__(self, driver): self._stream = GrpcWrapperAsyncIO(FromServer.from_proto) self._incoming = asyncio.Queue() - self._reader_task: Optional[asyncio.Task] = None + self._reader_task: asyncio.Task | None = None self._closed = False - self.session_id: Optional[int] = None + self.session_id: int | None = None async def start_session( self, path: str, timeout_millis: int, - session_id: Optional[int] = None, + session_id: int | None = None, ): await self._stream.start( self._driver, @@ -100,7 +99,7 @@ async def send(self, req: IToProto): raise issues.Error("Coordination stream closed") self._stream.write(req) - async def receive(self, timeout: Optional[float] = None): + async def receive(self, timeout: float | None = None): if self._closed: raise issues.Error("Coordination stream closed") diff --git a/ydb/aio/driver.py b/ydb/aio/driver.py index 88221e947..5d23b5326 100644 --- a/ydb/aio/driver.py +++ b/ydb/aio/driver.py @@ -14,8 +14,8 @@ class DriverConfig(ydb.DriverConfig): def default_from_endpoint_and_database( cls, endpoint: str, - database: Optional[str] = None, - root_certificates: Optional[bytes] = None, + database: str | None = None, + root_certificates: bytes | None = None, credentials: Optional["Credentials"] = None, **kwargs: Any, ) -> "DriverConfig": @@ -31,7 +31,7 @@ def default_from_endpoint_and_database( def default_from_connection_string( cls, connection_string: str, - root_certificates: Optional[bytes] = None, + root_certificates: bytes | None = None, credentials: Optional["Credentials"] = None, **kwargs: Any, ) -> "DriverConfig": @@ -50,11 +50,11 @@ class Driver(pool.ConnectionPool): def __init__( self, - driver_config: Optional[ydb.DriverConfig] = None, - connection_string: Optional[str] = None, - endpoint: Optional[str] = None, - database: Optional[str] = None, - root_certificates: Optional[bytes] = None, + driver_config: ydb.DriverConfig | None = None, + connection_string: str | None = None, + endpoint: str | None = None, + database: str | None = None, + root_certificates: bytes | None = None, credentials: Optional["Credentials"] = None, **kwargs: Any, ) -> None: diff --git a/ydb/aio/oauth2_token_exchange.py b/ydb/aio/oauth2_token_exchange.py index f01ce94fb..77cf00e4c 100644 --- a/ydb/aio/oauth2_token_exchange.py +++ b/ydb/aio/oauth2_token_exchange.py @@ -15,11 +15,11 @@ class Oauth2TokenExchangeCredentials(AbstractExpiringTokenCredentials, Oauth2Tok def __init__( self, token_endpoint: str, - subject_token_source: typing.Optional[TokenSource] = None, - actor_token_source: typing.Optional[TokenSource] = None, - audience: typing.Union[typing.List[str], str, None] = None, - scope: typing.Union[typing.List[str], str, None] = None, - resource: typing.Optional[str] = None, + subject_token_source: TokenSource | None = None, + actor_token_source: TokenSource | None = None, + audience: list[str] | str | None = None, + scope: list[str] | str | None = None, + resource: str | None = None, grant_type: str = "urn:ietf:params:oauth:grant-type:token-exchange", requested_token_type: str = "urn:ietf:params:oauth:token-type:access_token", ): diff --git a/ydb/aio/pool.py b/ydb/aio/pool.py index 150274368..6bafe6708 100644 --- a/ydb/aio/pool.py +++ b/ydb/aio/pool.py @@ -3,7 +3,7 @@ import asyncio import logging import random -from typing import Any, Callable, Optional, Tuple, TYPE_CHECKING, cast +from typing import Any, Callable, TYPE_CHECKING, cast from ydb import issues from ydb.observability.tracing import SpanName, create_ydb_span @@ -26,11 +26,11 @@ def __init__(self, use_all_nodes: bool = False) -> None: self.lock = resolver._FakeLock() # Mock lock to emulate thread safety self._event: asyncio.Event = asyncio.Event() self._fast_fail_event: asyncio.Event = asyncio.Event() - self._fast_fail_error: Optional[Exception] = None + self._fast_fail_error: Exception | None = None async def get( # async version with different Connection type self, - preferred_endpoint: Optional[EndpointKey] = None, + preferred_endpoint: EndpointKey | None = None, fast_fail: bool = False, wait_timeout: float = 10.0, ) -> Connection: @@ -57,7 +57,7 @@ async def get( # async version with different Connection type raise issues.ConnectionLost("Couldn't find valid connection") - def add(self, connection: Optional[Connection], preferred: bool = False) -> bool: # type: ignore[override] # async Connection type + def add(self, connection: Connection | None, preferred: bool = False) -> bool: # type: ignore[override] # async Connection type if connection is None: return False @@ -76,7 +76,7 @@ def add(self, connection: Optional[Connection], preferred: bool = False) -> bool return True - def complete_discovery(self, error: Optional[Exception]) -> None: + def complete_discovery(self, error: Exception | None) -> None: self._fast_fail_error = error self._fast_fail_event.set() @@ -249,7 +249,7 @@ def __init__(self, driver_config: "DriverConfig") -> None: self._grpc_init = Connection(self._driver_config.endpoint, self._driver_config) self._stopped = False self._stopping = False - self._discovery: Optional[Discovery] = None + self._discovery: Discovery | None = None self._discovery_task: "asyncio.Task[None]" if driver_config.disable_discovery: @@ -316,11 +316,11 @@ def _pessimize_node(self, node_id: int) -> None: if node_id <= 0: return - connection = cast(Optional[Connection], self._store.get_connection_by_node_id(node_id)) + connection = cast(Connection | None, self._store.get_connection_by_node_id(node_id)) if connection is not None: asyncio.get_running_loop().create_task(self._on_disconnected(connection)()) - async def wait(self, timeout: Optional[float] = 7.0, fail_fast: bool = False) -> None: # type: ignore[override] # async override of sync method + async def wait(self, timeout: float | None = 7.0, fail_fast: bool = False) -> None: # type: ignore[override] # async override of sync method with create_ydb_span(SpanName.DRIVER_INITIALIZE, self._driver_config, kind="internal").attach_context(): await self._store.get(fast_fail=fail_fast, wait_timeout=timeout if timeout is not None else 7.0) @@ -340,10 +340,10 @@ async def __call__( request: Any, stub: Any, rpc_name: str, - wrap_result: Optional[Callable[..., Any]] = None, - settings: Optional["BaseRequestSettings"] = None, - wrap_args: Tuple[Any, ...] = (), - preferred_endpoint: Optional[EndpointKey] = None, + wrap_result: Callable[..., Any] | None = None, + settings: "BaseRequestSettings" | None = None, + wrap_args: tuple[Any, ...] = (), + preferred_endpoint: EndpointKey | None = None, fast_fail: bool = False, ) -> Any: if self._stopped: diff --git a/ydb/aio/query/pool.py b/ydb/aio/query/pool.py index 25f92703f..a648466a1 100644 --- a/ydb/aio/query/pool.py +++ b/ydb/aio/query/pool.py @@ -4,11 +4,7 @@ import logging from typing import ( Callable, - Optional, - List, - Dict, Any, - Union, ) from .session import ( @@ -37,9 +33,9 @@ def __init__( driver: common_utils.SupportedDriverType, size: int = 100, *, - query_client_settings: Optional[QueryClientSettings] = None, - loop: Optional[asyncio.AbstractEventLoop] = None, - name: Optional[str] = None, + query_client_settings: QueryClientSettings | None = None, + loop: asyncio.AbstractEventLoop | None = None, + name: str | None = None, ): """ :param driver: A driver instance @@ -65,7 +61,7 @@ async def _create_new_session(self): logger.debug(f"New session was created for pool. Session id: {session.session_id}") return session - async def acquire(self, timeout: Optional[float] = None) -> QuerySession: + async def acquire(self, timeout: float | None = None) -> QuerySession: """Acquire a session from Session Pool. :param timeout: Seconds to wait when pool is exhausted. Overrides the pool-level acquire_timeout. @@ -150,7 +146,7 @@ async def release(self, session: QuerySession) -> None: self._queue.put_nowait(session) logger.debug("Session returned to queue: %s", session.session_id) - def checkout(self, timeout: Optional[float] = None) -> "SimpleQuerySessionCheckoutAsync": + def checkout(self, timeout: float | None = None) -> "SimpleQuerySessionCheckoutAsync": """Return a Session context manager, that acquires session on enter and releases session on exit. :param timeout: Seconds to wait when pool is exhausted. Overrides the pool-level acquire_timeout. @@ -159,7 +155,7 @@ def checkout(self, timeout: Optional[float] = None) -> "SimpleQuerySessionChecko return SimpleQuerySessionCheckoutAsync(self, timeout) async def retry_operation_async( - self, callee: Callable, retry_settings: Optional[RetrySettings] = None, *args, **kwargs + self, callee: Callable, retry_settings: RetrySettings | None = None, *args, **kwargs ): """Special interface to execute a bunch of commands with session in a safe, retriable way. @@ -180,8 +176,8 @@ async def wrapped_callee(): async def retry_tx_async( self, callee: Callable, - tx_mode: Optional[BaseQueryTxMode] = None, - retry_settings: Optional[RetrySettings] = None, + tx_mode: BaseQueryTxMode | None = None, + retry_settings: RetrySettings | None = None, *args, **kwargs, ): @@ -216,12 +212,12 @@ async def wrapped_callee(): async def execute_with_retries( self, query: str, - parameters: Optional[dict] = None, - retry_settings: Optional[RetrySettings] = None, + parameters: dict | None = None, + retry_settings: RetrySettings | None = None, *args, - pool_id: Optional[str] = None, + pool_id: str | None = None, **kwargs, - ) -> List[convert.ResultSet]: + ) -> list[convert.ResultSet]: """Special interface to execute a one-shot queries in a safe, retriable way. Note: this method loads all data from stream before return, do not use this method with huge read queries. @@ -246,11 +242,11 @@ async def wrapped_callee(): async def explain_with_retries( self, query: str, - parameters: Optional[dict] = None, + parameters: dict | None = None, *, result_format: QueryExplainResultFormat = QueryExplainResultFormat.STR, - retry_settings: Optional[RetrySettings] = None, - ) -> Union[str, Dict[str, Any]]: + retry_settings: RetrySettings | None = None, + ) -> str | dict[str, Any]: """ Explain a query in retriable way. No real query execution will happen. @@ -290,9 +286,9 @@ async def __aexit__(self, exc_type, exc_val, exc_tb): class SimpleQuerySessionCheckoutAsync: - _session: Optional[QuerySession] + _session: QuerySession | None - def __init__(self, pool: QuerySessionPool, timeout: Optional[float] = None): + def __init__(self, pool: QuerySessionPool, timeout: float | None = None): self._pool = pool self._timeout = timeout self._session = None diff --git a/ydb/aio/query/session.py b/ydb/aio/query/session.py index 7f7654580..2b02a7701 100644 --- a/ydb/aio/query/session.py +++ b/ydb/aio/query/session.py @@ -2,10 +2,7 @@ import json from typing import ( - Optional, - Dict, Any, - Union, TYPE_CHECKING, ) @@ -36,13 +33,13 @@ class QuerySession(BaseQuerySession["AsyncDriver"]): """ _loop: asyncio.AbstractEventLoop - _status_stream: Optional[_utilities.AsyncResponseIterator] + _status_stream: _utilities.AsyncResponseIterator | None def __init__( self, driver: "AsyncDriver", - settings: Optional[base.QueryClientSettings] = None, - loop: Optional[asyncio.AbstractEventLoop] = None, + settings: base.QueryClientSettings | None = None, + loop: asyncio.AbstractEventLoop | None = None, ): super(QuerySession, self).__init__(driver, settings) self._loop = loop if loop is not None else asyncio.get_running_loop() @@ -78,7 +75,7 @@ async def _check_session_status_loop(self) -> None: logger.debug("Attach stream error: %s, session_id: %s", e, self._session_id) self._close_session(invalidate=True) - async def delete(self, settings: Optional[BaseRequestSettings] = None) -> None: + async def delete(self, settings: BaseRequestSettings | None = None) -> None: """Deletes a Session of Query Service on server side and releases resources. :return: None @@ -94,7 +91,7 @@ async def delete(self, settings: Optional[BaseRequestSettings] = None) -> None: self._close_session() - async def create(self, settings: Optional[BaseRequestSettings] = None) -> "QuerySession": + async def create(self, settings: BaseRequestSettings | None = None) -> "QuerySession": """Creates a Session of Query Service on server side and attaches it. :return: QuerySession object. @@ -130,13 +127,13 @@ async def execute( syntax: base.QuerySyntax = None, exec_mode: base.QueryExecMode = None, concurrent_result_sets: bool = False, - settings: Optional[BaseRequestSettings] = None, + settings: BaseRequestSettings | None = None, *, - stats_mode: Optional[base.QueryStatsMode] = None, - 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, + stats_mode: base.QueryStatsMode | None = None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None = None, + result_set_format: base.QueryResultSetFormat | None = None, + arrow_format_settings: base.ArrowFormatSettings | None = None, + pool_id: str | None = None, ) -> AsyncResponseContextIterator: """Sends a query to Query Service @@ -201,9 +198,9 @@ async def execute( async def explain( self, query: str, - parameters: Optional[dict] = None, + parameters: dict | None = None, result_format: base.QueryExplainResultFormat = base.QueryExplainResultFormat.STR, - ) -> Union[str, Dict[str, Any]]: + ) -> str | dict[str, Any]: """Explains query result :param query: YQL or SQL query. :param parameters: dict with parameters and YDB types; diff --git a/ydb/aio/query/transaction.py b/ydb/aio/query/transaction.py index 33a00dea4..8cea50e68 100644 --- a/ydb/aio/query/transaction.py +++ b/ydb/aio/query/transaction.py @@ -1,6 +1,5 @@ import logging from typing import ( - Optional, TYPE_CHECKING, ) @@ -81,7 +80,7 @@ async def _ensure_prev_stream_finished(self) -> None: pass self._prev_stream = None - async def begin(self, settings: Optional[BaseRequestSettings] = None) -> "QueryTxContext": + async def begin(self, settings: BaseRequestSettings | None = None) -> "QueryTxContext": """Explicitly begins a transaction :param settings: An additional request settings BaseRequestSettings; @@ -97,7 +96,7 @@ async def begin(self, settings: Optional[BaseRequestSettings] = None) -> "QueryT await self._begin_call(settings) return self - async def commit(self, settings: Optional[BaseRequestSettings] = None) -> None: + async def commit(self, settings: BaseRequestSettings | None = None) -> None: """Calls commit on a transaction if it is open otherwise is no-op. If transaction execution failed then this method raises PreconditionFailed. @@ -130,7 +129,7 @@ async def commit(self, settings: Optional[BaseRequestSettings] = None) -> None: await self._execute_callbacks_async(base.TxEvent.AFTER_COMMIT, exc=e) raise e - async def rollback(self, settings: Optional[BaseRequestSettings] = None) -> None: + async def rollback(self, settings: BaseRequestSettings | None = None) -> None: """Calls rollback on a transaction if it is open otherwise is no-op. If transaction execution failed then this method raises PreconditionFailed. @@ -166,18 +165,18 @@ async def rollback(self, settings: Optional[BaseRequestSettings] = None) -> None async def execute( self, query: str, - parameters: Optional[dict] = None, - commit_tx: Optional[bool] = False, - syntax: Optional[base.QuerySyntax] = None, - exec_mode: Optional[base.QueryExecMode] = None, - concurrent_result_sets: Optional[bool] = False, - settings: Optional[BaseRequestSettings] = None, + parameters: dict | None = None, + commit_tx: bool | None = False, + syntax: base.QuerySyntax | None = None, + exec_mode: base.QueryExecMode | None = None, + concurrent_result_sets: bool | None = False, + settings: BaseRequestSettings | None = None, *, - stats_mode: Optional[base.QueryStatsMode] = None, - 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, + stats_mode: base.QueryStatsMode | None = None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None = None, + result_set_format: base.QueryResultSetFormat | None = None, + arrow_format_settings: base.ArrowFormatSettings | None = None, + pool_id: str | None = None, ) -> AsyncResponseContextIterator: """Sends a query to Query Service diff --git a/ydb/aio/resolver.py b/ydb/aio/resolver.py index fb8c4417e..1da0a8217 100644 --- a/ydb/aio/resolver.py +++ b/ydb/aio/resolver.py @@ -1,4 +1,4 @@ -from typing import Any, Optional +from typing import Any from . import connection as conn_impl @@ -29,7 +29,7 @@ def __init__(self, driver_config: DriverConfig) -> None: super().__init__(driver_config) self._lock = _FakeLock() - async def resolve(self) -> Optional[DiscoveryResult]: # type: ignore[override] # async override of sync method + async def resolve(self) -> DiscoveryResult | None: # type: ignore[override] # async override of sync method self.logger.debug("Preparing initial endpoint to resolve endpoints") endpoint = next(self._endpoints_iter) connection = conn_impl.Connection(endpoint, self._driver_config) diff --git a/ydb/aio/table.py b/ydb/aio/table.py index f76a934bd..930e148a3 100644 --- a/ydb/aio/table.py +++ b/ydb/aio/table.py @@ -5,10 +5,7 @@ from typing import ( Any, - Dict, - List, Optional, - Tuple, TYPE_CHECKING, ) @@ -155,9 +152,9 @@ async def rename_tables(self, rename_items, settings=None): # pylint: disable=W class TableClient(BaseTableClient["AsyncDriver"]): - def __init__(self, driver: "AsyncDriver", table_client_settings: Optional[TableClientSettings] = None) -> None: + def __init__(self, driver: "AsyncDriver", table_client_settings: TableClientSettings | None = None) -> None: super().__init__(driver=driver, table_client_settings=table_client_settings) - self._pool: Optional[SessionPool] = None + self._pool: SessionPool | None = None def __del__(self): if self._pool is not None and not self._pool._terminating: @@ -246,42 +243,42 @@ async def callee(session: Session): async def alter_table( self, path: str, - add_columns: Optional[List["ydb.Column"]] = None, - drop_columns: Optional[List[str]] = None, + add_columns: list["ydb.Column"] | None = None, + drop_columns: list[str] | None = None, settings: Optional["settings_impl.BaseRequestSettings"] = None, - alter_attributes: Optional[Optional[Dict[str, str]]] = None, - add_indexes: Optional[List["ydb.TableIndex"]] = None, - drop_indexes: Optional[List[str]] = None, + alter_attributes: dict[str, str] | None = None, + add_indexes: list["ydb.TableIndex"] | None = None, + drop_indexes: list[str] | None = None, set_ttl_settings: Optional["ydb.TtlSettings"] = None, - drop_ttl_settings: Optional[Any] = None, - add_column_families: Optional[List["ydb.ColumnFamily"]] = None, - alter_column_families: Optional[List["ydb.ColumnFamily"]] = None, + drop_ttl_settings: Any | None = None, + add_column_families: list["ydb.ColumnFamily"] | None = None, + alter_column_families: list["ydb.ColumnFamily"] | None = None, alter_storage_settings: Optional["ydb.StorageSettings"] = None, - set_compaction_policy: Optional[str] = None, + set_compaction_policy: str | None = None, alter_partitioning_settings: Optional["ydb.PartitioningSettings"] = None, set_key_bloom_filter: Optional["ydb.FeatureFlag"] = None, set_read_replicas_settings: Optional["ydb.ReadReplicasSettings"] = None, - rename_indexes: Optional[List["ydb.RenameIndexItem"]] = None, + rename_indexes: list["ydb.RenameIndexItem"] | None = None, ) -> "ydb.Operation": """ Alter a YDB table. :param path: A table path - :param add_columns: List of ydb.Column to add - :param drop_columns: List of column names to drop + :param add_columns: list of ydb.Column to add + :param drop_columns: list of column names to drop :param settings: An instance of BaseRequestSettings that describes how rpc should be invoked. - :param alter_attributes: Dict of attributes to alter - :param add_indexes: List of ydb.TableIndex to add - :param drop_indexes: List of index names to drop + :param alter_attributes: dict of attributes to alter + :param add_indexes: list of ydb.TableIndex to add + :param drop_indexes: list of index names to drop :param set_ttl_settings: ydb.TtlSettings to set :param drop_ttl_settings: Any to drop - :param add_column_families: List of ydb.ColumnFamily to add - :param alter_column_families: List of ydb.ColumnFamily to alter + :param add_column_families: list of ydb.ColumnFamily to add + :param alter_column_families: list of ydb.ColumnFamily to alter :param alter_storage_settings: ydb.StorageSettings to alter :param set_compaction_policy: Compaction policy :param alter_partitioning_settings: ydb.PartitioningSettings to alter :param set_key_bloom_filter: ydb.FeatureFlag to set key bloom filter - :param rename_indexes: List of ydb.RenameIndexItem to rename + :param rename_indexes: list of ydb.RenameIndexItem to rename :return: Operation or YDB error otherwise. """ @@ -364,13 +361,13 @@ async def callee(session: Session): async def copy_tables( self, - source_destination_pairs: List[Tuple[str, str]], + source_destination_pairs: list[tuple[str, str]], settings: Optional["settings_impl.BaseRequestSettings"] = None, ) -> "ydb.Operation": """ Copy a YDB tables. - :param source_destination_pairs: List of tuples (source_path, destination_path) + :param source_destination_pairs: list of tuples (source_path, destination_path) :param settings: An instance of BaseRequestSettings that describes how rpc should be invoked. :return: Operation or YDB error otherwise. @@ -386,13 +383,13 @@ async def callee(session: Session): async def rename_tables( self, - rename_items: List[Tuple[str, str]], + rename_items: list[tuple[str, str]], settings: Optional["settings_impl.BaseRequestSettings"] = None, ) -> "ydb.Operation": """ Rename a YDB tables. - :param rename_items: List of tuples (current_name, desired_name) + :param rename_items: list of tuples (current_name, desired_name) :param settings: An instance of BaseRequestSettings that describes how rpc should be invoked. :return: Operation or YDB error otherwise. @@ -541,7 +538,7 @@ def _create(self) -> Session: self._logger.debug("Created session %s", session) return session - async def _init_session_logic(self, session: ydb.ISession) -> typing.Optional[ydb.ISession]: + async def _init_session_logic(self, session: ydb.ISession) -> ydb.ISession | None: try: await self._driver.wait(self._driver_await_timeout) session = await session.create(self._req_settings) @@ -556,7 +553,7 @@ async def _init_session_logic(self, session: ydb.ISession) -> typing.Optional[yd return None - async def _init_session(self, session: ydb.ISession, retry_num: int = None) -> typing.Optional[ydb.ISession]: + async def _init_session(self, session: ydb.ISession, retry_num: int = None) -> ydb.ISession | None: """ :param retry_num: Number of retries. If None - retries until success. :return: @@ -569,9 +566,7 @@ async def _init_session(self, session: ydb.ISession, retry_num: int = None) -> t i += 1 return None - async def _prepare_session( - self, timeout: typing.Optional[float], retry_num: typing.Optional[int] - ) -> typing.Optional[ydb.ISession]: + async def _prepare_session(self, timeout: float | None, retry_num: int | None) -> ydb.ISession | None: session = self._create() try: new_sess = await asyncio.wait_for(self._init_session(session, retry_num=retry_num), timeout=timeout) @@ -583,7 +578,7 @@ async def _prepare_session( self._destroy(session) raise e - async def _get_session_from_queue(self, timeout: typing.Optional[float]) -> Session: + async def _get_session_from_queue(self, timeout: float | None) -> Session: task_wait = asyncio.ensure_future(asyncio.wait_for(self._active_queue.get(), timeout=timeout)) task_should_stop = asyncio.ensure_future(self._should_stop.wait()) try: @@ -603,9 +598,9 @@ async def _get_session_from_queue(self, timeout: typing.Optional[float]) -> Sess async def acquire( self, - timeout: typing.Optional[float] = None, - retry_timeout: typing.Optional[float] = None, - retry_num: typing.Optional[int] = None, + timeout: float | None = None, + retry_timeout: float | None = None, + retry_num: int | None = None, ) -> Session: if self._should_stop.is_set(): self._logger.error("Take session from closed session pool") @@ -707,7 +702,7 @@ async def _pick_for_keepalive(self): await self._active_queue.put((priority, session)) return None - async def _send_keep_alive(self, session: typing.Optional[ydb.ISession]) -> bool: + async def _send_keep_alive(self, session: ydb.ISession | None) -> bool: if session is None: return False if self._should_stop.is_set(): diff --git a/ydb/auth_helpers.py b/ydb/auth_helpers.py index e90d95254..190f0ebca 100644 --- a/ydb/auth_helpers.py +++ b/ydb/auth_helpers.py @@ -1,6 +1,5 @@ # -*- coding: utf-8 -*- import os -from typing import Optional def read_bytes(f): @@ -8,7 +7,7 @@ def read_bytes(f): return fr.read() -def load_ydb_root_certificate(path: Optional[str] = None): +def load_ydb_root_certificate(path: str | None = None): path = path if path is not None else os.getenv("YDB_SSL_ROOT_CERTIFICATES_FILE") if path is not None and os.path.exists(path): return read_bytes(path) diff --git a/ydb/connection.py b/ydb/connection.py index 052b563ef..fa08b2441 100644 --- a/ydb/connection.py +++ b/ydb/connection.py @@ -8,10 +8,7 @@ from typing import ( Any, Callable, - Dict, - List, Optional, - Tuple, Union, TYPE_CHECKING, ) @@ -80,8 +77,8 @@ def _log_request(rpc_state: "_RpcState", request: Any) -> None: def _rpc_error_handler( rpc_state: Union["_RpcState", str], - rpc_error: Union[grpc.RpcError, grpc.aio.AioRpcError, grpc.Call, grpc.aio.Call], - on_disconnected: Optional[Callable[[], None]] = None, + rpc_error: grpc.RpcError | grpc.aio.AioRpcError | grpc.Call | grpc.aio.Call, + on_disconnected: Callable[[], None] | None = None, use_unavailable: bool = False, ) -> issues.Error: """ @@ -199,7 +196,7 @@ def _get_request_timeout(settings): return settings.timeout -class EndpointOptions(object): +class EndpointOptions: __slots__ = ("ssl_target_name_override", "node_id", "address", "port", "location") def __init__(self, ssl_target_name_override=None, node_id=None, address=None, port=None, location=None): @@ -273,7 +270,7 @@ class _RpcState: rpc_name: str endpoint: str rendezvous: Any - metadata_kv: Optional[Dict[str, set]] + metadata_kv: dict[str, set] | None endpoint_key: "EndpointKey" def __init__(self, stub_instance: Any, rpc_name: str, endpoint: str, endpoint_key: "EndpointKey") -> None: @@ -298,7 +295,7 @@ def __call__(self, *args: Any, **kwargs: Any) -> Any: except AttributeError: return self.rpc(*args, **kwargs) - def trailing_metadata(self) -> Dict[str, set]: + def trailing_metadata(self) -> dict[str, set]: """Trailing metadata of the call.""" if self.metadata_kv is None: self.metadata_kv = collections.defaultdict(set) @@ -307,7 +304,7 @@ def trailing_metadata(self) -> Dict[str, set]: return self.metadata_kv - def future(self, *args: Any, **kwargs: Any) -> Tuple[Any, "futures.Future[Any]"]: + def future(self, *args: Any, **kwargs: Any) -> tuple[Any, "futures.Future[Any]"]: self.rendezvous = self.rpc.future(*args, **kwargs) self.result_future = futures.Future() @@ -365,7 +362,7 @@ def channel_factory(endpoint, driver_config, channel_provider=None, endpoint_opt ) -class EndpointKey(object): +class EndpointKey: __slots__ = ("endpoint", "node_id") def __init__(self, endpoint, node_id): @@ -400,7 +397,7 @@ def __getattr__(self, item): return getattr(self.resp, item) -class Connection(object): +class Connection: __slots__ = ( "endpoint", "_channel", @@ -423,7 +420,7 @@ def __init__( self, endpoint: str, driver_config: Optional["DriverConfig"] = None, - endpoint_options: Optional[EndpointOptions] = None, + endpoint_options: EndpointOptions | None = None, ) -> None: """ Object that wraps gRPC channel and encapsulates gRPC request execution logic @@ -439,9 +436,9 @@ def __init__( self.endpoint_key = EndpointKey(endpoint, getattr(endpoint_options, "node_id", None)) self._channel = channel_factory(self.endpoint, driver_config, endpoint_options=endpoint_options) self._driver_config = driver_config - self._call_states: Dict[Any, "_RpcState"] = {} - self._stub_instances: Dict[Any, Any] = {} - self._cleanup_callbacks: List[Callable[["Connection"], None]] = [] + self._call_states: dict[Any, "_RpcState"] = {} + self._stub_instances: dict[Any, Any] = {} + self._cleanup_callbacks: list[Callable[["Connection"], None]] = [] # pre-initialize stubs for stub in _stubs_list: self._stub_instances[stub] = stub(self._channel) @@ -485,10 +482,10 @@ def future( request: Any, stub: Any, rpc_name: str, - wrap_result: Optional[Callable[..., Any]] = None, + wrap_result: Callable[..., Any] | None = None, settings: Optional["BaseRequestSettings"] = None, - wrap_args: Tuple[Any, ...] = (), - on_disconnected: Optional[Callable[[], None]] = None, + wrap_args: tuple[Any, ...] = (), + on_disconnected: Callable[[], None] | None = None, ) -> "futures.Future[Any]": """ Sends request constructed by client @@ -525,10 +522,10 @@ def __call__( request: Any, stub: Any, rpc_name: str, - wrap_result: Optional[Callable[..., Any]] = None, + wrap_result: Callable[..., Any] | None = None, settings: Optional["BaseRequestSettings"] = None, - wrap_args: Tuple[Any, ...] = (), - on_disconnected: Optional[Callable[[], None]] = None, + wrap_args: tuple[Any, ...] = (), + on_disconnected: Callable[[], None] | None = None, ) -> Any: """ Synchronously sends request constructed by client library diff --git a/ydb/coordination/client.py b/ydb/coordination/client.py index 783cea1b5..af55e0f86 100644 --- a/ydb/coordination/client.py +++ b/ydb/coordination/client.py @@ -1,5 +1,5 @@ import logging -from typing import Optional, TYPE_CHECKING +from typing import TYPE_CHECKING from .._grpc.grpcwrapper.ydb_coordination import ( CreateNodeRequest, @@ -18,7 +18,7 @@ class CoordinationClient(BaseCoordinationClient["SyncDriver"]): - def create_node(self, path: str, config: Optional[NodeConfig] = None, settings=None): + def create_node(self, path: str, config: NodeConfig | None = None, settings=None): self._log_experimental_api() return self._call_create( diff --git a/ydb/coordination/semaphore.py b/ydb/coordination/semaphore.py index 10e53e2f9..6ed38f3d7 100644 --- a/ydb/coordination/semaphore.py +++ b/ydb/coordination/semaphore.py @@ -1,5 +1,3 @@ -from typing import Optional - from .. import issues from .._topic_common.common import _get_shared_event_loop, CallFromSyncToAsync from ..aio.coordination.semaphore import CoordinationSemaphore as CoordinationSemaphoreAio @@ -32,14 +30,14 @@ def __exit__(self, exc_type, exc_val, exc_tb): except Exception: pass - def acquire(self, count: int = 1, timeout: Optional[float] = None): + def acquire(self, count: int = 1, timeout: float | None = None): self._check_closed() return self._caller.safe_call_with_result( self._async_semaphore.acquire(count), timeout, ) - def release(self, timeout: Optional[float] = None): + def release(self, timeout: float | None = None): if self._closed: return return self._caller.safe_call_with_result( @@ -47,21 +45,21 @@ def release(self, timeout: Optional[float] = None): timeout, ) - def describe(self, timeout: Optional[float] = None): + def describe(self, timeout: float | None = None): self._check_closed() return self._caller.safe_call_with_result( self._async_semaphore.describe(), timeout, ) - def update(self, new_data: bytes, timeout: Optional[float] = None): + def update(self, new_data: bytes, timeout: float | None = None): self._check_closed() return self._caller.safe_call_with_result( self._async_semaphore.update(new_data), timeout, ) - def close(self, timeout: Optional[float] = None): + def close(self, timeout: float | None = None): if self._closed: return try: diff --git a/ydb/credentials.py b/ydb/credentials.py index 28ea8b52f..58a9b0d7b 100644 --- a/ydb/credentials.py +++ b/ydb/credentials.py @@ -22,7 +22,7 @@ logger = logging.getLogger(__name__) -class AtMostOneExecution(object): +class AtMostOneExecution: def __init__(self): self._can_schedule = True self._lock = threading.Lock() diff --git a/ydb/driver.py b/ydb/driver.py index 12b8519b6..26638dc89 100644 --- a/ydb/driver.py +++ b/ydb/driver.py @@ -2,7 +2,7 @@ import grpc import logging import os -from typing import Any, List, Optional, Tuple, Type, TYPE_CHECKING +from typing import Any, Optional, TYPE_CHECKING from . import credentials as credentials_impl, table, scheme, pool from . import tracing @@ -27,7 +27,7 @@ class RPCCompression: def default_credentials( credentials: Optional["Credentials"] = None, - tracer: Optional[tracing.Tracer] = None, + tracer: tracing.Tracer | None = None, ) -> "Credentials": tracer = tracer if tracer is not None else tracing.Tracer(None) with tracer.trace("Driver.default_credentials") as ctx: @@ -39,7 +39,7 @@ def default_credentials( return credentials -def credentials_from_env_variables(tracer: Optional[tracing.Tracer] = None) -> "Credentials": +def credentials_from_env_variables(tracer: tracing.Tracer | None = None) -> "Credentials": tracer = tracer if tracer is not None else tracing.Tracer(None) with tracer.trace("Driver.credentials_from_env_variables") as ctx: service_account_key_file = os.getenv("YDB_SERVICE_ACCOUNT_KEY_FILE_CREDENTIALS") @@ -89,7 +89,7 @@ def credentials_from_env_variables(tracer: Optional[tracing.Tracer] = None) -> " return iam.MetadataUrlCredentials(tracer=tracer) -class DriverConfig(object): +class DriverConfig: __slots__ = ( "endpoint", "database", @@ -119,29 +119,29 @@ class DriverConfig(object): def __init__( self, endpoint: str, - database: Optional[str] = None, - ca_cert: Optional[str] = None, - auth_token: Optional[str] = None, - channel_options: Optional[List[Tuple[str, Any]]] = None, + database: str | None = None, + ca_cert: str | None = None, + auth_token: str | None = None, + channel_options: list[tuple[str, Any]] | None = None, credentials: Optional["Credentials"] = None, use_all_nodes: bool = True, - root_certificates: Optional[bytes] = None, - certificate_chain: Optional[bytes] = None, - private_key: Optional[bytes] = None, - grpc_keep_alive_timeout: Optional[int] = None, + root_certificates: bytes | None = None, + certificate_chain: bytes | None = None, + private_key: bytes | None = None, + grpc_keep_alive_timeout: int | None = None, table_client_settings: Optional["TableClientSettings"] = None, - topic_client_settings: Optional[Any] = None, + topic_client_settings: Any | None = None, query_client_settings: Optional["QueryClientSettings"] = None, - endpoints: Optional[List[str]] = None, + endpoints: list[str] | None = None, primary_user_agent: str = "python-library", - tracer: Optional[tracing.Tracer] = None, + tracer: tracing.Tracer | None = None, grpc_lb_policy_name: str = "round_robin", discovery_request_timeout: int = 10, - compression: Optional[grpc.Compression] = None, + compression: grpc.Compression | None = None, disable_discovery: bool = False, detect_local_dc: bool = False, *, - _additional_sdk_headers: Tuple[str, ...] = (), + _additional_sdk_headers: tuple[str, ...] = (), ) -> None: """ A driver config to initialize a driver instance @@ -206,8 +206,8 @@ def set_database(self, database: str) -> "DriverConfig": def default_from_endpoint_and_database( cls, endpoint: str, - database: Optional[str] = None, - root_certificates: Optional[bytes] = None, + database: str | None = None, + root_certificates: bytes | None = None, credentials: Optional["Credentials"] = None, **kwargs: Any, ) -> "DriverConfig": @@ -223,7 +223,7 @@ def default_from_endpoint_and_database( def default_from_connection_string( cls, connection_string: str, - root_certificates: Optional[bytes] = None, + root_certificates: bytes | None = None, credentials: Optional["Credentials"] = None, **kwargs: Any, ) -> "DriverConfig": @@ -254,13 +254,13 @@ def _update_attrs_by_kwargs(self, **kwargs: Any) -> None: def get_config( - driver_config: Optional[DriverConfig] = None, - connection_string: Optional[str] = None, - endpoint: Optional[str] = None, - database: Optional[str] = None, - root_certificates: Optional[bytes] = None, + driver_config: DriverConfig | None = None, + connection_string: str | None = None, + endpoint: str | None = None, + database: str | None = None, + root_certificates: bytes | None = None, credentials: Optional["Credentials"] = None, - config_class: Type[DriverConfig] = DriverConfig, + config_class: type[DriverConfig] = DriverConfig, **kwargs: Any, ) -> DriverConfig: if driver_config is None: @@ -295,11 +295,11 @@ class Driver(pool.ConnectionPool): def __init__( self, - driver_config: Optional[DriverConfig] = None, - connection_string: Optional[str] = None, - endpoint: Optional[str] = None, - database: Optional[str] = None, - root_certificates: Optional[bytes] = None, + driver_config: DriverConfig | None = None, + connection_string: str | None = None, + endpoint: str | None = None, + database: str | None = None, + root_certificates: bytes | None = None, credentials: Optional["Credentials"] = None, **kwargs: Any, ) -> None: diff --git a/ydb/export.py b/ydb/export.py index 060b36b90..3ae902915 100644 --- a/ydb/export.py +++ b/ydb/export.py @@ -241,7 +241,7 @@ def _export_to_s3_request_factory(settings): return request -class ExportClient(object): +class ExportClient: def __init__(self, driver): self._driver = driver diff --git a/ydb/import_client.py b/ydb/import_client.py index 0eead2005..9abb8e5ef 100644 --- a/ydb/import_client.py +++ b/ydb/import_client.py @@ -142,7 +142,7 @@ def _import_from_s3_request_factory(settings): return request -class ImportClient(object): +class ImportClient: def __init__(self, driver): self._driver = driver diff --git a/ydb/issues.py b/ydb/issues.py index fba0ee0c6..d4d65afce 100644 --- a/ydb/issues.py +++ b/ydb/issues.py @@ -5,7 +5,7 @@ import enum import queue import typing -from typing import ClassVar, Optional, Iterable, Any, Union, Protocol, runtime_checkable +from typing import ClassVar, Iterable, Any, Protocol, runtime_checkable from . import _apis @@ -21,7 +21,7 @@ class _StatusResponseProtocol(Protocol): """Protocol for objects that have status and issues attributes.""" @property - def status(self) -> Union[StatusCode, int]: ... + def status(self) -> StatusCode | int: ... @property def issues(self) -> Iterable[Any]: ... @@ -76,131 +76,131 @@ def __init__(self, message: str, issue_code: int, severity: int, issues) -> None class Error(Exception): - status: ClassVar[Optional[StatusCode]] = None + status: ClassVar[StatusCode | None] = None - def __init__(self, message: str, issues: typing.Optional[typing.Iterable[_IssueMessage]] = None): + def __init__(self, message: str, issues: typing.Iterable[_IssueMessage] | None = None): super(Error, self).__init__(message) self.issues = issues self.message = message class TruncatedResponseError(Error): - status: ClassVar[Optional[StatusCode]] = None + status: ClassVar[StatusCode | None] = None class ConnectionError(Error): - status: ClassVar[Optional[StatusCode]] = None + status: ClassVar[StatusCode | None] = None class ConnectionFailure(ConnectionError): - status: ClassVar[Optional[StatusCode]] = StatusCode.CONNECTION_FAILURE + status: ClassVar[StatusCode | None] = StatusCode.CONNECTION_FAILURE class ConnectionLost(ConnectionError): - status: ClassVar[Optional[StatusCode]] = StatusCode.CONNECTION_LOST + status: ClassVar[StatusCode | None] = StatusCode.CONNECTION_LOST class DeadlineExceed(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.DEADLINE_EXCEEDED + status: ClassVar[StatusCode | None] = StatusCode.DEADLINE_EXCEEDED class Unimplemented(ConnectionError): - status: ClassVar[Optional[StatusCode]] = StatusCode.UNIMPLEMENTED + status: ClassVar[StatusCode | None] = StatusCode.UNIMPLEMENTED class Unauthenticated(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.UNAUTHENTICATED + status: ClassVar[StatusCode | None] = StatusCode.UNAUTHENTICATED class BadRequest(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.BAD_REQUEST + status: ClassVar[StatusCode | None] = StatusCode.BAD_REQUEST class Unauthorized(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.UNAUTHORIZED + status: ClassVar[StatusCode | None] = StatusCode.UNAUTHORIZED class InternalError(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.INTERNAL_ERROR + status: ClassVar[StatusCode | None] = StatusCode.INTERNAL_ERROR class Aborted(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.ABORTED + status: ClassVar[StatusCode | None] = StatusCode.ABORTED class Unavailable(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.UNAVAILABLE + status: ClassVar[StatusCode | None] = StatusCode.UNAVAILABLE class Overloaded(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.OVERLOADED + status: ClassVar[StatusCode | None] = StatusCode.OVERLOADED class SchemeError(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.SCHEME_ERROR + status: ClassVar[StatusCode | None] = StatusCode.SCHEME_ERROR class GenericError(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.GENERIC_ERROR + status: ClassVar[StatusCode | None] = StatusCode.GENERIC_ERROR class BadSession(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.BAD_SESSION + status: ClassVar[StatusCode | None] = StatusCode.BAD_SESSION class Timeout(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.TIMEOUT + status: ClassVar[StatusCode | None] = StatusCode.TIMEOUT class PreconditionFailed(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.PRECONDITION_FAILED + status: ClassVar[StatusCode | None] = StatusCode.PRECONDITION_FAILED class NotFound(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.NOT_FOUND + status: ClassVar[StatusCode | None] = StatusCode.NOT_FOUND class AlreadyExists(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.ALREADY_EXISTS + status: ClassVar[StatusCode | None] = StatusCode.ALREADY_EXISTS class SessionExpired(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.SESSION_EXPIRED + status: ClassVar[StatusCode | None] = StatusCode.SESSION_EXPIRED class Cancelled(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.CANCELLED + status: ClassVar[StatusCode | None] = StatusCode.CANCELLED class Undetermined(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.UNDETERMINED + status: ClassVar[StatusCode | None] = StatusCode.UNDETERMINED class Unsupported(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.UNSUPPORTED + status: ClassVar[StatusCode | None] = StatusCode.UNSUPPORTED class SessionBusy(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.SESSION_BUSY + status: ClassVar[StatusCode | None] = StatusCode.SESSION_BUSY class ExternalError(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.EXTERNAL_ERROR + status: ClassVar[StatusCode | None] = StatusCode.EXTERNAL_ERROR class SessionPoolEmpty(Error, queue.Empty): - status: ClassVar[Optional[StatusCode]] = StatusCode.SESSION_POOL_EMPTY + status: ClassVar[StatusCode | None] = StatusCode.SESSION_POOL_EMPTY class SessionPoolClosed(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.SESSION_POOL_CLOSED + status: ClassVar[StatusCode | None] = StatusCode.SESSION_POOL_CLOSED def __init__(self): super().__init__("Session pool is closed.") class ClientInternalError(Error): - status: ClassVar[Optional[StatusCode]] = StatusCode.CLIENT_INTERNAL_ERROR + status: ClassVar[StatusCode | None] = StatusCode.CLIENT_INTERNAL_ERROR class UnexpectedGrpcMessage(Error): diff --git a/ydb/oauth2_token_exchange/token_exchange.py b/ydb/oauth2_token_exchange/token_exchange.py index 318f825d3..5735aaa55 100644 --- a/ydb/oauth2_token_exchange/token_exchange.py +++ b/ydb/oauth2_token_exchange/token_exchange.py @@ -36,11 +36,11 @@ class Oauth2TokenExchangeCredentialsBase(abc.ABC): def __init__( self, token_endpoint: str, - subject_token_source: typing.Optional[TokenSource] = None, - actor_token_source: typing.Optional[TokenSource] = None, - audience: typing.Union[typing.List[str], str, None] = None, - scope: typing.Union[typing.List[str], str, None] = None, - resource: typing.Union[typing.List[str], str, None] = None, + subject_token_source: TokenSource | None = None, + actor_token_source: TokenSource | None = None, + audience: list[str] | str | None = None, + scope: list[str] | str | None = None, + resource: list[str] | str | None = None, grant_type: str = "urn:ietf:params:oauth:grant-type:token-exchange", requested_token_type: str = "urn:ietf:params:oauth:token-type:access_token", ): @@ -83,7 +83,7 @@ def _process_response_json(self, response_json): ) return {"access_token": "Bearer " + access_token, "expires_in": expires_in} - def _get_scope_param(self) -> typing.Optional[str]: + def _get_scope_param(self) -> str | None: if self._scope is None: return None if isinstance(self._scope, str): @@ -306,11 +306,11 @@ class Oauth2TokenExchangeCredentials(credentials.AbstractExpiringTokenCredential def __init__( self, token_endpoint: str, - subject_token_source: typing.Optional[TokenSource] = None, - actor_token_source: typing.Optional[TokenSource] = None, - audience: typing.Union[typing.List[str], str, None] = None, - scope: typing.Union[typing.List[str], str, None] = None, - resource: typing.Union[typing.List[str], str, None] = None, + subject_token_source: TokenSource | None = None, + actor_token_source: TokenSource | None = None, + audience: list[str] | str | None = None, + scope: list[str] | str | None = None, + resource: list[str] | str | None = None, grant_type: str = "urn:ietf:params:oauth:grant-type:token-exchange", requested_token_type: str = "urn:ietf:params:oauth:token-type:access_token", tracer=None, diff --git a/ydb/oauth2_token_exchange/token_source.py b/ydb/oauth2_token_exchange/token_source.py index 061176ebb..1aa81f726 100644 --- a/ydb/oauth2_token_exchange/token_source.py +++ b/ydb/oauth2_token_exchange/token_source.py @@ -38,13 +38,13 @@ class JwtTokenSource(TokenSource): def __init__( self, signing_method: str, - private_key: typing.Optional[str] = None, - private_key_file: typing.Optional[str] = None, - key_id: typing.Optional[str] = None, - issuer: typing.Optional[str] = None, - subject: typing.Optional[str] = None, - audience: typing.Union[typing.List[str], str, None] = None, - id: typing.Optional[str] = None, + private_key: str | None = None, + private_key_file: str | None = None, + key_id: str | None = None, + issuer: str | None = None, + subject: str | None = None, + audience: list[str] | str | None = None, + id: str | None = None, token_ttl_seconds: int = 3600, ): assert jwt is not None, "Install pyjwt library to use jwt tokens" @@ -75,7 +75,7 @@ def token(self) -> Token: now = time.time() now_utc = datetime.utcfromtimestamp(now) exp_utc = datetime.utcfromtimestamp(now + self._token_ttl_seconds) - payload: typing.Dict[str, typing.Any] = { + payload: dict[str, typing.Any] = { "iat": now_utc, "exp": exp_utc, } diff --git a/ydb/observability/__init__.py b/ydb/observability/__init__.py index 8936f2bd4..ebe65b0de 100644 --- a/ydb/observability/__init__.py +++ b/ydb/observability/__init__.py @@ -14,8 +14,6 @@ dropped by a no-op registry. """ -from typing import List, Optional - from ydb.observability.tracing import ( NoopSpan, NoopTracingProvider, @@ -51,7 +49,7 @@ def disable_tracing() -> None: _registry.set_provider(None) -def get_active_provider() -> Optional[TracingProvider]: +def get_active_provider() -> TracingProvider | None: """Return the currently installed provider, or ``None`` if tracing is disabled.""" return _registry.get_provider() if _registry.is_active() else None @@ -70,14 +68,14 @@ def disable_metrics() -> None: _reset_metrics_provider() -def sdk_build_info_tokens() -> List[str]: +def sdk_build_info_tokens() -> list[str]: """All ``x-ydb-sdk-build-info`` feature tokens contributed by observability. Aggregated across observability features so the SDK build-info header advertises every capability the client has turned on: tracing contributes ``ydb-sdk-tracing/`` while active, metrics contribute ``ydb-sdk-metrics/``. """ - tokens: List[str] = [] + tokens: list[str] = [] tokens.extend(_tracing_build_info_tokens()) tokens.extend(_metrics_build_info_tokens()) return tokens diff --git a/ydb/observability/_endpoint.py b/ydb/observability/_endpoint.py index b88a4c09b..3120a21bf 100644 --- a/ydb/observability/_endpoint.py +++ b/ydb/observability/_endpoint.py @@ -1,9 +1,7 @@ """Shared endpoint parsing used by both tracing and metrics attributes.""" -from typing import Optional, Tuple - -def split_endpoint(endpoint: Optional[str]) -> Tuple[str, int]: +def split_endpoint(endpoint: str | None) -> tuple[str, int]: ep = endpoint or "" if ep.startswith("grpcs://"): ep = ep[len("grpcs://") :] diff --git a/ydb/observability/metrics.py b/ydb/observability/metrics.py index 95a4d321c..63feb0567 100644 --- a/ydb/observability/metrics.py +++ b/ydb/observability/metrics.py @@ -20,7 +20,7 @@ import itertools import functools import inspect -from typing import Any, Callable, Dict, Iterable, List, Optional, Protocol, Tuple +from typing import Any, Callable, Iterable, Protocol from ydb.observability._endpoint import split_endpoint @@ -95,7 +95,7 @@ } # A gauge callback returns the current ``(value, attributes)`` observations. -GaugeCallback = Callable[[], Iterable[Tuple[float, Dict[str, Any]]]] +GaugeCallback = Callable[[], Iterable[tuple[float, dict[str, Any]]]] class MetricsProvider(Protocol): @@ -105,9 +105,9 @@ class MetricsProvider(Protocol): three instrument kinds the SDK uses; see the module docstring for the semantics. """ - def record(self, name: str, value: float, attributes: Optional[Dict[str, Any]] = None) -> None: ... + def record(self, name: str, value: float, attributes: dict[str, Any] | None = None) -> None: ... - def add(self, name: str, value: int, attributes: Optional[Dict[str, Any]] = None) -> None: ... + def add(self, name: str, value: int, attributes: dict[str, Any] | None = None) -> None: ... def observe_gauge(self, name: str, callback: GaugeCallback) -> None: """Register an asynchronous gauge whose current values *callback* returns.""" @@ -133,25 +133,25 @@ def observe_gauge(self, name, callback): # Accumulated state for the asynchronous gauges, owned by the SDK (vendor-neutral) and # read by whatever provider is installed via ``observe_gauge``. _gauge_lock = threading.Lock() -_session_count_state: Dict[Tuple, int] = {} -_session_max_state: Dict[Tuple, int] = {} +_session_count_state: dict[tuple, int] = {} +_session_max_state: dict[tuple, int] = {} def is_metrics_enabled() -> bool: return _provider is not _NOOP_PROVIDER -def _observe_session_count() -> List[Tuple[float, Dict[str, Any]]]: +def _observe_session_count() -> list[tuple[float, dict[str, Any]]]: with _gauge_lock: return [(value, dict(attrs)) for attrs, value in _session_count_state.items()] -def _observe_session_max() -> List[Tuple[float, Dict[str, Any]]]: +def _observe_session_max() -> list[tuple[float, dict[str, Any]]]: with _gauge_lock: return [(value, dict(attrs)) for attrs, value in _session_max_state.items()] -def _observe_session_min() -> List[Tuple[float, Dict[str, Any]]]: +def _observe_session_min() -> list[tuple[float, dict[str, Any]]]: # The SDK never configures a pool minimum, so this is always 0 for every known pool. with _gauge_lock: return [(0, dict(attrs)) for attrs in _session_max_state] @@ -164,7 +164,7 @@ def _observe_session_min() -> List[Tuple[float, Dict[str, Any]]]: ) -def _set_metrics_provider(provider: Optional[MetricsProvider]) -> None: +def _set_metrics_provider(provider: MetricsProvider | None) -> None: global _provider _provider = provider if provider is not None else _NOOP_PROVIDER @@ -192,9 +192,9 @@ def next_query_session_pool_name() -> str: def query_session_pool_name( - name: Optional[str], - endpoint: Optional[str] = None, - database: Optional[str] = None, + name: str | None, + endpoint: str | None = None, + database: str | None = None, ) -> str: """Return a stable label for the ``ydb.query.session.pool.name`` metric attribute. @@ -217,7 +217,7 @@ def query_session_pool_name( return next_query_session_pool_name() -def _metrics_build_info_tokens() -> List[str]: +def _metrics_build_info_tokens() -> list[str]: """Metrics' contribution to the ``x-ydb-sdk-build-info`` header. Returns ``["ydb-sdk-metrics/0.1.0"]`` once a metrics backend is installed, @@ -227,11 +227,11 @@ def _metrics_build_info_tokens() -> List[str]: return [METRICS_SDK_BUILD_INFO] if is_metrics_enabled() else [] -def _pool_attrs(pool_name: Optional[str]) -> Dict[str, Any]: +def _pool_attrs(pool_name: str | None) -> dict[str, Any]: return {"ydb.query.session.pool.name": pool_name or _UNKNOWN_POOL} -def _build_ydb_metrics_attrs(driver_config) -> Dict[str, Any]: +def _build_ydb_metrics_attrs(driver_config) -> dict[str, Any]: host, port = split_endpoint(getattr(driver_config, "endpoint", None)) endpoint = "%s:%d" % (host, port) if port else host return { @@ -244,7 +244,7 @@ def _operation_name(operation_name: str) -> str: return _CLIENT_OPERATION_NAME_BY_INPUT.get(operation_name, operation_name) -def _operation_attrs(operation_name: str, attributes: Dict[str, Any]) -> Dict[str, Any]: +def _operation_attrs(operation_name: str, attributes: dict[str, Any]) -> dict[str, Any]: name = _operation_name(operation_name) return { "database": attributes.get("database", ""), @@ -269,11 +269,11 @@ class MetricsOperation: attached, and accepts only stable operation labels. """ - def __init__(self, name: str, attributes: Optional[Dict[str, Any]] = None) -> None: + def __init__(self, name: str, attributes: dict[str, Any] | None = None) -> None: self._name = name self._attributes = _operation_attrs(name, attributes or {}) self._start_time = time.monotonic() - self._exception: Optional[BaseException] = None + self._exception: BaseException | None = None self._ended = False self._end_lock = threading.Lock() @@ -366,13 +366,13 @@ def __exit__(self, exc_type, exc_val, exc_tb): _NOOP_METRICS_OPERATION = _NoopMetricsOperation() -def create_metrics_operation(name: str, attributes: Optional[Dict[str, Any]] = None): +def create_metrics_operation(name: str, attributes: dict[str, Any] | None = None): if _provider is _NOOP_PROVIDER or _operation_name(name) not in _CLIENT_OPERATION_NAMES: return _NOOP_METRICS_OPERATION return MetricsOperation(name, attributes) -def record_query_session_count(delta: int, pool_name: Optional[str] = None, state: str = "used") -> None: +def record_query_session_count(delta: int, pool_name: str | None = None, state: str = "used") -> None: if not is_metrics_enabled(): return attrs = _pool_attrs(pool_name) @@ -382,25 +382,25 @@ def record_query_session_count(delta: int, pool_name: Optional[str] = None, stat _session_count_state[key] = _session_count_state.get(key, 0) + delta -def record_query_session_create_time(duration: float, pool_name: Optional[str]) -> None: +def record_query_session_create_time(duration: float, pool_name: str | None) -> None: if not is_metrics_enabled(): return _provider.record(QUERY_SESSION_CREATE_TIME, duration, _pool_attrs(pool_name)) -def record_query_session_pending_requests(delta: int, pool_name: Optional[str]) -> None: +def record_query_session_pending_requests(delta: int, pool_name: str | None) -> None: if not is_metrics_enabled(): return _provider.add(QUERY_SESSION_PENDING_REQUESTS, delta, _pool_attrs(pool_name)) -def record_query_session_timeout(pool_name: Optional[str]) -> None: +def record_query_session_timeout(pool_name: str | None) -> None: if not is_metrics_enabled(): return _provider.add(QUERY_SESSION_TIMEOUTS, 1, _pool_attrs(pool_name)) -def record_query_session_max(value: int, pool_name: Optional[str]) -> None: +def record_query_session_max(value: int, pool_name: str | None) -> None: if not is_metrics_enabled(): return key = tuple(sorted(_pool_attrs(pool_name).items())) @@ -408,7 +408,7 @@ def record_query_session_max(value: int, pool_name: Optional[str]) -> None: _session_max_state[key] = value -def remove_query_session_pool_metrics(pool_name: Optional[str]) -> None: +def remove_query_session_pool_metrics(pool_name: str | None) -> None: if not is_metrics_enabled(): return base = list(_pool_attrs(pool_name).items()) @@ -450,7 +450,7 @@ class SessionMetrics: __slots__ = ("pool_name", "state", "_counted") def __init__(self) -> None: - self.pool_name: Optional[str] = None + self.pool_name: str | None = None self.state: str = "used" self._counted = False @@ -483,7 +483,7 @@ def count_closed(self) -> None: class _CreateTimer: __slots__ = ("_pool_name", "_start") - def __init__(self, pool_name: Optional[str]) -> None: + def __init__(self, pool_name: str | None) -> None: self._pool_name = pool_name self._start = 0.0 @@ -499,7 +499,7 @@ def __exit__(self, exc_type, exc_val, exc_tb): class _PendingTracker: __slots__ = ("_pool_name",) - def __init__(self, pool_name: Optional[str]) -> None: + def __init__(self, pool_name: str | None) -> None: self._pool_name = pool_name def __enter__(self): @@ -518,7 +518,7 @@ class QuerySessionPoolMetrics: counters live here, and stay cheap no-ops while metrics are disabled. """ - def __init__(self, name: Optional[str], driver, size: int) -> None: + def __init__(self, name: str | None, driver, size: int) -> None: driver_config = getattr(driver, "_driver_config", None) self._pool_name = query_session_pool_name( name, diff --git a/ydb/observability/tracing.py b/ydb/observability/tracing.py index 87cb8e582..4e096db28 100644 --- a/ydb/observability/tracing.py +++ b/ydb/observability/tracing.py @@ -8,7 +8,7 @@ """ import enum -from typing import Any, Callable, ContextManager, Iterable, List, Optional, Protocol, Tuple +from typing import Any, Callable, ContextManager, Iterable, Protocol from ydb.observability import metrics as _metrics from ydb.observability._endpoint import split_endpoint as _split_endpoint @@ -61,11 +61,11 @@ class TracingProvider(Protocol): def create_span( self, name: str, - attributes: Optional[dict] = None, - kind: Optional[str] = None, + attributes: dict | None = None, + kind: str | None = None, ) -> Span: ... - def get_trace_metadata(self) -> Iterable[Tuple[str, str]]: + def get_trace_metadata(self) -> Iterable[tuple[str, str]]: """Return ``(key, value)`` pairs to inject into outgoing RPC metadata.""" ... @@ -128,7 +128,7 @@ def __init__(self) -> None: def is_active(self) -> bool: return self._provider is not _NOOP_PROVIDER - def set_provider(self, provider: Optional[TracingProvider]) -> None: + def set_provider(self, provider: TracingProvider | None) -> None: self._provider = provider if provider is not None else _NOOP_PROVIDER def get_provider(self) -> TracingProvider: @@ -137,14 +137,14 @@ def get_provider(self) -> TracingProvider: def create_span(self, name, attributes=None, kind=None) -> Span: return self._provider.create_span(name, attributes, kind=kind) - def get_trace_metadata(self) -> Iterable[Tuple[str, str]]: + def get_trace_metadata(self) -> Iterable[tuple[str, str]]: return self._provider.get_trace_metadata() _registry = _TracingRegistry() -def get_trace_metadata() -> List[Tuple[str, str]]: +def get_trace_metadata() -> list[tuple[str, str]]: """Return tracing metadata for gRPC calls (empty list when no provider).""" return list(_registry.get_trace_metadata()) @@ -152,7 +152,7 @@ def get_trace_metadata() -> List[Tuple[str, str]]: TRACING_SDK_BUILD_INFO = "ydb-sdk-tracing/0.1.0" -def _tracing_build_info_tokens() -> List[str]: +def _tracing_build_info_tokens() -> list[str]: """Tracing's contribution to the ``x-ydb-sdk-build-info`` header. Returns ``["ydb-sdk-tracing/0.1.0"]`` once a provider is installed, otherwise an diff --git a/ydb/opentelemetry/metrics_plugin.py b/ydb/opentelemetry/metrics_plugin.py index e1df71d99..1025e7932 100644 --- a/ydb/opentelemetry/metrics_plugin.py +++ b/ydb/opentelemetry/metrics_plugin.py @@ -5,7 +5,7 @@ user calls :func:`ydb.opentelemetry.enable_metrics`. """ -from typing import Any, Dict, Optional +from typing import Any, Optional from opentelemetry import metrics as otel_metrics from opentelemetry.metrics import ( @@ -54,7 +54,7 @@ class OtelMetricsProvider: def __init__(self, meter: Meter) -> None: self._meter = meter - self._histograms: Dict[str, Histogram] = { + self._histograms: dict[str, Histogram] = { CLIENT_OPERATION_DURATION: _create_histogram( meter, CLIENT_OPERATION_DURATION, @@ -90,7 +90,7 @@ def __init__(self, meter: Meter) -> None: bucket_boundaries=ATTEMPT_BUCKETS, ), } - self._counters: Dict[str, Any] = { + self._counters: dict[str, Any] = { CLIENT_OPERATION_FAILED: meter.create_counter( CLIENT_OPERATION_FAILED, unit="{command}", @@ -108,12 +108,12 @@ def __init__(self, meter: Meter) -> None: ), } - def record(self, name: str, value: float, attributes: Optional[Dict[str, Any]] = None) -> None: + def record(self, name: str, value: float, attributes: Optional[dict[str, Any]] = None) -> None: instrument = self._histograms.get(name) if instrument is not None: instrument.record(value, attributes=attributes or {}) - def add(self, name: str, value: int, attributes: Optional[Dict[str, Any]] = None) -> None: + def add(self, name: str, value: int, attributes: Optional[dict[str, Any]] = None) -> None: instrument = self._counters.get(name) if instrument is not None: instrument.add(value, attributes=attributes or {}) diff --git a/ydb/opentelemetry/plugin.py b/ydb/opentelemetry/plugin.py index ccdee9961..75d8571e0 100644 --- a/ydb/opentelemetry/plugin.py +++ b/ydb/opentelemetry/plugin.py @@ -5,7 +5,7 @@ dependency is only pulled in when a user calls :func:`ydb.opentelemetry.enable_tracing`. """ -from typing import Dict, Iterable, Optional, Tuple +from typing import Iterable from opentelemetry import context as otel_context from opentelemetry import trace @@ -114,13 +114,13 @@ def create_span(self, name, attributes=None, kind=None): ) return TracingSpan(span) - def get_trace_metadata(self) -> Iterable[Tuple[str, str]]: - headers: Dict[str, str] = {} + def get_trace_metadata(self) -> Iterable[tuple[str, str]]: + headers: dict[str, str] = {} inject(headers) return tuple(headers.items()) -def _enable_tracing(tracer: Optional[object] = None) -> None: +def _enable_tracing(tracer: object | None = None) -> None: """Install an :class:`OtelTracingProvider` (idempotent replace). Called by :func:`ydb.opentelemetry.enable_tracing`. Any previously diff --git a/ydb/pool.py b/ydb/pool.py index 2ba1c33ab..4a5a506e7 100644 --- a/ydb/pool.py +++ b/ydb/pool.py @@ -7,7 +7,7 @@ from concurrent import futures import collections import random -from typing import Any, Callable, ContextManager, List, Optional, Set, Tuple, TYPE_CHECKING +from typing import Any, Callable, ContextManager, TYPE_CHECKING from . import connection as connection_impl, issues, resolver, _utilities, tracing from .observability.tracing import SpanName, create_ydb_span @@ -29,16 +29,16 @@ def __init__(self, use_all_nodes: bool = False, tracer: tracing.Tracer = tracing self.tracer = tracer self.lock = threading.RLock() self.connections: collections.OrderedDict[str, Connection] = collections.OrderedDict() - self.connections_by_node_id: collections.OrderedDict[Optional[int], Connection] = collections.OrderedDict() + self.connections_by_node_id: collections.OrderedDict[int | None, Connection] = collections.OrderedDict() self.outdated: collections.OrderedDict[str, Connection] = collections.OrderedDict() - self.subscriptions: Set["futures.Future[None]"] = set() + self.subscriptions: set["futures.Future[None]"] = set() self.preferred: collections.OrderedDict[str, Connection] = collections.OrderedDict() self.logger = logging.getLogger(__name__) self.use_all_nodes = use_all_nodes self.conn_lst_order = (self.connections,) if self.use_all_nodes else (self.preferred, self.connections) - self.fast_fail_subscriptions: Set["futures.Future[None]"] = set() + self.fast_fail_subscriptions: set["futures.Future[None]"] = set() - def add(self, connection: Optional[Connection], preferred: bool = False) -> bool: + def add(self, connection: Connection | None, preferred: bool = False) -> bool: if connection is None: return False @@ -59,7 +59,7 @@ def add(self, connection: Optional[Connection], preferred: bool = False) -> bool subscription.set_result(None) return True - def _on_done_callback(self, subscription: "futures.Future[None]") -> Optional["futures.Future[None]"]: + def _on_done_callback(self, subscription: "futures.Future[None]") -> "futures.Future[None]" | None: """ A done callback for the subscription future :param subscription: A subscription @@ -81,7 +81,7 @@ def already_exists(self, endpoint: str) -> bool: with self.lock: return endpoint in self.connections - def values(self) -> List[Connection]: + def values(self) -> list[Connection]: with self.lock: return list(self.connections.values()) @@ -103,7 +103,7 @@ def cleanup(self) -> None: for connection in actual_connections: connection.close() - def complete_discovery(self, error: Optional[Exception]) -> None: + def complete_discovery(self, error: Exception | None) -> None: with self.lock: for subscription in self.fast_fail_subscriptions: if error is None: @@ -134,7 +134,7 @@ def subscribe(self) -> "futures.Future[None]": return subscription @tracing.with_trace() - def get(self, preferred_endpoint: Optional[EndpointKey] = None) -> Connection: + def get(self, preferred_endpoint: EndpointKey | None = None) -> Connection: with self.lock: if preferred_endpoint is not None and preferred_endpoint.node_id in self.connections_by_node_id: return self.connections_by_node_id[preferred_endpoint.node_id] @@ -160,7 +160,7 @@ def remove(self, connection: Connection) -> None: self.connections.pop(connection.endpoint, None) self.outdated.pop(connection.endpoint, None) - def get_connection_by_node_id(self, node_id: Optional[int]) -> Optional[Connection]: + def get_connection_by_node_id(self, node_id: int | None) -> Connection | None: with self.lock: return self.connections_by_node_id.get(node_id) @@ -352,7 +352,7 @@ def stop(self, timeout: int = 10) -> None: pass @abstractmethod - def wait(self, timeout: Optional[float] = None, fail_fast: bool = False) -> None: + def wait(self, timeout: float | None = None, fail_fast: bool = False) -> None: """ Waits for endpoints to be are available to serve user requests :param timeout: A timeout to wait in seconds @@ -374,10 +374,10 @@ def __call__( request: Any, stub: Any, rpc_name: str, - wrap_result: Optional[Callable[..., Any]] = None, - settings: Optional["BaseRequestSettings"] = None, - wrap_args: Tuple[Any, ...] = (), - preferred_endpoint: Optional[EndpointKey] = None, + wrap_result: Callable[..., Any] | None = None, + settings: "BaseRequestSettings" | None = None, + wrap_args: tuple[Any, ...] = (), + preferred_endpoint: EndpointKey | None = None, ) -> Any: """ Sends request constructed by client library @@ -408,7 +408,7 @@ def __init__(self, driver_config: "DriverConfig") -> None: self._stopped = False self._stop_guard = threading.Lock() self._stop_event = threading.Event() - self._init_thread: Optional[threading.Thread] = None + self._init_thread: threading.Thread | None = None if driver_config.disable_discovery: # If discovery is disabled, establish the initial connection in a @@ -474,7 +474,7 @@ def async_wait(self, fail_fast: bool = False) -> "futures.Future[None]": return self._store.add_fast_fail() return self._store.subscribe() - def wait(self, timeout: Optional[float] = None, fail_fast: bool = False) -> None: + def wait(self, timeout: float | None = None, fail_fast: bool = False) -> None: """ Waits for endpoints to be are available to serve user requests @@ -522,10 +522,10 @@ def __call__( request: Any, stub: Any, rpc_name: str, - wrap_result: Optional[Callable[..., Any]] = None, - settings: Optional["BaseRequestSettings"] = None, - wrap_args: Tuple[Any, ...] = (), - preferred_endpoint: Optional[EndpointKey] = None, + wrap_result: Callable[..., Any] | None = None, + settings: "BaseRequestSettings" | None = None, + wrap_args: tuple[Any, ...] = (), + preferred_endpoint: EndpointKey | None = None, ) -> Any: """ Synchronously sends request constructed by client library @@ -571,10 +571,10 @@ def future( request: Any, stub: Any, rpc_name: str, - wrap_result: Optional[Callable[..., Any]] = None, - settings: Optional["BaseRequestSettings"] = None, - wrap_args: Tuple[Any, ...] = (), - preferred_endpoint: Optional[EndpointKey] = None, + wrap_result: Callable[..., Any] | None = None, + settings: "BaseRequestSettings" | None = None, + wrap_args: tuple[Any, ...] = (), + preferred_endpoint: EndpointKey | None = None, ) -> "futures.Future[Any]": """ Sends request constructed by client diff --git a/ydb/query/__init__.py b/ydb/query/__init__.py index 9325709d1..8d47f7737 100644 --- a/ydb/query/__init__.py +++ b/ydb/query/__init__.py @@ -20,7 +20,7 @@ ] import logging -from typing import Optional, TYPE_CHECKING +from typing import TYPE_CHECKING from .base import ( QueryClientSettings, @@ -57,7 +57,7 @@ class QueryClientSync: _driver: "SyncDriver" - def __init__(self, driver: "SyncDriver", query_client_settings: Optional[QueryClientSettings] = None): + def __init__(self, driver: "SyncDriver", query_client_settings: QueryClientSettings | None = None): self._driver = driver self._settings = query_client_settings diff --git a/ydb/query/base.py b/ydb/query/base.py index 12752fe23..b03563d8e 100644 --- a/ydb/query/base.py +++ b/ydb/query/base.py @@ -8,9 +8,6 @@ Optional, Any, Callable, - List, - DefaultDict, - Union, ) from .._grpc.grpcwrapper import ydb_query @@ -154,18 +151,18 @@ def with_native_datetime_in_result_sets(self, enabled: bool) -> "QueryClientSett def create_execute_query_request( query: str, session_id: str, - tx_id: Optional[str], - commit_tx: Optional[bool], - tx_mode: Optional[BaseQueryTxMode], - syntax: Optional[QuerySyntax], - exec_mode: Optional[QueryExecMode], - stats_mode: Optional[QueryStatsMode], - schema_inclusion_mode: Optional[QuerySchemaInclusionMode], - result_set_format: Optional[QueryResultSetFormat], - arrow_format_settings: Optional[ArrowFormatSettings], - parameters: Optional[dict], - concurrent_result_sets: Optional[bool], - pool_id: Optional[str], + tx_id: str | None, + commit_tx: bool | None, + tx_mode: BaseQueryTxMode | None, + syntax: QuerySyntax | None, + exec_mode: QueryExecMode | None, + stats_mode: QueryStatsMode | None, + schema_inclusion_mode: QuerySchemaInclusionMode | None, + result_set_format: QueryResultSetFormat | None, + arrow_format_settings: ArrowFormatSettings | None, + parameters: dict | None, + concurrent_result_sets: bool | None, + pool_id: str | None, ) -> ydb_query.ExecuteQueryRequest: try: syntax = QuerySyntax.YQL_V1 if not syntax else syntax @@ -232,9 +229,9 @@ def wrap_execute_query_response( response_pb: _apis.ydb_query.ExecuteQueryResponsePart, session: "BaseQuerySession", tx: Optional["BaseQueryTxContext"] = None, - commit_tx: Optional[bool] = False, - settings: Optional[QueryClientSettings] = None, -) -> Optional[convert.ResultSet]: + commit_tx: bool | None = False, + settings: QueryClientSettings | None = None, +) -> convert.ResultSet | None: issues._process_response(response_pb) if tx and commit_tx: tx._move_to_commited() @@ -269,7 +266,7 @@ class CallbackHandlerMode(enum.Enum): ASYNC = "ASYNC" -def _get_sync_callback(method: typing.Callable, loop: Optional[asyncio.AbstractEventLoop]): +def _get_sync_callback(method: typing.Callable, loop: asyncio.AbstractEventLoop | None): if asyncio.iscoroutinefunction(method): if loop is None: loop = _get_shared_event_loop() @@ -293,19 +290,19 @@ async def sync_to_async_callback(*args, **kwargs): class CallbackHandler: - _callbacks: DefaultDict[str, List[Callable[..., Any]]] + _callbacks: defaultdict[str, list[Callable[..., Any]]] _callback_mode: CallbackHandlerMode def _init_callback_handler(self, mode: CallbackHandlerMode) -> None: self._callbacks = defaultdict(list) self._callback_mode = mode - def _execute_callbacks_sync(self, event_name: Union[str, TxEvent], *args: Any, **kwargs: Any) -> None: + def _execute_callbacks_sync(self, event_name: str | TxEvent, *args: Any, **kwargs: Any) -> None: key = event_name.value if isinstance(event_name, TxEvent) else event_name for callback in self._callbacks[key]: callback(self, *args, **kwargs) - async def _execute_callbacks_async(self, event_name: Union[str, TxEvent], *args: Any, **kwargs: Any) -> None: + async def _execute_callbacks_async(self, event_name: str | TxEvent, *args: Any, **kwargs: Any) -> None: key = event_name.value if isinstance(event_name, TxEvent) else event_name tasks = [asyncio.create_task(callback(self, *args, **kwargs)) for callback in self._callbacks[key]] if not tasks: @@ -313,7 +310,7 @@ async def _execute_callbacks_async(self, event_name: Union[str, TxEvent], *args: await asyncio.gather(*tasks) def _prepare_callback( - self, callback: typing.Callable[..., Any], loop: Optional[asyncio.AbstractEventLoop] + self, callback: typing.Callable[..., Any], loop: asyncio.AbstractEventLoop | None ) -> typing.Callable[..., Any]: if self._callback_mode == CallbackHandlerMode.SYNC: return _get_sync_callback(callback, loop) @@ -321,9 +318,9 @@ def _prepare_callback( def _add_callback( self, - event_name: Union[str, TxEvent], + event_name: str | TxEvent, callback: typing.Callable[..., Any], - loop: Optional[asyncio.AbstractEventLoop], + loop: asyncio.AbstractEventLoop | None, ) -> None: key = event_name.value if isinstance(event_name, TxEvent) else event_name self._callbacks[key].append(self._prepare_callback(callback, loop)) diff --git a/ydb/query/pool.py b/ydb/query/pool.py index 79f5051a7..95053aca2 100644 --- a/ydb/query/pool.py +++ b/ydb/query/pool.py @@ -4,11 +4,7 @@ from concurrent import futures from typing import ( Callable, - Optional, - List, - Dict, Any, - Union, TYPE_CHECKING, ) import time @@ -46,9 +42,9 @@ def __init__( driver: "SyncDriver", size: int = 100, *, - query_client_settings: Optional[QueryClientSettings] = None, + query_client_settings: QueryClientSettings | None = None, workers_threads_count: int = 4, - name: Optional[str] = None, + name: str | None = None, ): """ :param driver: A driver instance. @@ -68,7 +64,7 @@ def __init__( self._query_client_settings = query_client_settings self._metrics = QuerySessionPoolMetrics(name, driver, self._size) - def _create_new_session(self, timeout: Optional[float]): + def _create_new_session(self, timeout: float | None): session = QuerySession(self._driver, settings=self._query_client_settings) self._metrics.attach(session) with self._metrics.measure_create(): @@ -76,7 +72,7 @@ def _create_new_session(self, timeout: Optional[float]): logger.debug(f"New session was created for pool. Session id: {session.session_id}") return session - def acquire(self, timeout: Optional[float] = None) -> QuerySession: + def acquire(self, timeout: float | None = None) -> QuerySession: """Acquire a session from Session Pool. :param timeout: Seconds to wait when pool is exhausted. Overrides the pool-level acquire_timeout. @@ -138,7 +134,7 @@ def release(self, session: QuerySession) -> None: self._queue.put_nowait(session) logger.debug("Session returned to queue: %s", session.session_id) - def checkout(self, timeout: Optional[float] = None) -> "SimpleQuerySessionCheckout": + def checkout(self, timeout: float | None = None) -> "SimpleQuerySessionCheckout": """Return a Session context manager, that acquires session on enter and releases session on exit. :param timeout: A timeout to wait in seconds. @@ -146,7 +142,7 @@ def checkout(self, timeout: Optional[float] = None) -> "SimpleQuerySessionChecko return SimpleQuerySessionCheckout(self, timeout) - def retry_operation_sync(self, callee: Callable, retry_settings: Optional[RetrySettings] = None, *args, **kwargs): + def retry_operation_sync(self, callee: Callable, retry_settings: RetrySettings | None = None, *args, **kwargs): """Special interface to execute a bunch of commands with session in a safe, retriable way. :param callee: A function, that works with session. @@ -169,8 +165,8 @@ def wrapped_callee(): def retry_tx_async( self, callee: Callable, - tx_mode: Optional[BaseQueryTxMode] = None, - retry_settings: Optional[RetrySettings] = None, + tx_mode: BaseQueryTxMode | None = None, + retry_settings: RetrySettings | None = None, *args, **kwargs, ) -> futures.Future: @@ -189,7 +185,7 @@ def retry_tx_async( ) def retry_operation_async( - self, callee: Callable, retry_settings: Optional[RetrySettings] = None, *args, **kwargs + self, callee: Callable, retry_settings: RetrySettings | None = None, *args, **kwargs ) -> futures.Future: """Asynchronously execute a retryable operation.""" @@ -201,8 +197,8 @@ def retry_operation_async( def retry_tx_sync( self, callee: Callable, - tx_mode: Optional[BaseQueryTxMode] = None, - retry_settings: Optional[RetrySettings] = None, + tx_mode: BaseQueryTxMode | None = None, + retry_settings: RetrySettings | None = None, *args, **kwargs, ): @@ -240,12 +236,12 @@ def wrapped_callee(): def execute_with_retries( self, query: str, - parameters: Optional[dict] = None, - retry_settings: Optional[RetrySettings] = None, + parameters: dict | None = None, + retry_settings: RetrySettings | None = None, *args, - pool_id: Optional[str] = None, + pool_id: str | None = None, **kwargs, - ) -> List[convert.ResultSet]: + ) -> list[convert.ResultSet]: """Special interface to execute a one-shot queries in a safe, retriable way. Note: this method loads all data from stream before return, do not use this method with huge read queries. @@ -273,10 +269,10 @@ def wrapped_callee(): def execute_with_retries_async( self, query: str, - parameters: Optional[dict] = None, - retry_settings: Optional[RetrySettings] = None, + parameters: dict | None = None, + retry_settings: RetrySettings | None = None, *args, - pool_id: Optional[str] = None, + pool_id: str | None = None, **kwargs, ) -> futures.Future: """Asynchronously execute a query with retries.""" @@ -297,11 +293,11 @@ def execute_with_retries_async( def explain_with_retries( self, query: str, - parameters: Optional[dict] = None, + parameters: dict | None = None, *, result_format: QueryExplainResultFormat = QueryExplainResultFormat.STR, - retry_settings: Optional[RetrySettings] = None, - ) -> Union[str, Dict[str, Any]]: + retry_settings: RetrySettings | None = None, + ) -> str | dict[str, Any]: """ Explain a query in retriable way. No real query execution will happen. @@ -344,9 +340,9 @@ def __exit__(self, exc_type, exc_val, exc_tb): class SimpleQuerySessionCheckout: - _session: Optional[QuerySession] + _session: QuerySession | None - def __init__(self, pool: QuerySessionPool, timeout: Optional[float]): + def __init__(self, pool: QuerySessionPool, timeout: float | None): self._pool = pool self._timeout = timeout self._session = None diff --git a/ydb/query/session.py b/ydb/query/session.py index b8a7d82e1..df358db75 100644 --- a/ydb/query/session.py +++ b/ydb/query/session.py @@ -7,10 +7,8 @@ Generic, Iterable, Optional, - Dict, Any, TYPE_CHECKING, - Union, overload, ) @@ -87,17 +85,17 @@ class BaseQuerySession(abc.ABC, Generic[DriverT]): _driver: DriverT _settings: base.QueryClientSettings - _stream: Optional[GrpcStreamCall[_apis.ydb_query.SessionState]] = None + _stream: GrpcStreamCall[_apis.ydb_query.SessionState] | None = None # Session data - _session_id: Optional[str] = None - _node_id: Optional[int] = None - _peer: Optional[tuple] = None + _session_id: str | None = None + _node_id: int | None = None + _peer: tuple | None = None _closed: bool = False _invalidated: bool = False _session_metrics: SessionMetrics = _NOOP_SESSION_METRICS - def __init__(self, driver: DriverT, settings: Optional[base.QueryClientSettings] = None): + def __init__(self, driver: DriverT, settings: base.QueryClientSettings | None = None): self._driver = driver self._settings = self._get_client_settings(driver, settings) self._attach_settings: BaseRequestSettings = ( @@ -115,11 +113,11 @@ def _driver_config(self) -> Optional["DriverConfig"]: return getattr(self._driver, "_driver_config", None) @property - def session_id(self) -> Optional[str]: + def session_id(self) -> str | None: return self._session_id @property - def node_id(self) -> Optional[int]: + def node_id(self) -> int | None: return self._node_id @property @@ -127,7 +125,7 @@ def is_active(self) -> bool: return self._session_id is not None and not self._closed @property - def _endpoint_key(self) -> Optional[EndpointKey]: + def _endpoint_key(self) -> EndpointKey | None: if self._node_id is None: return None return EndpointKey(endpoint=None, node_id=self._node_id) @@ -143,7 +141,7 @@ def last_query_stats(self): def _get_client_settings( self, driver: SupportedDriverType, - settings: Optional[base.QueryClientSettings] = None, + settings: base.QueryClientSettings | None = None, ) -> base.QueryClientSettings: if settings is not None: return settings @@ -223,17 +221,17 @@ def _on_execute_stream_error(self, e: BaseException) -> None: # Overloads for _create_call @overload def _create_call( - self: "BaseQuerySession[SyncDriver]", settings: Optional[BaseRequestSettings] = None + self: "BaseQuerySession[SyncDriver]", settings: BaseRequestSettings | None = None ) -> "BaseQuerySession[SyncDriver]": ... @overload def _create_call( - self: "BaseQuerySession[AsyncDriver]", settings: Optional[BaseRequestSettings] = None + self: "BaseQuerySession[AsyncDriver]", settings: BaseRequestSettings | None = None ) -> Awaitable["BaseQuerySession[AsyncDriver]"]: ... def _create_call( - self, settings: Optional[BaseRequestSettings] = None - ) -> "Union[BaseQuerySession[Any], Awaitable[BaseQuerySession[Any]]]": + self, settings: BaseRequestSettings | None = None + ) -> "BaseQuerySession[Any] | Awaitable[BaseQuerySession[Any]]": """Create session. Returns Awaitable in async context.""" return self._driver( _apis.ydb_query.CreateSessionRequest(), @@ -247,17 +245,17 @@ def _create_call( # Overloads for _delete_call @overload def _delete_call( - self: "BaseQuerySession[SyncDriver]", settings: Optional[BaseRequestSettings] = None + self: "BaseQuerySession[SyncDriver]", settings: BaseRequestSettings | None = None ) -> "BaseQuerySession[SyncDriver]": ... @overload def _delete_call( - self: "BaseQuerySession[AsyncDriver]", settings: Optional[BaseRequestSettings] = None + self: "BaseQuerySession[AsyncDriver]", settings: BaseRequestSettings | None = None ) -> Awaitable["BaseQuerySession[AsyncDriver]"]: ... def _delete_call( - self, settings: Optional[BaseRequestSettings] = None - ) -> "Union[BaseQuerySession[Any], Awaitable[BaseQuerySession[Any]]]": + self, settings: BaseRequestSettings | None = None + ) -> "BaseQuerySession[Any] | Awaitable[BaseQuerySession[Any]]": """Delete session. Returns Awaitable in async context.""" return self._driver( _apis.ydb_query.DeleteSessionRequest(session_id=self._session_id), @@ -282,7 +280,7 @@ def _attach_call( def _attach_call( self, - ) -> Union[GrpcStreamCall[_apis.ydb_query.SessionState], Awaitable[GrpcStreamCall[_apis.ydb_query.SessionState]]]: + ) -> GrpcStreamCall[_apis.ydb_query.SessionState] | Awaitable[GrpcStreamCall[_apis.ydb_query.SessionState]]: """Attach to session. Returns Awaitable in async context.""" return self._driver( _apis.ydb_query.AttachSessionRequest(session_id=self._session_id), @@ -297,54 +295,54 @@ def _attach_call( def _execute_call( self: "BaseQuerySession[SyncDriver]", query: str, - parameters: Optional[dict] = None, + parameters: dict | None = None, commit_tx: bool = False, - syntax: Optional[base.QuerySyntax] = None, - exec_mode: Optional[base.QueryExecMode] = None, - stats_mode: Optional[base.QueryStatsMode] = None, - schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode] = None, - result_set_format: Optional[base.QueryResultSetFormat] = None, - arrow_format_settings: Optional[base.ArrowFormatSettings] = None, + syntax: base.QuerySyntax | None = None, + exec_mode: base.QueryExecMode | None = None, + stats_mode: base.QueryStatsMode | None = None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None = None, + result_set_format: base.QueryResultSetFormat | None = None, + arrow_format_settings: base.ArrowFormatSettings | None = None, concurrent_result_sets: bool = False, - settings: Optional[BaseRequestSettings] = None, - pool_id: Optional[str] = None, + settings: BaseRequestSettings | None = None, + pool_id: str | None = None, ) -> Iterable[_apis.ydb_query.ExecuteQueryResponsePart]: ... @overload def _execute_call( self: "BaseQuerySession[AsyncDriver]", query: str, - parameters: Optional[dict] = None, + parameters: dict | None = None, commit_tx: bool = False, - syntax: Optional[base.QuerySyntax] = None, - exec_mode: Optional[base.QueryExecMode] = None, - stats_mode: Optional[base.QueryStatsMode] = None, - schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode] = None, - result_set_format: Optional[base.QueryResultSetFormat] = None, - arrow_format_settings: Optional[base.ArrowFormatSettings] = None, + syntax: base.QuerySyntax | None = None, + exec_mode: base.QueryExecMode | None = None, + stats_mode: base.QueryStatsMode | None = None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None = None, + result_set_format: base.QueryResultSetFormat | None = None, + arrow_format_settings: base.ArrowFormatSettings | None = None, concurrent_result_sets: bool = False, - settings: Optional[BaseRequestSettings] = None, - pool_id: Optional[str] = None, + settings: BaseRequestSettings | None = None, + pool_id: str | None = None, ) -> Awaitable[Iterable[_apis.ydb_query.ExecuteQueryResponsePart]]: ... def _execute_call( self, query: str, - parameters: Optional[dict] = None, + parameters: dict | None = None, commit_tx: bool = False, - syntax: Optional[base.QuerySyntax] = None, - exec_mode: Optional[base.QueryExecMode] = None, - stats_mode: Optional[base.QueryStatsMode] = None, - schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode] = None, - result_set_format: Optional[base.QueryResultSetFormat] = None, - arrow_format_settings: Optional[base.ArrowFormatSettings] = None, + syntax: base.QuerySyntax | None = None, + exec_mode: base.QueryExecMode | None = None, + stats_mode: base.QueryStatsMode | None = None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None = None, + result_set_format: base.QueryResultSetFormat | None = None, + arrow_format_settings: base.ArrowFormatSettings | None = 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]], - ]: + settings: BaseRequestSettings | None = None, + pool_id: str | None = None, + ) -> ( + Iterable[_apis.ydb_query.ExecuteQueryResponsePart] + | Awaitable[Iterable[_apis.ydb_query.ExecuteQueryResponsePart]] + ): self._last_query_stats = None if self._session_id is None: @@ -381,7 +379,7 @@ class QuerySession(BaseQuerySession["SyncDriver"]): session's lifecycle manually - use a QuerySessionPool is always a better choice. """ - def __init__(self, driver: "SyncDriver", settings: Optional[base.QueryClientSettings] = None): + def __init__(self, driver: "SyncDriver", settings: base.QueryClientSettings | None = None): super().__init__(driver, settings) def _attach(self, first_resp_timeout: int = DEFAULT_INITIAL_RESPONSE_TIMEOUT) -> None: @@ -417,7 +415,7 @@ def _check_session_status_loop(self, status_stream: _utilities.SyncResponseItera logger.debug("Attach stream error: %s, session_id: %s", e, self._session_id) self._close_session(invalidate=True) - def delete(self, settings: Optional[BaseRequestSettings] = None) -> None: + def delete(self, settings: BaseRequestSettings | None = None) -> None: """Deletes a Session of Query Service on server side and releases resources. :return: None @@ -433,7 +431,7 @@ def delete(self, settings: Optional[BaseRequestSettings] = None) -> None: self._close_session() - def create(self, settings: Optional[BaseRequestSettings] = None) -> "QuerySession": + def create(self, settings: BaseRequestSettings | None = None) -> "QuerySession": """Creates a Session of Query Service on server side and attaches it. :return: QuerySession object. @@ -452,7 +450,7 @@ def create(self, settings: Optional[BaseRequestSettings] = None) -> "QuerySessio return self - def transaction(self, tx_mode: Optional[base.BaseQueryTxMode] = None) -> QueryTxContext: + def transaction(self, tx_mode: base.BaseQueryTxMode | None = None) -> QueryTxContext: """Creates a transaction context manager with specified transaction mode. :param tx_mode: Transaction mode, which is a one from the following choices: @@ -482,13 +480,13 @@ def execute( syntax: base.QuerySyntax = None, exec_mode: base.QueryExecMode = None, concurrent_result_sets: bool = False, - settings: Optional[BaseRequestSettings] = None, + settings: BaseRequestSettings | None = None, *, - stats_mode: Optional[base.QueryStatsMode] = None, - 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, + stats_mode: base.QueryStatsMode | None = None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None = None, + result_set_format: base.QueryResultSetFormat | None = None, + arrow_format_settings: base.ArrowFormatSettings | None = None, + pool_id: str | None = None, ) -> base.SyncResponseContextIterator: """Sends a query to Query Service @@ -556,7 +554,7 @@ def explain( parameters: dict = None, *, result_format: QueryExplainResultFormat = QueryExplainResultFormat.STR, - ) -> Union[str, Dict[str, Any]]: + ) -> str | dict[str, Any]: """Explains query result :param query: YQL or SQL query. :param parameters: dict with parameters and YDB types; diff --git a/ydb/query/transaction.py b/ydb/query/transaction.py index 692b2a4c4..977f5f32d 100644 --- a/ydb/query/transaction.py +++ b/ydb/query/transaction.py @@ -7,9 +7,7 @@ Awaitable, Generic, Iterable, - Optional, TYPE_CHECKING, - Union, overload, ) @@ -91,7 +89,7 @@ def decorator(rpc_state, response_pb, session: "BaseQuerySession", tx_state: "Qu class QueryTxState: - tx_id: Optional[str] + tx_id: str | None tx_mode: base.BaseQueryTxMode _state: QueryTxStateEnum @@ -216,7 +214,7 @@ class BaseQueryTxContext(base.CallbackHandler, Generic[DriverT]): _driver: DriverT _prev_stream: Any # SyncResponseContextIterator or AsyncResponseContextIterator - _external_error: Optional[BaseException] + _external_error: BaseException | None def __init__(self, driver: DriverT, session: "BaseQuerySession", tx_mode: base.BaseQueryTxMode): """ @@ -250,7 +248,7 @@ def _driver_config(self): return getattr(self._driver, "_driver_config", None) @property - def session_id(self) -> Optional[str]: + def session_id(self) -> str | None: """ A transaction's session id @@ -259,7 +257,7 @@ def session_id(self) -> Optional[str]: return self.session.session_id @property - def tx_id(self) -> Optional[str]: + def tx_id(self) -> str | None: """ Returns an id of open transaction or None otherwise @@ -289,17 +287,17 @@ def _check_external_error_set(self): # Overloads for _begin_call - sync driver returns value, async driver returns Awaitable @overload def _begin_call( - self: "BaseQueryTxContext[SyncDriver]", settings: Optional[BaseRequestSettings] + self: "BaseQueryTxContext[SyncDriver]", settings: BaseRequestSettings | None ) -> "BaseQueryTxContext[SyncDriver]": ... @overload def _begin_call( - self: "BaseQueryTxContext[AsyncDriver]", settings: Optional[BaseRequestSettings] + self: "BaseQueryTxContext[AsyncDriver]", settings: BaseRequestSettings | None ) -> Awaitable["BaseQueryTxContext[AsyncDriver]"]: ... def _begin_call( - self, settings: Optional[BaseRequestSettings] - ) -> "Union[BaseQueryTxContext[Any], Awaitable[BaseQueryTxContext[Any]]]": + self, settings: BaseRequestSettings | None + ) -> "BaseQueryTxContext[Any] | Awaitable[BaseQueryTxContext[Any]]": """Begin transaction. Returns Awaitable in async context.""" self._tx_state._check_invalid_transition(QueryTxStateEnum.BEGINED) @@ -316,17 +314,17 @@ def _begin_call( # Overloads for _commit_call @overload def _commit_call( - self: "BaseQueryTxContext[SyncDriver]", settings: Optional[BaseRequestSettings] + self: "BaseQueryTxContext[SyncDriver]", settings: BaseRequestSettings | None ) -> "BaseQueryTxContext[SyncDriver]": ... @overload def _commit_call( - self: "BaseQueryTxContext[AsyncDriver]", settings: Optional[BaseRequestSettings] + self: "BaseQueryTxContext[AsyncDriver]", settings: BaseRequestSettings | None ) -> Awaitable["BaseQueryTxContext[AsyncDriver]"]: ... def _commit_call( - self, settings: Optional[BaseRequestSettings] - ) -> "Union[BaseQueryTxContext[Any], Awaitable[BaseQueryTxContext[Any]]]": + self, settings: BaseRequestSettings | None + ) -> "BaseQueryTxContext[Any] | Awaitable[BaseQueryTxContext[Any]]": """Commit transaction. Returns Awaitable in async context.""" self._check_external_error_set() self._tx_state._check_invalid_transition(QueryTxStateEnum.COMMITTED) @@ -344,17 +342,17 @@ def _commit_call( # Overloads for _rollback_call @overload def _rollback_call( - self: "BaseQueryTxContext[SyncDriver]", settings: Optional[BaseRequestSettings] + self: "BaseQueryTxContext[SyncDriver]", settings: BaseRequestSettings | None ) -> "BaseQueryTxContext[SyncDriver]": ... @overload def _rollback_call( - self: "BaseQueryTxContext[AsyncDriver]", settings: Optional[BaseRequestSettings] + self: "BaseQueryTxContext[AsyncDriver]", settings: BaseRequestSettings | None ) -> Awaitable["BaseQueryTxContext[AsyncDriver]"]: ... def _rollback_call( - self, settings: Optional[BaseRequestSettings] - ) -> "Union[BaseQueryTxContext[Any], Awaitable[BaseQueryTxContext[Any]]]": + self, settings: BaseRequestSettings | None + ) -> "BaseQueryTxContext[Any] | Awaitable[BaseQueryTxContext[Any]]": """Rollback transaction. Returns Awaitable in async context.""" self._check_external_error_set() self._tx_state._check_invalid_transition(QueryTxStateEnum.ROLLBACKED) @@ -374,54 +372,54 @@ def _rollback_call( def _execute_call( self: "BaseQueryTxContext[SyncDriver]", query: str, - parameters: Optional[dict], - commit_tx: Optional[bool], - syntax: Optional[base.QuerySyntax], - exec_mode: Optional[base.QueryExecMode], - stats_mode: Optional[base.QueryStatsMode], - schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode], - result_set_format: Optional[base.QueryResultSetFormat], - arrow_format_settings: Optional[base.ArrowFormatSettings], - concurrent_result_sets: Optional[bool], - settings: Optional[BaseRequestSettings], - pool_id: Optional[str], + parameters: dict | None, + commit_tx: bool | None, + syntax: base.QuerySyntax | None, + exec_mode: base.QueryExecMode | None, + stats_mode: base.QueryStatsMode | None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None, + result_set_format: base.QueryResultSetFormat | None, + arrow_format_settings: base.ArrowFormatSettings | None, + concurrent_result_sets: bool | None, + settings: BaseRequestSettings | None, + pool_id: str | None, ) -> Iterable[_apis.ydb_query.ExecuteQueryResponsePart]: ... @overload def _execute_call( self: "BaseQueryTxContext[AsyncDriver]", query: str, - parameters: Optional[dict], - commit_tx: Optional[bool], - syntax: Optional[base.QuerySyntax], - exec_mode: Optional[base.QueryExecMode], - stats_mode: Optional[base.QueryStatsMode], - schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode], - result_set_format: Optional[base.QueryResultSetFormat], - arrow_format_settings: Optional[base.ArrowFormatSettings], - concurrent_result_sets: Optional[bool], - settings: Optional[BaseRequestSettings], - pool_id: Optional[str], + parameters: dict | None, + commit_tx: bool | None, + syntax: base.QuerySyntax | None, + exec_mode: base.QueryExecMode | None, + stats_mode: base.QueryStatsMode | None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None, + result_set_format: base.QueryResultSetFormat | None, + arrow_format_settings: base.ArrowFormatSettings | None, + concurrent_result_sets: bool | None, + settings: BaseRequestSettings | None, + pool_id: str | None, ) -> Awaitable[Iterable[_apis.ydb_query.ExecuteQueryResponsePart]]: ... def _execute_call( self, query: str, - parameters: Optional[dict], - commit_tx: Optional[bool], - syntax: Optional[base.QuerySyntax], - exec_mode: Optional[base.QueryExecMode], - stats_mode: Optional[base.QueryStatsMode], - schema_inclusion_mode: Optional[base.QuerySchemaInclusionMode], - result_set_format: Optional[base.QueryResultSetFormat], - 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]], - ]: + parameters: dict | None, + commit_tx: bool | None, + syntax: base.QuerySyntax | None, + exec_mode: base.QueryExecMode | None, + stats_mode: base.QueryStatsMode | None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None, + result_set_format: base.QueryResultSetFormat | None, + arrow_format_settings: base.ArrowFormatSettings | None, + concurrent_result_sets: bool | None, + settings: BaseRequestSettings | None, + pool_id: str | None, + ) -> ( + Iterable[_apis.ydb_query.ExecuteQueryResponsePart] + | Awaitable[Iterable[_apis.ydb_query.ExecuteQueryResponsePart]] + ): self._tx_state._check_tx_ready_to_use() self._check_external_error_set() @@ -525,7 +523,7 @@ def _ensure_prev_stream_finished(self) -> None: pass self._prev_stream = None - def begin(self, settings: Optional[BaseRequestSettings] = None) -> "QueryTxContext": + def begin(self, settings: BaseRequestSettings | None = None) -> "QueryTxContext": """Explicitly begins a transaction :param settings: An additional request settings BaseRequestSettings; @@ -542,7 +540,7 @@ def begin(self, settings: Optional[BaseRequestSettings] = None) -> "QueryTxConte return self - def commit(self, settings: Optional[BaseRequestSettings] = None) -> None: + def commit(self, settings: BaseRequestSettings | None = None) -> None: """Calls commit on a transaction if it is open otherwise is no-op. If transaction execution failed then this method raises PreconditionFailed. @@ -574,7 +572,7 @@ def commit(self, settings: Optional[BaseRequestSettings] = None) -> None: self._execute_callbacks_sync(base.TxEvent.AFTER_COMMIT, exc=e) raise e - def rollback(self, settings: Optional[BaseRequestSettings] = None) -> None: + def rollback(self, settings: BaseRequestSettings | None = None) -> None: """Calls rollback on a transaction if it is open otherwise is no-op. If transaction execution failed then this method raises PreconditionFailed. @@ -609,18 +607,18 @@ def rollback(self, settings: Optional[BaseRequestSettings] = None) -> None: def execute( self, query: str, - parameters: Optional[dict] = None, - commit_tx: Optional[bool] = False, - syntax: Optional[base.QuerySyntax] = None, - exec_mode: Optional[base.QueryExecMode] = None, - concurrent_result_sets: Optional[bool] = False, - settings: Optional[BaseRequestSettings] = None, + parameters: dict | None = None, + commit_tx: bool | None = False, + syntax: base.QuerySyntax | None = None, + exec_mode: base.QueryExecMode | None = None, + concurrent_result_sets: bool | None = False, + settings: BaseRequestSettings | None = None, *, - stats_mode: Optional[base.QueryStatsMode] = None, - 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, + stats_mode: base.QueryStatsMode | None = None, + schema_inclusion_mode: base.QuerySchemaInclusionMode | None = None, + result_set_format: base.QueryResultSetFormat | None = None, + arrow_format_settings: base.ArrowFormatSettings | None = None, + pool_id: str | None = None, ) -> base.SyncResponseContextIterator: """Sends a query to Query Service diff --git a/ydb/resolver.py b/ydb/resolver.py index d55de3895..30cf42f5b 100644 --- a/ydb/resolver.py +++ b/ydb/resolver.py @@ -7,7 +7,7 @@ import random import itertools import typing -from typing import Any, ContextManager, List, Optional, Iterator +from typing import Any, ContextManager, Iterator from . import connection as conn_impl, driver, issues, settings as settings_impl, _apis @@ -21,7 +21,7 @@ logger = logging.getLogger(__name__) -class EndpointInfo(object): +class EndpointInfo: __slots__ = ( "address", "endpoint", @@ -45,7 +45,7 @@ def __init__(self, endpoint_info: ydb_discovery_pb2.EndpointInfo): self.ssl_target_name_override = endpoint_info.ssl_target_name_override self.node_id = endpoint_info.node_id - def endpoints_with_options(self) -> typing.Generator[typing.Tuple[str, conn_impl.EndpointOptions], None, None]: + def endpoints_with_options(self) -> typing.Generator[tuple[str, conn_impl.EndpointOptions], None, None]: ssl_target_name_override = None if self.ssl: if self.ssl_target_name_override: @@ -95,7 +95,7 @@ def _list_endpoints_request_factory(connection_params: driver.DriverConfig) -> _ return request -class DiscoveryResult(object): +class DiscoveryResult: def __init__(self, self_location: str, endpoints: "list[EndpointInfo]"): self.self_location = self_location self.endpoints = endpoints @@ -127,7 +127,7 @@ def from_response( else: unique_different_set.add(EndpointInfo(info)) - result: List[EndpointInfo] = [] + result: list[EndpointInfo] = [] local_endpoints = list(unique_local_set) different_endpoints = list(unique_different_set) if use_all_nodes: @@ -143,7 +143,7 @@ def from_response( return cls(message.self_location, result) -class DiscoveryEndpointsResolver(object): +class DiscoveryEndpointsResolver: _lock: ContextManager[Any] # Can be threading.Lock or _FakeLock in async subclass def __init__(self, driver_config: driver.DriverConfig): @@ -152,7 +152,7 @@ def __init__(self, driver_config: driver.DriverConfig): self._ready_timeout = getattr(self._driver_config, "discovery_request_timeout", 10) self._lock = threading.Lock() self._debug_details_history_size = 20 - self._debug_details_items: List[str] = [] + self._debug_details_items: list[str] = [] self._endpoints = [] self._endpoints.append(driver_config.endpoint) self._endpoints.extend(driver_config.endpoints) @@ -174,12 +174,12 @@ def debug_details(self) -> str: with self._lock: return "\n".join(self._debug_details_items) - def resolve(self) -> Optional[DiscoveryResult]: + def resolve(self) -> DiscoveryResult | None: with self.context_resolve() as result: return result @contextlib.contextmanager - def context_resolve(self) -> Iterator[Optional[DiscoveryResult]]: + def context_resolve(self) -> Iterator[DiscoveryResult | None]: self.logger.debug("Preparing initial endpoint to resolve endpoints") endpoint = next(self._endpoints_iter) initial = conn_impl.Connection.ready_factory(endpoint, self._driver_config, ready_timeout=self._ready_timeout) diff --git a/ydb/retries.py b/ydb/retries.py index 4e352b5dc..fb8b9288e 100644 --- a/ydb/retries.py +++ b/ydb/retries.py @@ -3,7 +3,7 @@ import inspect import random import time -from typing import Any, Callable, Generator, Optional, Union +from typing import Any, Callable, Generator from . import issues from ._errors import check_retriable_error @@ -11,7 +11,7 @@ from .observability.tracing import SpanName, create_span as _create_span -def _try_span_attrs(backoff_ms: Optional[int]): +def _try_span_attrs(backoff_ms: int | None): return {"ydb.retry.backoff_ms": backoff_ms} if backoff_ms is not None else None @@ -38,13 +38,13 @@ class RetrySettings: def __init__( self, max_retries: int = 10, - max_session_acquire_timeout: Optional[float] = None, - on_ydb_error_callback: Optional[Callable[[issues.Error], None]] = None, + max_session_acquire_timeout: float | None = None, + on_ydb_error_callback: Callable[[issues.Error], None] | None = None, backoff_ceiling: int = 6, backoff_slot_duration: float = 1, get_session_client_timeout: float = 5, - fast_backoff_settings: Optional[BackoffSettings] = None, - slow_backoff_settings: Optional[BackoffSettings] = None, + fast_backoff_settings: BackoffSettings | None = None, + slow_backoff_settings: BackoffSettings | None = None, idempotent: bool = False, retry_cancelled: bool = False, ) -> None: @@ -93,7 +93,7 @@ def __repr__(self) -> str: class YdbRetryOperationFinalResult: def __init__(self, result: Any) -> None: self.result = result - self.exc: Optional[BaseException] = None + self.exc: BaseException | None = None def __eq__(self, other: object) -> bool: return ( @@ -112,12 +112,12 @@ def set_exception(self, exc: BaseException) -> None: def retry_operation_impl( callee: Callable[..., Any], - retry_settings: Optional[RetrySettings] = None, + retry_settings: RetrySettings | None = None, *args: Any, **kwargs: Any, -) -> Generator[Union[YdbRetryOperationSleepOpt, YdbRetryOperationFinalResult], None, None]: +) -> Generator[YdbRetryOperationSleepOpt | YdbRetryOperationFinalResult, None, None]: retry_settings = RetrySettings() if retry_settings is None else retry_settings - status: Optional[issues.Error] = None + status: issues.Error | None = None for attempt in range(retry_settings.max_retries + 1): try: @@ -161,11 +161,11 @@ def retry_operation_impl( @observe_retry_metrics def retry_operation_sync( callee: Callable[..., Any], - retry_settings: Optional[RetrySettings] = None, + retry_settings: RetrySettings | None = None, *args: Any, **kwargs: Any, ) -> Any: - backoff_ms: Optional[int] = None + backoff_ms: int | None = None @functools.wraps(callee) def traced_callee(*a: Any, **kw: Any) -> Any: @@ -186,7 +186,7 @@ def traced_callee(*a: Any, **kw: Any) -> Any: @observe_retry_metrics async def retry_operation_async( # pylint: disable=W1113 callee: Callable[..., Any], - retry_settings: Optional[RetrySettings] = None, + retry_settings: RetrySettings | None = None, *args: Any, **kwargs: Any, ) -> Any: @@ -202,7 +202,7 @@ async def retry_operation_async( # pylint: disable=W1113 Returns awaitable result of coroutine. If retries are not successful exception is raised. """ - backoff_ms: Optional[int] = None + backoff_ms: int | None = None @functools.wraps(callee) async def traced_callee(*a: Any, **kw: Any) -> Any: @@ -225,13 +225,13 @@ async def traced_callee(*a: Any, **kw: Any) -> Any: def ydb_retry( max_retries: int = 10, - max_session_acquire_timeout: Optional[float] = None, - on_ydb_error_callback: Optional[Callable[[issues.Error], None]] = None, + max_session_acquire_timeout: float | None = None, + on_ydb_error_callback: Callable[[issues.Error], None] | None = None, backoff_ceiling: int = 6, backoff_slot_duration: float = 1, get_session_client_timeout: float = 5, - fast_backoff_settings: Optional[BackoffSettings] = None, - slow_backoff_settings: Optional[BackoffSettings] = None, + fast_backoff_settings: BackoffSettings | None = None, + slow_backoff_settings: BackoffSettings | None = None, idempotent: bool = False, retry_cancelled: bool = False, ) -> Callable[[Callable[..., Any]], Callable[..., Any]]: diff --git a/ydb/scheme.py b/ydb/scheme.py index fd0630f8b..1c2db10e9 100644 --- a/ydb/scheme.py +++ b/ydb/scheme.py @@ -173,7 +173,7 @@ def is_secret(entry): return entry == SchemeEntryType.SECRET -class SchemeEntry(object): +class SchemeEntry: __slots__ = ( "name", "owner", @@ -395,7 +395,7 @@ def to_pb(self): return self._pb -class Permissions(object): +class Permissions: __slots__ = ("subject", "permission_names") def __init__(self, subject, permission_names): diff --git a/ydb/scripting.py b/ydb/scripting.py index 595f4e169..d40663147 100644 --- a/ydb/scripting.py +++ b/ydb/scripting.py @@ -12,13 +12,13 @@ from . import issues, convert, settings -class TypedParameters(object): +class TypedParameters: def __init__(self, parameters_types, parameters_values): self.parameters_types = parameters_types self.parameters_values = parameters_values -class ScriptingClientSettings(object): +class ScriptingClientSettings: def __init__(self): self._native_date_in_result_sets = False self._native_datetime_in_result_sets = False @@ -52,12 +52,12 @@ def _execute_yql_query_request_factory(script, tp=None, settings=None): return ydb_scripting_pb2.ExecuteYqlRequest(script=script, parameters=params) -class YqlQueryResult(object): +class YqlQueryResult: def __init__(self, result, scripting_client_settings=None): self.result_sets = convert.ResultSets(result.result_sets, scripting_client_settings) -class YqlExplainResult(object): +class YqlExplainResult: def __init__(self, result): self.plan = result.plan @@ -76,7 +76,7 @@ def _wrap_explain_response(rpc_state, response): return YqlExplainResult(message) -class ScriptingClient(object): +class ScriptingClient: def __init__(self, driver, scripting_client_settings=None): self.driver = driver self.scripting_client_settings = ( diff --git a/ydb/settings.py b/ydb/settings.py index ec124ca3c..08b2c9e10 100644 --- a/ydb/settings.py +++ b/ydb/settings.py @@ -1,5 +1,5 @@ # -*- coding: utf-8 -*- -from typing import Any, List, Optional, Tuple +from typing import Any class BaseRequestSettings: @@ -15,14 +15,14 @@ class BaseRequestSettings: "need_rpc_auth", ) - trace_id: Optional[str] - request_type: Optional[str] - timeout: Optional[float] - cancel_after: Optional[float] - operation_timeout: Optional[float] + trace_id: str | None + request_type: str | None + timeout: float | None + cancel_after: float | None + operation_timeout: float | None tracer: Any compression: Any - headers: List[Tuple[str, str]] + headers: list[tuple[str, str]] need_rpc_auth: bool def __init__(self) -> None: @@ -73,7 +73,7 @@ def with_header(self, key: str, value: str) -> "BaseRequestSettings": self.headers.append((key, value)) return self - def with_trace_id(self, trace_id: Optional[str]) -> "BaseRequestSettings": + def with_trace_id(self, trace_id: str | None) -> "BaseRequestSettings": """ Includes trace id for RPC headers :param trace_id: A trace id string @@ -82,7 +82,7 @@ def with_trace_id(self, trace_id: Optional[str]) -> "BaseRequestSettings": self.trace_id = trace_id return self - def with_request_type(self, request_type: Optional[str]) -> "BaseRequestSettings": + def with_request_type(self, request_type: str | None) -> "BaseRequestSettings": """ Includes request type for RPC headers :param request_type: A request type string @@ -91,7 +91,7 @@ def with_request_type(self, request_type: Optional[str]) -> "BaseRequestSettings self.request_type = request_type return self - def with_operation_timeout(self, timeout: Optional[float]) -> "BaseRequestSettings": + def with_operation_timeout(self, timeout: float | None) -> "BaseRequestSettings": """ Indicates that client is no longer interested in the result of operation after the specified duration starting from the time operation arrives at the server. @@ -105,7 +105,7 @@ def with_operation_timeout(self, timeout: Optional[float]) -> "BaseRequestSettin self.operation_timeout = timeout return self - def with_cancel_after(self, timeout: Optional[float]) -> "BaseRequestSettings": + def with_cancel_after(self, timeout: float | None) -> "BaseRequestSettings": """ Server will try to cancel the operation after the specified duration starting from the time the operation arrives at server. @@ -118,7 +118,7 @@ def with_cancel_after(self, timeout: Optional[float]) -> "BaseRequestSettings": self.cancel_after = timeout return self - def with_timeout(self, timeout: Optional[float]) -> "BaseRequestSettings": + def with_timeout(self, timeout: float | None) -> "BaseRequestSettings": """ Client-side timeout to complete request. Since YDB doesn't support request cancellation at this moment, this feature should be diff --git a/ydb/table.py b/ydb/table.py index 8fbc5e780..feb7f6119 100644 --- a/ydb/table.py +++ b/ydb/table.py @@ -9,11 +9,8 @@ from typing import ( Any, - Dict, Generic, - List, Optional, - Tuple, TYPE_CHECKING, ) @@ -86,7 +83,7 @@ def with_keep_in_cache(self, value): return self -class KeyBound(object): +class KeyBound: __slots__ = ("_equal", "value", "type") def __init__(self, key_value, key_type=None, inclusive=False): @@ -129,7 +126,7 @@ def exclusive(cls, key_value, key_type): return cls(key_value, key_type, False) -class KeyRange(object): +class KeyRange: __slots__ = ("from_bound", "to_bound") def __init__(self, from_bound, to_bound): @@ -143,7 +140,7 @@ def __str__(self): return "KeyRange(%s, %s)" % (str(self.from_bound), str(self.to_bound)) -class Column(object): +class Column: def __init__(self, name, type, family=None): self._name = name self._type = type @@ -194,7 +191,7 @@ class IndexStatus(enum.IntEnum): BUILDING = 2 -class CachingPolicy(object): +class CachingPolicy: def __init__(self): self._pb = _apis.ydb_table.CachingPolicy() self.preset_name = None @@ -208,7 +205,7 @@ def to_pb(self): return self._pb -class ExecutionPolicy(object): +class ExecutionPolicy: def __init__(self): self._pb = _apis.ydb_table.ExecutionPolicy() self.preset_name = None @@ -222,7 +219,7 @@ def to_pb(self): return self._pb -class CompactionPolicy(object): +class CompactionPolicy: def __init__(self): self._pb = _apis.ydb_table.CompactionPolicy() self.preset_name = None @@ -236,7 +233,7 @@ def to_pb(self): return self._pb -class SplitPoint(object): +class SplitPoint: def __init__(self, *args): self._value = tuple(args) @@ -245,12 +242,12 @@ def value(self): return self._value -class ExplicitPartitions(object): +class ExplicitPartitions: def __init__(self, split_points): self.split_points = split_points -class PartitioningPolicy(object): +class PartitioningPolicy: def __init__(self): self._pb = _apis.ydb_table.PartitioningPolicy() self.preset_name = None @@ -301,7 +298,7 @@ def to_pb(self, table_description): return self._pb -class TableIndex(object): +class TableIndex: def __init__(self, name): self._pb = _apis.ydb_table.TableIndex() self._pb.name = name @@ -349,7 +346,7 @@ def to_pb(self): ) -class ReplicationPolicy(object): +class ReplicationPolicy: def __init__(self): self._pb = _apis.ydb_table.ReplicationPolicy() self.preset_name = None @@ -381,7 +378,7 @@ def to_pb(self): return self._pb -class StoragePool(object): +class StoragePool: def __init__(self, media): self.media = media @@ -389,7 +386,7 @@ def to_pb(self): return _apis.ydb_table.StoragePool(media=self.media) -class StoragePolicy(object): +class StoragePolicy: def __init__(self): self._pb = _apis.ydb_table.StoragePolicy() self.preset_name = None @@ -433,7 +430,7 @@ def to_pb(self): return self._pb -class TableProfile(object): +class TableProfile: def __init__(self): self.preset_name = None self.compaction_policy = None @@ -498,7 +495,7 @@ def to_pb(self, table_description): return pb -class DateTypeColumnModeSettings(object): +class DateTypeColumnModeSettings: def __init__(self, column_name, expire_after_seconds=0): self.column_name = column_name self.expire_after_seconds = expire_after_seconds @@ -521,7 +518,7 @@ class ColumnUnit(enum.IntEnum): UNIT_NANOSECONDS = 4 -class ValueSinceUnixEpochModeSettings(object): +class ValueSinceUnixEpochModeSettings: def __init__(self, column_name, column_unit, expire_after_seconds=0): self.column_name = column_name self.column_unit = column_unit @@ -537,7 +534,7 @@ def to_pb(self): return pb -class TtlSettings(object): +class TtlSettings: def __init__(self): self.date_type_column = None self.value_since_unix_epoch = None @@ -563,7 +560,7 @@ def to_pb(self): return pb -class TableStats(object): +class TableStats: def __init__(self): self.partitions = None self.store_size = 0 @@ -592,7 +589,7 @@ def with_modification_time(self, modification_time): return self -class ReadReplicasSettings(object): +class ReadReplicasSettings: def __init__(self): self.per_az_read_replicas_count = 0 self.any_az_read_replicas_count = 0 @@ -614,7 +611,7 @@ def to_pb(self): return pb -class PartitioningSettings(object): +class PartitioningSettings: def __init__(self): self.partitioning_by_size = 0 self.partition_size_mb = 0 @@ -652,7 +649,7 @@ def to_pb(self): return pb -class StorageSettings(object): +class StorageSettings: def __init__(self): self.tablet_commit_log0 = None self.tablet_commit_log1 = None @@ -694,7 +691,7 @@ class Compression(enum.IntEnum): LZ4 = 2 -class ColumnFamily(object): +class ColumnFamily: def __init__(self): self.compression = 0 self.name = None @@ -728,7 +725,7 @@ def to_pb(self): return cm -class TableDescription(object): +class TableDescription: def __init__(self): self.columns = [] self.primary_key = [] @@ -903,7 +900,7 @@ def name(self): return self._name -class TableClientSettings(object): +class TableClientSettings: def __init__(self): self._client_query_cache_enabled = False self._native_datetime_in_result_sets = False @@ -955,7 +952,7 @@ def with_allow_truncated_result(self, enabled): return self -class ScanQueryResult(object): +class ScanQueryResult: def __init__(self, result, table_client_settings): self._result = result self.query_stats = result.query_stats @@ -979,7 +976,7 @@ def with_collect_stats(self, collect_stats_mode): return self -class ScanQuery(object): +class ScanQuery: def __init__(self, yql_text, parameters_types): self.yql_text = yql_text self.parameters_types = parameters_types @@ -1198,7 +1195,7 @@ def bulk_upsert(self, table_path, rows, column_types, settings=None): class BaseTableClient(ITableClient, Generic[DriverT]): _driver: DriverT - def __init__(self, driver: DriverT, table_client_settings: Optional[TableClientSettings] = None) -> None: + def __init__(self, driver: DriverT, table_client_settings: TableClientSettings | None = None) -> None: self._driver = driver self._table_client_settings = TableClientSettings() if table_client_settings is None else table_client_settings @@ -1260,9 +1257,9 @@ def describe_system_view(self, path, settings=None): class TableClient(BaseTableClient["SyncDriver"]): - def __init__(self, driver: "SyncDriver", table_client_settings: Optional[TableClientSettings] = None) -> None: + def __init__(self, driver: "SyncDriver", table_client_settings: TableClientSettings | None = None) -> None: super().__init__(driver=driver, table_client_settings=table_client_settings) - self._pool: Optional[SessionPool] = None + self._pool: SessionPool | None = None def __del__(self): self._stop_pool_if_needed() @@ -1361,42 +1358,42 @@ def callee(session: Session): def alter_table( self, path: str, - add_columns: Optional[List["ydb.Column"]] = None, - drop_columns: Optional[List[str]] = None, + add_columns: list["ydb.Column"] | None = None, + drop_columns: list[str] | None = None, settings: Optional["settings_impl.BaseRequestSettings"] = None, - alter_attributes: Optional[Optional[Dict[str, str]]] = None, - add_indexes: Optional[List["ydb.TableIndex"]] = None, - drop_indexes: Optional[List[str]] = None, + alter_attributes: dict[str, str] | None = None, + add_indexes: list["ydb.TableIndex"] | None = None, + drop_indexes: list[str] | None = None, set_ttl_settings: Optional["ydb.TtlSettings"] = None, - drop_ttl_settings: Optional[Any] = None, - add_column_families: Optional[List["ydb.ColumnFamily"]] = None, - alter_column_families: Optional[List["ydb.ColumnFamily"]] = None, + drop_ttl_settings: Any | None = None, + add_column_families: list["ydb.ColumnFamily"] | None = None, + alter_column_families: list["ydb.ColumnFamily"] | None = None, alter_storage_settings: Optional["ydb.StorageSettings"] = None, - set_compaction_policy: Optional[str] = None, + set_compaction_policy: str | None = None, alter_partitioning_settings: Optional["ydb.PartitioningSettings"] = None, set_key_bloom_filter: Optional["ydb.FeatureFlag"] = None, set_read_replicas_settings: Optional["ydb.ReadReplicasSettings"] = None, - rename_indexes: Optional[List["ydb.RenameIndexItem"]] = None, + rename_indexes: list["ydb.RenameIndexItem"] | None = None, ) -> "ydb.Operation": """ Alter a YDB table. :param path: A table path - :param add_columns: List of ydb.Column to add - :param drop_columns: List of column names to drop + :param add_columns: list of ydb.Column to add + :param drop_columns: list of column names to drop :param settings: An instance of BaseRequestSettings that describes how rpc should be invoked. - :param alter_attributes: Dict of attributes to alter - :param add_indexes: List of ydb.TableIndex to add - :param drop_indexes: List of index names to drop + :param alter_attributes: dict of attributes to alter + :param add_indexes: list of ydb.TableIndex to add + :param drop_indexes: list of index names to drop :param set_ttl_settings: ydb.TtlSettings to set :param drop_ttl_settings: Any to drop - :param add_column_families: List of ydb.ColumnFamily to add - :param alter_column_families: List of ydb.ColumnFamily to alter + :param add_column_families: list of ydb.ColumnFamily to add + :param alter_column_families: list of ydb.ColumnFamily to alter :param alter_storage_settings: ydb.StorageSettings to alter :param set_compaction_policy: Compaction policy :param alter_partitioning_settings: ydb.PartitioningSettings to alter :param set_key_bloom_filter: ydb.FeatureFlag to set key bloom filter - :param rename_indexes: List of ydb.RenameIndexItem to rename + :param rename_indexes: list of ydb.RenameIndexItem to rename :return: Operation or YDB error otherwise. """ @@ -1479,13 +1476,13 @@ def callee(session: Session): def copy_tables( self, - source_destination_pairs: List[Tuple[str, str]], + source_destination_pairs: list[tuple[str, str]], settings: Optional["settings_impl.BaseRequestSettings"] = None, ) -> "ydb.Operation": """ Copy a YDB tables. - :param source_destination_pairs: List of tuples (source_path, destination_path) + :param source_destination_pairs: list of tuples (source_path, destination_path) :param settings: An instance of BaseRequestSettings that describes how rpc should be invoked. :return: Operation or YDB error otherwise. @@ -1501,13 +1498,13 @@ def callee(session: Session): def rename_tables( self, - rename_items: List[Tuple[str, str]], + rename_items: list[tuple[str, str]], settings: Optional["settings_impl.BaseRequestSettings"] = None, ) -> "ydb.Operation": """ Rename a YDB tables. - :param rename_items: List of tuples (current_name, desired_name) + :param rename_items: list of tuples (current_name, desired_name) :param settings: An instance of BaseRequestSettings that describes how rpc should be invoked. :return: Operation or YDB error otherwise. @@ -2685,7 +2682,7 @@ def async_begin(self, settings=None): ) -class SessionPool(object): +class SessionPool: def __init__( self, driver, @@ -2781,7 +2778,7 @@ def __exit__(self, exc_type, exc_val, exc_tb): self.stop() -class AsyncSessionCheckout(object): +class AsyncSessionCheckout: __slots__ = ("subscription", "pool") def __init__(self, pool): @@ -2802,7 +2799,7 @@ def __exit__(self, exc_type, exc_val, exc_tb): self.pool.unsubscribe(self.subscription) -class SessionCheckout(object): +class SessionCheckout: __slots__ = ("_acquired", "_pool", "_blocking", "_timeout") def __init__(self, pool, blocking, timeout): diff --git a/ydb/topic.py b/ydb/topic.py index 988592937..16d8a0f69 100644 --- a/ydb/topic.py +++ b/ydb/topic.py @@ -39,7 +39,7 @@ import datetime from dataclasses import dataclass import logging -from typing import List, Union, Mapping, Optional, Dict, Callable +from typing import Mapping, Callable from . import aio, Credentials, _apis, issues @@ -110,11 +110,11 @@ class TopicClientAsyncIO: _closed: bool _driver: aio.Driver - _credentials: Union[Credentials, None] + _credentials: Credentials | None _settings: TopicClientSettings _executor: concurrent.futures.Executor - def __init__(self, driver: aio.Driver, settings: Optional[TopicClientSettings] = None): + def __init__(self, driver: aio.Driver, settings: TopicClientSettings | None = None): if not settings: settings = TopicClientSettings() self._closed = False @@ -136,18 +136,18 @@ def __del__(self): async def create_topic( self, path: str, - min_active_partitions: Optional[int] = None, - max_active_partitions: Optional[int] = None, - partition_count_limit: Optional[int] = None, - retention_period: Optional[datetime.timedelta] = None, - retention_storage_mb: Optional[int] = None, - supported_codecs: Optional[List[Union[TopicCodec, int]]] = None, - partition_write_speed_bytes_per_second: Optional[int] = None, - partition_write_burst_bytes: Optional[int] = None, - attributes: Optional[Dict[str, str]] = None, - consumers: Optional[List[Union[TopicConsumer, str]]] = None, - metering_mode: Optional[TopicMeteringMode] = None, - auto_partitioning_settings: Optional[TopicAutoPartitioningSettings] = None, + min_active_partitions: int | None = None, + max_active_partitions: int | None = None, + partition_count_limit: int | None = None, + retention_period: datetime.timedelta | None = None, + retention_storage_mb: int | None = None, + supported_codecs: list[TopicCodec | int] | None = None, + partition_write_speed_bytes_per_second: int | None = None, + partition_write_burst_bytes: int | None = None, + attributes: dict[str, str] | None = None, + consumers: list[TopicConsumer | str] | None = None, + metering_mode: TopicMeteringMode | None = None, + auto_partitioning_settings: TopicAutoPartitioningSettings | None = None, ): """ create topic command @@ -158,13 +158,13 @@ async def create_topic( and read-only partitions. :param retention_period: How long data in partition should be stored :param retention_storage_mb: How much data in partition should be stored - :param supported_codecs: List of allowed codecs for writers. Writes with codec not from this list are forbidden. + :param supported_codecs: list of allowed codecs for writers. Writes with codec not from this list are forbidden. Empty list mean disable codec compatibility checks for the topic. :param partition_write_speed_bytes_per_second: Partition write speed in bytes per second :param partition_write_burst_bytes: Burst size for write in partition, in bytes :param attributes: User and server attributes of topic. Server attributes starts from "_" and will be validated by server. - :param consumers: List of consumers for this topic + :param consumers: list of consumers for this topic :param metering_mode: Metering mode for the topic in a serverless database """ logger.debug("Create topic request: path=%s", path) @@ -182,20 +182,20 @@ async def create_topic( async def alter_topic( self, path: str, - set_min_active_partitions: Optional[int] = None, - set_max_active_partitions: Optional[int] = None, - set_partition_count_limit: Optional[int] = None, - add_consumers: Optional[List[Union[TopicConsumer, str]]] = None, - alter_consumers: Optional[List[Union[TopicAlterConsumer, str]]] = None, - drop_consumers: Optional[List[str]] = None, - alter_attributes: Optional[Dict[str, str]] = None, - set_metering_mode: Optional[TopicMeteringMode] = None, - set_partition_write_speed_bytes_per_second: Optional[int] = None, - set_partition_write_burst_bytes: Optional[int] = None, - set_retention_period: Optional[datetime.timedelta] = None, - set_retention_storage_mb: Optional[int] = None, - set_supported_codecs: Optional[List[Union[TopicCodec, int]]] = None, - alter_auto_partitioning_settings: Optional[TopicAlterAutoPartitioningSettings] = None, + set_min_active_partitions: int | None = None, + set_max_active_partitions: int | None = None, + set_partition_count_limit: int | None = None, + add_consumers: list[TopicConsumer | str] | None = None, + alter_consumers: list[TopicAlterConsumer | str] | None = None, + drop_consumers: list[str] | None = None, + alter_attributes: dict[str, str] | None = None, + set_metering_mode: TopicMeteringMode | None = None, + set_partition_write_speed_bytes_per_second: int | None = None, + set_partition_write_burst_bytes: int | None = None, + set_retention_period: datetime.timedelta | None = None, + set_retention_storage_mb: int | None = None, + set_supported_codecs: list[TopicCodec | int] | None = None, + alter_auto_partitioning_settings: TopicAlterAutoPartitioningSettings | None = None, ): """ alter topic command @@ -204,9 +204,9 @@ async def alter_topic( :param set_min_active_partitions: Minimum partition count auto merge would stop working at. :param set_partition_count_limit: Limit for total partition count, including active (open for write) and read-only partitions. - :param add_consumers: List of consumers for this topic to add - :param alter_consumers: List of consumers for this topic to alter - :param drop_consumers: List of consumer names for this topic to drop + :param add_consumers: list of consumers for this topic to add + :param alter_consumers: list of consumers for this topic to alter + :param drop_consumers: list of consumer names for this topic to drop :param alter_attributes: User and server attributes of topic. Server attributes starts from "_" and will be validated by server. :param set_metering_mode: Metering mode for the topic in a serverless database @@ -214,7 +214,7 @@ async def alter_topic( :param set_partition_write_burst_bytes: Burst size for write in partition, in bytes :param set_retention_period: How long data in partition should be stored :param set_retention_storage_mb: How much data in partition should be stored - :param set_supported_codecs: List of allowed codecs for writers. Writes with codec not from this list are forbidden. + :param set_supported_codecs: list of allowed codecs for writers. Writes with codec not from this list are forbidden. Empty list mean disable codec compatibility checks for the topic. """ logger.debug("Alter topic request: path=%s", path) @@ -283,17 +283,17 @@ async def drop_topic(self, path: str): def reader( self, - topic: Union[str, TopicReaderSelector, List[Union[str, TopicReaderSelector]]], - consumer: Optional[str], + topic: str | TopicReaderSelector | list[str | TopicReaderSelector], + consumer: str | None, buffer_size_bytes: int = 50 * 1024 * 1024, # decoders: map[codec_code] func(encoded_bytes)->decoded_bytes # the func will be called from multiply threads in parallel - decoders: Union[Mapping[int, Callable[[bytes], bytes]], None] = None, + decoders: Mapping[int, Callable[[bytes], bytes]] | None = None, # custom decoder executor for call builtin and custom decoders. If None - use shared executor pool. # if max_worker in the executor is 1 - then decoders will be called from the thread without parallel - decoder_executor: Optional[concurrent.futures.Executor] = None, - auto_partitioning_support: Optional[bool] = True, # Auto partitioning feature flag. Default - True. - event_handler: Optional[TopicReaderEvents.EventHandler] = None, + decoder_executor: concurrent.futures.Executor | None = None, + auto_partitioning_support: bool | None = True, # Auto partitioning feature flag. Default - True. + event_handler: TopicReaderEvents.EventHandler | None = None, buffer_release_threshold: float = 0.5, ) -> TopicReaderAsyncIO: @@ -330,21 +330,21 @@ def writer( self, topic, *, - producer_id: Optional[str] = None, # default - random + producer_id: str | None = None, # default - random session_metadata: Mapping[str, str] = None, - partition_id: Union[int, None] = None, + partition_id: int | None = None, auto_seqno: bool = True, auto_created_at: bool = True, - codec: Optional[TopicCodec] = None, # default mean auto-select + codec: TopicCodec | None = None, # default mean auto-select # encoders: map[codec_code] func(encoded_bytes)->decoded_bytes # the func will be called from multiply threads in parallel. - encoders: Optional[Mapping[_ydb_topic_public_types.PublicCodec, Callable[[bytes], bytes]]] = None, + encoders: Mapping[_ydb_topic_public_types.PublicCodec, Callable[[bytes], bytes]] | None = None, # custom encoder executor for call builtin and custom decoders. If None - use shared executor pool. # If max_worker in the executor is 1 - then encoders will be called from the thread without parallel. - encoder_executor: Optional[concurrent.futures.Executor] = None, - max_buffer_size_bytes: Optional[int] = None, - max_buffer_messages: Optional[int] = None, - buffer_wait_timeout_sec: Optional[float] = None, + encoder_executor: concurrent.futures.Executor | None = None, + max_buffer_size_bytes: int | None = None, + max_buffer_messages: int | None = None, + buffer_wait_timeout_sec: float | None = None, ) -> TopicWriterAsyncIO: logger.debug("Create writer for topic=%s producer_id=%s", topic, producer_id) args = locals().copy() @@ -362,21 +362,21 @@ def tx_writer( tx, topic, *, - producer_id: Optional[str] = None, # default - random + producer_id: str | None = None, # default - random session_metadata: Mapping[str, str] = None, - partition_id: Union[int, None] = None, + partition_id: int | None = None, auto_seqno: bool = True, auto_created_at: bool = True, - codec: Optional[TopicCodec] = None, # default mean auto-select + codec: TopicCodec | None = None, # default mean auto-select # encoders: map[codec_code] func(encoded_bytes)->decoded_bytes # the func will be called from multiply threads in parallel. - encoders: Optional[Mapping[_ydb_topic_public_types.PublicCodec, Callable[[bytes], bytes]]] = None, + encoders: Mapping[_ydb_topic_public_types.PublicCodec, Callable[[bytes], bytes]] | None = None, # custom encoder executor for call builtin and custom decoders. If None - use shared executor pool. # If max_worker in the executor is 1 - then encoders will be called from the thread without parallel. - encoder_executor: Optional[concurrent.futures.Executor] = None, - max_buffer_size_bytes: Optional[int] = None, - max_buffer_messages: Optional[int] = None, - buffer_wait_timeout_sec: Optional[float] = None, + encoder_executor: concurrent.futures.Executor | None = None, + max_buffer_size_bytes: int | None = None, + max_buffer_messages: int | None = None, + buffer_wait_timeout_sec: float | None = None, ) -> TopicTxWriterAsyncIO: logger.debug("Create tx writer for topic=%s tx=%s", topic, tx) args = locals().copy() @@ -392,7 +392,7 @@ def tx_writer( @ydb_retry(retry_cancelled=True, idempotent=True) async def commit_offset( - self, path: str, consumer: str, partition_id: int, offset: int, read_session_id: Optional[str] = None + self, path: str, consumer: str, partition_id: int, offset: int, read_session_id: str | None = None ) -> None: logger.debug( "Commit offset: path=%s partition_id=%s offset=%s consumer=%s", @@ -434,11 +434,11 @@ def _check_closed(self): class TopicClient: _closed: bool _driver: driver.Driver - _credentials: Union[Credentials, None] + _credentials: Credentials | None _settings: TopicClientSettings _executor: concurrent.futures.Executor - def __init__(self, driver: driver.Driver, settings: Optional[TopicClientSettings]): + def __init__(self, driver: driver.Driver, settings: TopicClientSettings | None): if not settings: settings = TopicClientSettings() @@ -461,18 +461,18 @@ def __del__(self): def create_topic( self, path: str, - min_active_partitions: Optional[int] = None, - max_active_partitions: Optional[int] = None, - partition_count_limit: Optional[int] = None, - retention_period: Optional[datetime.timedelta] = None, - retention_storage_mb: Optional[int] = None, - supported_codecs: Optional[List[Union[TopicCodec, int]]] = None, - partition_write_speed_bytes_per_second: Optional[int] = None, - partition_write_burst_bytes: Optional[int] = None, - attributes: Optional[Dict[str, str]] = None, - consumers: Optional[List[Union[TopicConsumer, str]]] = None, - metering_mode: Optional[TopicMeteringMode] = None, - auto_partitioning_settings: Optional[TopicAutoPartitioningSettings] = None, + min_active_partitions: int | None = None, + max_active_partitions: int | None = None, + partition_count_limit: int | None = None, + retention_period: datetime.timedelta | None = None, + retention_storage_mb: int | None = None, + supported_codecs: list[TopicCodec | int] | None = None, + partition_write_speed_bytes_per_second: int | None = None, + partition_write_burst_bytes: int | None = None, + attributes: dict[str, str] | None = None, + consumers: list[TopicConsumer | str] | None = None, + metering_mode: TopicMeteringMode | None = None, + auto_partitioning_settings: TopicAutoPartitioningSettings | None = None, ): """ create topic command @@ -483,13 +483,13 @@ def create_topic( and read-only partitions. :param retention_period: How long data in partition should be stored :param retention_storage_mb: How much data in partition should be stored - :param supported_codecs: List of allowed codecs for writers. Writes with codec not from this list are forbidden. + :param supported_codecs: list of allowed codecs for writers. Writes with codec not from this list are forbidden. Empty list mean disable codec compatibility checks for the topic. :param partition_write_speed_bytes_per_second: Partition write speed in bytes per second :param partition_write_burst_bytes: Burst size for write in partition, in bytes :param attributes: User and server attributes of topic. Server attributes starts from "_" and will be validated by server. - :param consumers: List of consumers for this topic + :param consumers: list of consumers for this topic :param metering_mode: Metering mode for the topic in a serverless database """ logger.debug("Create topic request: path=%s", path) @@ -509,20 +509,20 @@ def create_topic( def alter_topic( self, path: str, - set_min_active_partitions: Optional[int] = None, - set_max_active_partitions: Optional[int] = None, - set_partition_count_limit: Optional[int] = None, - add_consumers: Optional[List[Union[TopicConsumer, str]]] = None, - alter_consumers: Optional[List[Union[TopicAlterConsumer, str]]] = None, - drop_consumers: Optional[List[str]] = None, - alter_attributes: Optional[Dict[str, str]] = None, - set_metering_mode: Optional[TopicMeteringMode] = None, - set_partition_write_speed_bytes_per_second: Optional[int] = None, - set_partition_write_burst_bytes: Optional[int] = None, - set_retention_period: Optional[datetime.timedelta] = None, - set_retention_storage_mb: Optional[int] = None, - set_supported_codecs: Optional[List[Union[TopicCodec, int]]] = None, - alter_auto_partitioning_settings: Optional[TopicAlterAutoPartitioningSettings] = None, + set_min_active_partitions: int | None = None, + set_max_active_partitions: int | None = None, + set_partition_count_limit: int | None = None, + add_consumers: list[TopicConsumer | str] | None = None, + alter_consumers: list[TopicAlterConsumer | str] | None = None, + drop_consumers: list[str] | None = None, + alter_attributes: dict[str, str] | None = None, + set_metering_mode: TopicMeteringMode | None = None, + set_partition_write_speed_bytes_per_second: int | None = None, + set_partition_write_burst_bytes: int | None = None, + set_retention_period: datetime.timedelta | None = None, + set_retention_storage_mb: int | None = None, + set_supported_codecs: list[TopicCodec | int] | None = None, + alter_auto_partitioning_settings: TopicAlterAutoPartitioningSettings | None = None, ): """ alter topic command @@ -531,9 +531,9 @@ def alter_topic( :param set_min_active_partitions: Minimum partition count auto merge would stop working at. :param set_partition_count_limit: Limit for total partition count, including active (open for write) and read-only partitions. - :param add_consumers: List of consumers for this topic to add - :param alter_consumers: List of consumers for this topic to alter - :param drop_consumers: List of consumer names for this topic to drop + :param add_consumers: list of consumers for this topic to add + :param alter_consumers: list of consumers for this topic to alter + :param drop_consumers: list of consumer names for this topic to drop :param alter_attributes: User and server attributes of topic. Server attributes starts from "_" and will be validated by server. :param set_metering_mode: Metering mode for the topic in a serverless database @@ -541,7 +541,7 @@ def alter_topic( :param set_partition_write_burst_bytes: Burst size for write in partition, in bytes :param set_retention_period: How long data in partition should be stored :param set_retention_storage_mb: How much data in partition should be stored - :param set_supported_codecs: List of allowed codecs for writers. Writes with codec not from this list are forbidden. + :param set_supported_codecs: list of allowed codecs for writers. Writes with codec not from this list are forbidden. Empty list mean disable codec compatibility checks for the topic. """ logger.debug("Alter topic request: path=%s", path) @@ -619,17 +619,17 @@ def drop_topic(self, path: str): def reader( self, - topic: Union[str, TopicReaderSelector, List[Union[str, TopicReaderSelector]]], - consumer: Optional[str], + topic: str | TopicReaderSelector | list[str | TopicReaderSelector], + consumer: str | None, buffer_size_bytes: int = 50 * 1024 * 1024, # decoders: map[codec_code] func(encoded_bytes)->decoded_bytes # the func will be called from multiply threads in parallel - decoders: Union[Mapping[int, Callable[[bytes], bytes]], None] = None, + decoders: Mapping[int, Callable[[bytes], bytes]] | None = None, # custom decoder executor for call builtin and custom decoders. If None - use shared executor pool. # if max_worker in the executor is 1 - then decoders will be called from the thread without parallel - decoder_executor: Optional[concurrent.futures.Executor] = None, # default shared client executor pool - auto_partitioning_support: Optional[bool] = True, # Auto partitioning feature flag. Default - True. - event_handler: Optional[TopicReaderEvents.EventHandler] = None, + decoder_executor: concurrent.futures.Executor | None = None, # default shared client executor pool + auto_partitioning_support: bool | None = True, # Auto partitioning feature flag. Default - True. + event_handler: TopicReaderEvents.EventHandler | None = None, buffer_release_threshold: float = 0.5, ) -> TopicReader: logger.debug("Create reader for topic=%s consumer=%s", topic, consumer) @@ -664,21 +664,21 @@ def writer( self, topic, *, - producer_id: Optional[str] = None, # default - random + producer_id: str | None = None, # default - random session_metadata: Mapping[str, str] = None, - partition_id: Union[int, None] = None, + partition_id: int | None = None, auto_seqno: bool = True, auto_created_at: bool = True, - codec: Optional[TopicCodec] = None, # default mean auto-select + codec: TopicCodec | None = None, # default mean auto-select # encoders: map[codec_code] func(encoded_bytes)->decoded_bytes # the func will be called from multiply threads in parallel. - encoders: Optional[Mapping[_ydb_topic_public_types.PublicCodec, Callable[[bytes], bytes]]] = None, + encoders: Mapping[_ydb_topic_public_types.PublicCodec, Callable[[bytes], bytes]] | None = None, # custom encoder executor for call builtin and custom decoders. If None - use shared executor pool. # If max_worker in the executor is 1 - then encoders will be called from the thread without parallel. - encoder_executor: Optional[concurrent.futures.Executor] = None, # default shared client executor pool - max_buffer_size_bytes: Optional[int] = None, - max_buffer_messages: Optional[int] = None, - buffer_wait_timeout_sec: Optional[float] = None, + encoder_executor: concurrent.futures.Executor | None = None, # default shared client executor pool + max_buffer_size_bytes: int | None = None, + max_buffer_messages: int | None = None, + buffer_wait_timeout_sec: float | None = None, ) -> TopicWriter: logger.debug("Create writer for topic=%s producer_id=%s", topic, producer_id) args = locals().copy() @@ -697,21 +697,21 @@ def tx_writer( tx, topic, *, - producer_id: Optional[str] = None, # default - random + producer_id: str | None = None, # default - random session_metadata: Mapping[str, str] = None, - partition_id: Union[int, None] = None, + partition_id: int | None = None, auto_seqno: bool = True, auto_created_at: bool = True, - codec: Optional[TopicCodec] = None, # default mean auto-select + codec: TopicCodec | None = None, # default mean auto-select # encoders: map[codec_code] func(encoded_bytes)->decoded_bytes # the func will be called from multiply threads in parallel. - encoders: Optional[Mapping[_ydb_topic_public_types.PublicCodec, Callable[[bytes], bytes]]] = None, + encoders: Mapping[_ydb_topic_public_types.PublicCodec, Callable[[bytes], bytes]] | None = None, # custom encoder executor for call builtin and custom decoders. If None - use shared executor pool. # If max_worker in the executor is 1 - then encoders will be called from the thread without parallel. - encoder_executor: Optional[concurrent.futures.Executor] = None, # default shared client executor pool - max_buffer_size_bytes: Optional[int] = None, - max_buffer_messages: Optional[int] = None, - buffer_wait_timeout_sec: Optional[float] = None, + encoder_executor: concurrent.futures.Executor | None = None, # default shared client executor pool + max_buffer_size_bytes: int | None = None, + max_buffer_messages: int | None = None, + buffer_wait_timeout_sec: float | None = None, ) -> TopicWriter: logger.debug("Create tx writer for topic=%s tx=%s", topic, tx) args = locals().copy() @@ -728,7 +728,7 @@ def tx_writer( @ydb_retry(retry_cancelled=True, idempotent=True) def commit_offset( - self, path: str, consumer: str, partition_id: int, offset: int, read_session_id: Optional[str] = None + self, path: str, consumer: str, partition_id: int, offset: int, read_session_id: str | None = None ) -> None: logger.debug( "Commit offset: path=%s partition_id=%s offset=%s consumer=%s", diff --git a/ydb/tracing.py b/ydb/tracing.py index 78590d153..6f88a06bd 100644 --- a/ydb/tracing.py +++ b/ydb/tracing.py @@ -1,6 +1,6 @@ from enum import IntEnum import functools -from typing import Any, Callable, Dict, Optional, Type +from typing import Any, Callable from types import TracebackType @@ -37,12 +37,12 @@ def enabled(self) -> bool: """ return self._enabled - def trace(self, tags: Dict[str, Any], trace_level: TraceLevel = TraceLevel.INFO) -> None: + def trace(self, tags: dict[str, Any], trace_level: TraceLevel = TraceLevel.INFO) -> None: """ Add tags to current span :param ydb.TraceLevel trace_level: level of tracing - :param dict tags: Dict of tags + :param dict tags: dict of tags """ if self._tracer._verbose_level < trace_level: return @@ -53,9 +53,9 @@ def trace(self, tags: Dict[str, Any], trace_level: TraceLevel = TraceLevel.INFO) def __exit__( self, - exc_type: Optional[Type[BaseException]], - exc_val: Optional[BaseException], - exc_tb: Optional[TracebackType], + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, ) -> None: if not self.enabled: return @@ -68,7 +68,7 @@ def __exit__( self._scope = None -def with_trace(span_name: Optional[str] = None) -> Callable[[Callable[..., Any]], Callable[..., Any]]: +def with_trace(span_name: str | None = None) -> Callable[[Callable[..., Any]], Callable[..., Any]]: def decorator(f: Callable[..., Any]) -> Callable[..., Any]: @functools.wraps(f) def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any: @@ -81,7 +81,7 @@ def wrapper(self: Any, *args: Any, **kwargs: Any) -> Any: return decorator -def trace(tracer: "Tracer", tags: Dict[str, Any], trace_level: TraceLevel = TraceLevel.INFO) -> Optional[bool]: +def trace(tracer: "Tracer", tags: dict[str, Any], trace_level: TraceLevel = TraceLevel.INFO) -> bool | None: if tracer.enabled: scope = tracer._open_tracer.scope_manager.active if not scope: @@ -104,9 +104,9 @@ def __init__(self, tracer: Any) -> None: :param opentracing.Tracer tracer: opentracing.Tracer implementation. If None - tracing not enabled """ self._open_tracer: Any = tracer - self._pre_tags: Dict[str, Any] = {} - self._post_tags_ok: Dict[str, Any] = {} - self._post_tags_err: Dict[str, Any] = {} + self._pre_tags: dict[str, Any] = {} + self._post_tags_ok: dict[str, Any] = {} + self._post_tags_err: dict[str, Any] = {} self._on_err: Callable[..., None] = lambda *args, **kwargs: None self._verbose_level: TraceLevel = TraceLevel.NONE @@ -125,7 +125,7 @@ def trace(self, span_name: str) -> _TracingCtx: """ return _TracingCtx(self, span_name) - def with_pre_tags(self, tags: Dict[str, Any]) -> "Tracer": + def with_pre_tags(self, tags: dict[str, Any]) -> "Tracer": """ Add `tags` to every span immediately after creation @@ -136,7 +136,7 @@ def with_pre_tags(self, tags: Dict[str, Any]) -> "Tracer": self._pre_tags = tags return self - def with_post_tags(self, ok_tags: Dict[str, Any], err_tags: Dict[str, Any]) -> "Tracer": + def with_post_tags(self, ok_tags: dict[str, Any], err_tags: dict[str, Any]) -> "Tracer": """ Add some tags before span close @@ -184,9 +184,9 @@ def default(cls, tracer: Any) -> "Tracer": def _default_on_error_callback( ctx: _TracingCtx, - exc_type: Optional[Type[BaseException]], - exc_val: Optional[BaseException], - exc_tb: Optional[TracebackType], + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, ) -> None: ctx.trace( { diff --git a/ydb/types.py b/ydb/types.py index cfdfa809c..9d5ede9a3 100644 --- a/ydb/types.py +++ b/ydb/types.py @@ -28,13 +28,13 @@ _EPOCH_UTC = datetime(1970, 1, 1, tzinfo=timezone.utc) -def _from_date(x: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings) -> typing.Union[date, int]: +def _from_date(x: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings) -> date | int: if table_client_settings is not None and table_client_settings._native_date_in_result_sets: return _EPOCH.date() + timedelta(days=x.uint32_value) return x.uint32_value -def _to_date(pb: ydb_value_pb2.Value, value: typing.Union[date, datetime, int]) -> None: +def _to_date(pb: ydb_value_pb2.Value, value: date | datetime | int) -> None: if isinstance(value, datetime): pb.uint32_value = (value.date() - _EPOCH.date()).days elif isinstance(value, date): @@ -43,29 +43,27 @@ def _to_date(pb: ydb_value_pb2.Value, value: typing.Union[date, datetime, int]) pb.uint32_value = value -def _from_date32(x: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings) -> typing.Union[date, int]: +def _from_date32(x: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings) -> date | int: if table_client_settings is not None and table_client_settings._native_date_in_result_sets: return _EPOCH.date() + timedelta(days=x.int32_value) return x.int32_value -def _to_date32(pb: ydb_value_pb2.Value, value: typing.Union[date, int]) -> None: +def _to_date32(pb: ydb_value_pb2.Value, value: date | int) -> None: if isinstance(value, date): pb.int32_value = (value - _EPOCH.date()).days else: pb.int32_value = value -def _from_datetime_number( - x: typing.Union[float, datetime], table_client_settings: table.TableClientSettings -) -> typing.Union[float, datetime]: +def _from_datetime_number(x: float | datetime, table_client_settings: table.TableClientSettings) -> float | datetime: if table_client_settings is not None and table_client_settings._native_datetime_in_result_sets: # x is float when native_datetime_in_result_sets is True return datetime.utcfromtimestamp(typing.cast(float, x)) return x -def _to_datetime(pb: ydb_value_pb2.Value, value: typing.Union[datetime, int]) -> None: +def _to_datetime(pb: ydb_value_pb2.Value, value: datetime | int) -> None: if isinstance(value, datetime): epoch = _EPOCH_UTC if value.tzinfo else _EPOCH pb.uint32_value = (value - epoch) // timedelta(seconds=1) @@ -73,7 +71,7 @@ def _to_datetime(pb: ydb_value_pb2.Value, value: typing.Union[datetime, int]) -> pb.uint32_value = value -def _to_datetime64(pb: ydb_value_pb2.Value, value: typing.Union[datetime, int]) -> None: +def _to_datetime64(pb: ydb_value_pb2.Value, value: datetime | int) -> None: if isinstance(value, datetime): epoch = _EPOCH_UTC if value.tzinfo else _EPOCH pb.int64_value = (value - epoch) // timedelta(seconds=1) @@ -102,39 +100,39 @@ def _parse_tz(value: str) -> datetime: return naive.replace(tzinfo=tz) -def _from_tz_date(x: str, table_client_settings: table.TableClientSettings) -> typing.Union[datetime, str]: +def _from_tz_date(x: str, table_client_settings: table.TableClientSettings) -> datetime | str: if table_client_settings is not None and table_client_settings._native_date_in_result_sets: return _parse_tz(x) return x -def _to_tz_date(pb: ydb_value_pb2.Value, value: typing.Union[datetime, str]) -> None: +def _to_tz_date(pb: ydb_value_pb2.Value, value: datetime | str) -> None: if isinstance(value, datetime): pb.text_value = value.strftime("%Y-%m-%d") + "," + _tz_name(value) else: pb.text_value = value -def _from_tz_datetime(x: str, table_client_settings: table.TableClientSettings) -> typing.Union[datetime, str]: +def _from_tz_datetime(x: str, table_client_settings: table.TableClientSettings) -> datetime | str: if table_client_settings is not None and table_client_settings._native_datetime_in_result_sets: return _parse_tz(x) return x -def _to_tz_datetime(pb: ydb_value_pb2.Value, value: typing.Union[datetime, str]) -> None: +def _to_tz_datetime(pb: ydb_value_pb2.Value, value: datetime | str) -> None: if isinstance(value, datetime): pb.text_value = value.strftime("%Y-%m-%dT%H:%M:%S") + "," + _tz_name(value) else: pb.text_value = value -def _from_tz_timestamp(x: str, table_client_settings: table.TableClientSettings) -> typing.Union[datetime, str]: +def _from_tz_timestamp(x: str, table_client_settings: table.TableClientSettings) -> datetime | str: if table_client_settings is not None and table_client_settings._native_timestamp_in_result_sets: return _parse_tz(x) return x -def _to_tz_timestamp(pb: ydb_value_pb2.Value, value: typing.Union[datetime, str]) -> None: +def _to_tz_timestamp(pb: ydb_value_pb2.Value, value: datetime | str) -> None: if isinstance(value, datetime): # isoformat() matches YDB's canonical form: 6-digit microseconds when # non-zero, omitted when zero (YDB strips a trailing ".000000"). @@ -143,7 +141,7 @@ def _to_tz_timestamp(pb: ydb_value_pb2.Value, value: typing.Union[datetime, str] pb.text_value = value -def _from_json(x: typing.Union[str, bytearray, bytes], table_client_settings: table.TableClientSettings) -> typing.Any: +def _from_json(x: str | bytearray | bytes, table_client_settings: table.TableClientSettings) -> typing.Any: if table_client_settings is not None and table_client_settings._native_json_in_result_sets: return json.loads(x) return x @@ -162,30 +160,26 @@ def _timedelta_to_microseconds(value: timedelta) -> int: return (value.days * _SECONDS_IN_DAY + value.seconds) * 1000000 + value.microseconds -def _from_interval( - value_pb: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings -) -> typing.Union[timedelta, int]: +def _from_interval(value_pb: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings) -> timedelta | int: if table_client_settings is not None and table_client_settings._native_interval_in_result_sets: return timedelta(microseconds=value_pb.int64_value) return value_pb.int64_value -def _to_interval(pb: ydb_value_pb2.Value, value: typing.Union[timedelta, int]) -> None: +def _to_interval(pb: ydb_value_pb2.Value, value: timedelta | int) -> None: if isinstance(value, timedelta): pb.int64_value = _timedelta_to_microseconds(value) else: pb.int64_value = value -def _from_timestamp( - value_pb: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings -) -> typing.Union[datetime, int]: +def _from_timestamp(value_pb: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings) -> datetime | int: if table_client_settings is not None and table_client_settings._native_timestamp_in_result_sets: return _EPOCH + timedelta(microseconds=value_pb.uint64_value) return value_pb.uint64_value -def _to_timestamp(pb: ydb_value_pb2.Value, value: typing.Union[datetime, int]) -> None: +def _to_timestamp(pb: ydb_value_pb2.Value, value: datetime | int) -> None: if isinstance(value, datetime): if value.tzinfo: epoch = _EPOCH_UTC @@ -198,13 +192,13 @@ def _to_timestamp(pb: ydb_value_pb2.Value, value: typing.Union[datetime, int]) - def _from_timestamp64( value_pb: ydb_value_pb2.Value, table_client_settings: table.TableClientSettings -) -> typing.Union[datetime, int]: +) -> datetime | int: if table_client_settings is not None and table_client_settings._native_timestamp_in_result_sets: return _EPOCH + timedelta(microseconds=value_pb.int64_value) return value_pb.int64_value -def _to_timestamp64(pb: ydb_value_pb2.Value, value: typing.Union[datetime, int]) -> None: +def _to_timestamp64(pb: ydb_value_pb2.Value, value: datetime | int) -> None: if isinstance(value, datetime): if value.tzinfo: epoch = _EPOCH_UTC @@ -313,9 +307,9 @@ class PrimitiveType(enum.Enum): def __init__( self, idn: ydb_value_pb2.Type.PrimitiveTypeId, - proto_field: typing.Optional[str], - to_obj: typing.Optional[typing.Callable[..., typing.Any]] = None, - from_obj: typing.Optional[typing.Callable[..., None]] = None, + proto_field: str | None, + to_obj: typing.Callable[..., typing.Any] | None = None, + from_obj: typing.Callable[..., None] | None = None, ) -> None: self._idn_ = idn self._to_obj = to_obj @@ -365,9 +359,7 @@ def proto(self) -> ydb_value_pb2.Type: class DataQuery(object): __slots__ = ("yql_text", "parameters_types", "name") - def __init__( - self, query_id: str, parameters_types: "dict[str, ydb_value_pb2.Type]", name: typing.Optional[str] = None - ): + def __init__(self, query_id: str, parameters_types: "dict[str, ydb_value_pb2.Type]", name: str | None = None): self.yql_text = query_id self.parameters_types = parameters_types self.name = _utilities.get_query_hash(self.yql_text) if name is None else name @@ -450,7 +442,7 @@ def __str__(self) -> str: class OptionalType(AbstractTypeBuilder): __slots__ = ("_repr", "_proto", "_item") - def __init__(self, optional_type: typing.Union[AbstractTypeBuilder, PrimitiveType]) -> None: + def __init__(self, optional_type: AbstractTypeBuilder | PrimitiveType) -> None: """ Represents optional type that wraps inner type :param optional_type: An instance of an inner type @@ -461,7 +453,7 @@ def __init__(self, optional_type: typing.Union[AbstractTypeBuilder, PrimitiveTyp self._proto.optional_type.MergeFrom(_apis.ydb_value.OptionalType(item=optional_type.proto)) @property - def item(self) -> typing.Union[AbstractTypeBuilder, PrimitiveType]: + def item(self) -> AbstractTypeBuilder | PrimitiveType: return self._item @property @@ -484,7 +476,7 @@ def __str__(self) -> str: class ListType(AbstractTypeBuilder): __slots__ = ("_repr", "_proto") - def __init__(self, list_type: typing.Union[AbstractTypeBuilder, PrimitiveType]) -> None: + def __init__(self, list_type: AbstractTypeBuilder | PrimitiveType) -> None: """ :param list_type: List item type builder """ @@ -508,8 +500,8 @@ class DictType(AbstractTypeBuilder): def __init__( self, - key_type: typing.Union[AbstractTypeBuilder, PrimitiveType], - payload_type: typing.Union[AbstractTypeBuilder, PrimitiveType], + key_type: AbstractTypeBuilder | PrimitiveType, + payload_type: AbstractTypeBuilder | PrimitiveType, ) -> None: """ :param key_type: Key type builder @@ -536,7 +528,7 @@ class SetType(AbstractTypeBuilder): def __init__( self, - key_type: typing.Union[AbstractTypeBuilder, PrimitiveType], + key_type: AbstractTypeBuilder | PrimitiveType, ) -> None: """ :param key_type: Key type builder @@ -561,10 +553,10 @@ class TupleType(AbstractTypeBuilder): __slots__ = ("__elements_repr", "__proto") def __init__(self) -> None: - self.__elements_repr: typing.List[str] = [] + self.__elements_repr: list[str] = [] self.__proto = _apis.ydb_value.Type(tuple_type=_apis.ydb_value.TupleType()) - def add_element(self, element_type: typing.Union[AbstractTypeBuilder, PrimitiveType]) -> "TupleType": + def add_element(self, element_type: AbstractTypeBuilder | PrimitiveType) -> "TupleType": """ :param element_type: Adds additional element of tuple :return: self @@ -586,10 +578,10 @@ class StructType(AbstractTypeBuilder): __slots__ = ("__members_repr", "__proto") def __init__(self) -> None: - self.__members_repr: typing.List[str] = [] + self.__members_repr: list[str] = [] self.__proto = _apis.ydb_value.Type(struct_type=_apis.ydb_value.StructType()) - def add_member(self, name: str, member_type: typing.Union[AbstractTypeBuilder, PrimitiveType]) -> "StructType": + def add_member(self, name: str, member_type: AbstractTypeBuilder | PrimitiveType) -> "StructType": """ :param name: :param member_type: @@ -613,12 +605,10 @@ class BulkUpsertColumns(AbstractTypeBuilder): __slots__ = ("__columns_repr", "__proto") def __init__(self) -> None: - self.__columns_repr: typing.List[str] = [] + self.__columns_repr: list[str] = [] self.__proto = _apis.ydb_value.Type(struct_type=_apis.ydb_value.StructType()) - def add_column( - self, name: str, column_type: typing.Union[AbstractTypeBuilder, PrimitiveType] - ) -> "BulkUpsertColumns": + def add_column(self, name: str, column_type: AbstractTypeBuilder | PrimitiveType) -> "BulkUpsertColumns": """ :param name: A column name :param column_type: A column type @@ -640,4 +630,4 @@ def __str__(self) -> str: @dataclass class TypedValue: value: typing.Any - value_type: typing.Optional[typing.Union[PrimitiveType, AbstractTypeBuilder]] = None + value_type: PrimitiveType | AbstractTypeBuilder | None = None