diff --git a/CHANGELOG.md b/CHANGELOG.md index 04c5304a..c18e547a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,5 @@ +* Mark the package as typed so type checkers use the SDK's inline annotations + ## 3.32.0 ## * Add `TableClient.read_rows` (sync and async) to read rows by primary key without a transaction * Add the `ydb.query.session.closed` counter for query session pool closures, labeled by pool name and a standardized closure reason; metrics-enabled clients now advertise `ydb-sdk-metrics/0.2.0` in `x-ydb-sdk-build-info` diff --git a/MANIFEST.in b/MANIFEST.in index f9bd1455..5806adf8 100644 --- a/MANIFEST.in +++ b/MANIFEST.in @@ -1 +1,2 @@ include requirements.txt +include ydb/py.typed diff --git a/pyproject.toml b/pyproject.toml index bac01b5a..75b1d98e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,8 +9,8 @@ strict_equality = true show_error_codes = true # Roadmap: these can be enabled after fixing the errors -no_implicit_optional = false # 24 errors -disallow_untyped_decorators = false # 10 errors +no_implicit_optional = true +disallow_untyped_decorators = true check_untyped_defs = false # 58 errors warn_return_any = false # 62 errors @@ -50,6 +50,16 @@ ignore_missing_imports = true module = ["ydb.types", "ydb.table"] disable_error_code = ["attr-defined"] +# Generated protobuf stubs are checked when they are generated, not as SDK source. +[[tool.mypy.overrides]] +module = [ + "ydb._grpc.v3.*", + "ydb._grpc.v4.*", + "ydb._grpc.v5.*", + "ydb._grpc.v6.*", +] +ignore_errors = true + [tool.coverage.run] branch = true source = ["ydb"] diff --git a/setup.py b/setup.py index a2273279..6465139d 100644 --- a/setup.py +++ b/setup.py @@ -23,6 +23,7 @@ long_description=long_description, long_description_content_type='text/markdown', packages=setuptools.find_packages("."), + package_data={"ydb": ["py.typed"]}, classifiers=[ "Programming Language :: Python :: 3", "Programming Language :: Python :: 3.10", @@ -31,6 +32,7 @@ "Programming Language :: Python :: 3.13", "Programming Language :: Python :: 3.14", "Programming Language :: Python :: 3 :: Only", + "Typing :: Typed", ], python_requires=">=3.10", install_requires=requirements, # requirements.txt diff --git a/test-requirements.txt b/test-requirements.txt index 84b71649..e52a7db9 100644 --- a/test-requirements.txt +++ b/test-requirements.txt @@ -1,6 +1,7 @@ attrs==21.2.0 bcrypt==3.2.0 black==25.11.0 +build cached-property==1.5.2 certifi==2024.7.4 cffi>=1.17.0 diff --git a/ydb/_utilities.py b/ydb/_utilities.py index df152987..66888bf4 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 Any, Callable, Dict, List, Optional, TYPE_CHECKING, TypeVar, cast from . import ydb_version import typing @@ -32,6 +32,8 @@ _grpcs_protocol = "grpcs://" _grpc_protocol = "grpc://" +CallableT = TypeVar("CallableT", bound=Callable[..., Any]) + def wrap_result_in_future(result): f = futures.Future() @@ -82,7 +84,7 @@ def parse_connection_string(connection_string): # Decorator that ensures no exceptions are leaked from decorated async call -def wrap_async_call_exceptions(f): +def wrap_async_call_exceptions(f: CallableT) -> CallableT: @functools.wraps(f) def decorator(*args, **kwargs): try: @@ -90,7 +92,7 @@ def decorator(*args, **kwargs): except Exception as e: return wrap_exception_in_future(e) - return decorator + return cast(CallableT, decorator) def check_module_exists(path: str) -> bool: diff --git a/ydb/_utilities_test.py b/ydb/_utilities_test.py index df48bb5f..34a669b8 100644 --- a/ydb/_utilities_test.py +++ b/ydb/_utilities_test.py @@ -1,5 +1,8 @@ +from pathlib import Path, PurePosixPath import subprocess import sys +import tarfile +import zipfile import pytest @@ -46,3 +49,26 @@ def test_iam_is_loaded_lazily(): output = subprocess.check_output([sys.executable, "-c", code], text=True) assert output.splitlines() == ["True", "False", "False", "True"] + + +@pytest.fixture(scope="module") +def built_distributions(tmp_path_factory): + dist_dir = tmp_path_factory.mktemp("dist") + subprocess.run( + [sys.executable, "-m", "build", "--outdir", str(dist_dir)], + cwd=Path(__file__).resolve().parent.parent, + check=True, + ) + return dist_dir + + +def test_py_typed_is_in_wheel(built_distributions): + wheel_path = next(built_distributions.glob("*.whl")) + with zipfile.ZipFile(wheel_path) as wheel: + assert "ydb/py.typed" in wheel.namelist() + + +def test_py_typed_is_in_sdist(built_distributions): + sdist_path = next(built_distributions.glob("*.tar.gz")) + with tarfile.open(sdist_path) as sdist: + assert any(PurePosixPath(member.name).parts[-2:] == ("ydb", "py.typed") for member in sdist.getmembers()) diff --git a/ydb/aio/query/session.py b/ydb/aio/query/session.py index fa7201ab..d369ad72 100644 --- a/ydb/aio/query/session.py +++ b/ydb/aio/query/session.py @@ -135,9 +135,9 @@ def transaction(self, tx_mode=None) -> QueryTxContext: async def execute( self, query: str, - parameters: dict = None, - syntax: base.QuerySyntax = None, - exec_mode: base.QueryExecMode = None, + parameters: Optional[dict] = None, + syntax: Optional[base.QuerySyntax] = None, + exec_mode: Optional[base.QueryExecMode] = None, concurrent_result_sets: bool = False, settings: Optional[BaseRequestSettings] = None, *, diff --git a/ydb/aio/table.py b/ydb/aio/table.py index decd1d37..ef9c53a6 100644 --- a/ydb/aio/table.py +++ b/ydb/aio/table.py @@ -528,7 +528,7 @@ def __init__(self, driver: "ydb.aio.Driver", size: int, min_pool_size: int = 0): self._min_pool_tasks.append(asyncio.ensure_future(self._init_and_put(self._init_session_timeout))) async def retry_operation( - self, callee: typing.Callable, *args, retry_settings: table.RetrySettings = None, **kwargs + self, callee: typing.Callable, *args, retry_settings: typing.Optional[table.RetrySettings] = None, **kwargs ): if retry_settings is None: @@ -561,7 +561,9 @@ 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: typing.Optional[int] = None + ) -> typing.Optional[ydb.ISession]: """ :param retry_num: Number of retries. If None - retries until success. :return: @@ -772,5 +774,5 @@ async def __aexit__(self, exc_type, exc, tb): async def wait_until_min_size(self): await asyncio.gather(*self._min_pool_tasks) - def checkout(self, timeout: float = None, retry_timeout: float = None): + def checkout(self, timeout: typing.Optional[float] = None, retry_timeout: typing.Optional[float] = None): return SessionCheckout(self, timeout, retry_timeout=retry_timeout) diff --git a/ydb/observability/metrics.py b/ydb/observability/metrics.py index 63cc7008..dbc0f835 100644 --- a/ydb/observability/metrics.py +++ b/ydb/observability/metrics.py @@ -23,7 +23,7 @@ import functools import inspect import weakref -from typing import Any, Callable, Dict, Iterable, List, Optional, Protocol, Tuple +from typing import Any, Callable, Dict, Iterable, List, Optional, Protocol, Tuple, TypeVar, cast from ydb.observability._endpoint import split_endpoint @@ -67,6 +67,7 @@ ) ATTEMPT_BUCKETS = (1, 2, 3, 4, 5, 7, 10, 20) _UNKNOWN_POOL = "unknown" +CallableT = TypeVar("CallableT", bound=Callable[..., Any]) _pool_name_counter = itertools.count(1) _pool_metrics_counter = itertools.count(1) _OPERATION_ATTR_KEYS = frozenset( @@ -736,7 +737,7 @@ def finish(self) -> None: self._provider.record(RETRY_ATTEMPTS, self._attempts) -def observe_retry_metrics(retry_func: Callable) -> Callable: +def observe_retry_metrics(retry_func: CallableT) -> CallableT: """Decorator recording retry duration and attempt count around a retry helper. Wraps the retried callee to count attempts and times the whole operation — but only @@ -756,7 +757,7 @@ async def awrapper(callee, retry_settings=None, *args, **kwargs): finally: metrics.finish() - return awrapper + return cast(CallableT, awrapper) @functools.wraps(retry_func) def wrapper(callee, retry_settings=None, *args, **kwargs): @@ -769,4 +770,4 @@ def wrapper(callee, retry_settings=None, *args, **kwargs): finally: metrics.finish() - return wrapper + return cast(CallableT, wrapper) diff --git a/ydb/py.typed b/ydb/py.typed new file mode 100644 index 00000000..e69de29b diff --git a/ydb/query/base.py b/ydb/query/base.py index a78825dc..2a592aad 100644 --- a/ydb/query/base.py +++ b/ydb/query/base.py @@ -10,7 +10,9 @@ Callable, List, DefaultDict, + TypeVar, Union, + cast, ) from .._grpc.grpcwrapper import ydb_query @@ -31,6 +33,8 @@ from .transaction import BaseQueryTxContext from .session import BaseQuerySession +CallableT = TypeVar("CallableT", bound=Callable[..., Any]) + class QuerySyntax(enum.IntEnum): UNSPECIFIED = 0 @@ -222,7 +226,7 @@ def create_execute_query_request( raise issues.ClientInternalError("Unable to prepare execute request") from e -def bad_session_handler(func): +def bad_session_handler(func: CallableT) -> CallableT: @functools.wraps(func) def decorator(rpc_state, response_pb, session: "BaseQuerySession", *args, **kwargs): try: @@ -231,12 +235,12 @@ def decorator(rpc_state, response_pb, session: "BaseQuerySession", *args, **kwar session._close_session(invalidate=True, reason="bad_session") raise - return decorator + return cast(CallableT, decorator) @bad_session_handler def wrap_execute_query_response( - rpc_state: RpcState, + rpc_state: Optional[RpcState], response_pb: _apis.ydb_query.ExecuteQueryResponsePart, session: "BaseQuerySession", tx: Optional["BaseQueryTxContext"] = None, diff --git a/ydb/query/session.py b/ydb/query/session.py index 4a2d7bda..70cb858b 100644 --- a/ydb/query/session.py +++ b/ydb/query/session.py @@ -503,9 +503,9 @@ def transaction(self, tx_mode: Optional[base.BaseQueryTxMode] = None) -> QueryTx def execute( self, query: str, - parameters: dict = None, - syntax: base.QuerySyntax = None, - exec_mode: base.QueryExecMode = None, + parameters: Optional[dict] = None, + syntax: Optional[base.QuerySyntax] = None, + exec_mode: Optional[base.QueryExecMode] = None, concurrent_result_sets: bool = False, settings: Optional[BaseRequestSettings] = None, *, @@ -578,7 +578,7 @@ def execute( def explain( self, query: str, - parameters: dict = None, + parameters: Optional[dict] = None, *, result_format: QueryExplainResultFormat = QueryExplainResultFormat.STR, ) -> Union[str, Dict[str, Any]]: diff --git a/ydb/query/transaction.py b/ydb/query/transaction.py index 692b2a4c..0008ac71 100644 --- a/ydb/query/transaction.py +++ b/ydb/query/transaction.py @@ -9,7 +9,10 @@ Iterable, Optional, TYPE_CHECKING, + TypeVar, Union, + Callable, + cast, overload, ) @@ -32,6 +35,7 @@ from ..aio.driver import Driver as AsyncDriver logger = logging.getLogger(__name__) +CallableT = TypeVar("CallableT", bound=Callable[..., Any]) class QueryTxStateEnum(enum.Enum): @@ -77,7 +81,7 @@ def terminal(cls, state: QueryTxStateEnum) -> bool: return len(cls._VALID_TRANSITIONS[state]) == 0 -def reset_tx_id_handler(func): +def reset_tx_id_handler(func: CallableT) -> CallableT: @functools.wraps(func) def decorator(rpc_state, response_pb, session: "BaseQuerySession", tx_state: "QueryTxState", *args, **kwargs): try: @@ -87,7 +91,7 @@ def decorator(rpc_state, response_pb, session: "BaseQuerySession", tx_state: "Qu tx_state.tx_id = None raise - return decorator + return cast(CallableT, decorator) class QueryTxState: diff --git a/ydb/retries.py b/ydb/retries.py index 4e352b5d..80bd66c1 100644 --- a/ydb/retries.py +++ b/ydb/retries.py @@ -3,13 +3,15 @@ import inspect import random import time -from typing import Any, Callable, Generator, Optional, Union +from typing import Any, Callable, Generator, Optional, TypeVar, Union, cast from . import issues from ._errors import check_retriable_error from .observability.metrics import observe_retry_metrics from .observability.tracing import SpanName, create_span as _create_span +CallableT = TypeVar("CallableT", bound=Callable[..., Any]) + def _try_span_attrs(backoff_ms: Optional[int]): return {"ydb.retry.backoff_ms": backoff_ms} if backoff_ms is not None else None @@ -234,7 +236,7 @@ def ydb_retry( slow_backoff_settings: Optional[BackoffSettings] = None, idempotent: bool = False, retry_cancelled: bool = False, -) -> Callable[[Callable[..., Any]], Callable[..., Any]]: +) -> Callable[[CallableT], CallableT]: """ Decorator for automatic function retry in case of YDB errors. @@ -252,7 +254,7 @@ def ydb_retry( :param retry_cancelled: Whether to retry cancelled operations (default: False) """ - def decorator(func: Callable[..., Any]) -> Callable[..., Any]: + def decorator(func: CallableT) -> CallableT: retry_settings = RetrySettings( max_retries=max_retries, max_session_acquire_timeout=max_session_acquire_timeout, @@ -272,13 +274,13 @@ def decorator(func: Callable[..., Any]) -> Callable[..., Any]: async def async_wrapper(*args: Any, **kwargs: Any) -> Any: return await retry_operation_async(func, retry_settings, *args, **kwargs) - return async_wrapper + return cast(CallableT, async_wrapper) else: @functools.wraps(func) def sync_wrapper(*args: Any, **kwargs: Any) -> Any: return retry_operation_sync(func, retry_settings, *args, **kwargs) - return sync_wrapper + return cast(CallableT, sync_wrapper) return decorator diff --git a/ydb/table.py b/ydb/table.py index dad5788f..966bae82 100644 --- a/ydb/table.py +++ b/ydb/table.py @@ -1,5 +1,6 @@ # -*- coding: utf-8 -*- import abc +from concurrent import futures from dataclasses import dataclass import ydb from abc import abstractmethod @@ -12,9 +13,11 @@ Dict, Generic, List, + Mapping, Optional, Tuple, TYPE_CHECKING, + Union, ) from ._typing import DriverT @@ -1222,7 +1225,7 @@ def session(self): return Session(self._driver, self._table_client_settings) def scan_query(self, query, parameters=None, settings=None): - # type: (ydb.ScanQuery, tuple, ydb.BaseRequestSettings) -> _utilities.SyncResponseIterator + # type: (Union[str, ydb.ScanQuery], Optional[Mapping[str, Any]], Optional[ydb.BaseRequestSettings]) -> _utilities.SyncResponseIterator request = _scan_query_request_factory(query, parameters, settings) stream_it = self._driver( request, @@ -1236,7 +1239,7 @@ def scan_query(self, query, parameters=None, settings=None): ) def bulk_upsert(self, table_path, rows, column_types, settings=None): - # type: (str, list, typing.Union[ydb.AbstractTypeBuilder, ydb.PrimitiveType], ydb.BaseRequestSettings) -> Any + # type: (str, list, typing.Union[ydb.AbstractTypeBuilder, ydb.PrimitiveType], Optional[ydb.BaseRequestSettings]) -> Any """ Bulk upsert data @@ -1255,7 +1258,7 @@ def bulk_upsert(self, table_path, rows, column_types, settings=None): ) def read_rows(self, table_path, keys, key_types, columns=None, settings=None): - # type: (str, list, ydb.AbstractTypeBuilder, typing.Optional[list], ydb.BaseRequestSettings) -> Any + # type: (str, list, ydb.AbstractTypeBuilder, typing.Optional[list], Optional[ydb.BaseRequestSettings]) -> Any """ Read specified keys non-transactionally from a single table. @@ -1277,7 +1280,7 @@ def read_rows(self, table_path, keys, key_types, columns=None, settings=None): ) def describe_system_view(self, path, settings=None): - # type: (str, ydb.BaseRequestSettings) -> Any + # type: (str, Optional[ydb.BaseRequestSettings]) -> Any """ Returns a full description of a system view by the provided path. @@ -1305,7 +1308,7 @@ def __del__(self): self._stop_pool_if_needed() def async_scan_query(self, query, parameters=None, settings=None): - # type: (ydb.ScanQuery, tuple, ydb.BaseRequestSettings) -> _utilities.AsyncResponseIterator + # type: (Union[str, ydb.ScanQuery], Optional[Mapping[str, Any]], Optional[ydb.BaseRequestSettings]) -> _utilities.AsyncResponseIterator request = _scan_query_request_factory(query, parameters, settings) stream_it = self._driver( request, @@ -1320,7 +1323,7 @@ def async_scan_query(self, query, parameters=None, settings=None): @_utilities.wrap_async_call_exceptions def async_bulk_upsert(self, table_path, rows, column_types, settings=None): - # type: (str, list, typing.Union[ydb.AbstractTypeBuilder, ydb.PrimitiveType], ydb.BaseRequestSettings) -> None + # type: (str, list, typing.Union[ydb.AbstractTypeBuilder, ydb.PrimitiveType], Optional[ydb.BaseRequestSettings]) -> futures.Future[ydb.Operation] return self._driver.future( _session_impl.bulk_upsert_request_factory(table_path, rows, column_types), _apis.TableService.Stub, @@ -1332,7 +1335,7 @@ def async_bulk_upsert(self, table_path, rows, column_types, settings=None): @_utilities.wrap_async_call_exceptions def async_read_rows(self, table_path, keys, key_types, columns=None, settings=None): - # type: (str, list, ydb.AbstractTypeBuilder, typing.Optional[list], ydb.BaseRequestSettings) -> Any + # type: (str, list, ydb.AbstractTypeBuilder, typing.Optional[list], Optional[ydb.BaseRequestSettings]) -> Any return self._driver.future( _session_impl.read_rows_request_factory(table_path, keys, key_types, columns), _apis.TableService.Stub, diff --git a/ydb/topic.py b/ydb/topic.py index 98859293..e76956ab 100644 --- a/ydb/topic.py +++ b/ydb/topic.py @@ -331,7 +331,7 @@ def writer( topic, *, producer_id: Optional[str] = None, # default - random - session_metadata: Mapping[str, str] = None, + session_metadata: Optional[Mapping[str, str]] = None, partition_id: Union[int, None] = None, auto_seqno: bool = True, auto_created_at: bool = True, @@ -363,7 +363,7 @@ def tx_writer( topic, *, producer_id: Optional[str] = None, # default - random - session_metadata: Mapping[str, str] = None, + session_metadata: Optional[Mapping[str, str]] = None, partition_id: Union[int, None] = None, auto_seqno: bool = True, auto_created_at: bool = True, @@ -665,7 +665,7 @@ def writer( topic, *, producer_id: Optional[str] = None, # default - random - session_metadata: Mapping[str, str] = None, + session_metadata: Optional[Mapping[str, str]] = None, partition_id: Union[int, None] = None, auto_seqno: bool = True, auto_created_at: bool = True, @@ -698,7 +698,7 @@ def tx_writer( topic, *, producer_id: Optional[str] = None, # default - random - session_metadata: Mapping[str, str] = None, + session_metadata: Optional[Mapping[str, str]] = None, partition_id: Union[int, None] = None, auto_seqno: bool = True, auto_created_at: bool = True,