Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
* 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`
Expand Down
1 change: 1 addition & 0 deletions MANIFEST.in
Original file line number Diff line number Diff line change
@@ -1 +1,2 @@
include requirements.txt
include ydb/py.typed
14 changes: 12 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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"]
Expand Down
2 changes: 2 additions & 0 deletions setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
long_description=long_description,
long_description_content_type='text/markdown',
packages=setuptools.find_packages("."),
package_data={"ydb": ["py.typed"]},
Comment thread
0x4e3 marked this conversation as resolved.
classifiers=[
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
Expand All @@ -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
Expand Down
1 change: 1 addition & 0 deletions test-requirements.txt
Original file line number Diff line number Diff line change
@@ -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
Expand Down
8 changes: 5 additions & 3 deletions ydb/_utilities.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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()
Expand Down Expand Up @@ -82,15 +84,15 @@ 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:
return f(*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:
Expand Down
26 changes: 26 additions & 0 deletions ydb/_utilities_test.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,8 @@
from pathlib import Path, PurePosixPath
import subprocess
import sys
import tarfile
import zipfile

import pytest

Expand Down Expand Up @@ -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())
6 changes: 3 additions & 3 deletions ydb/aio/query/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
*,
Expand Down
8 changes: 5 additions & 3 deletions ydb/aio/table.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
9 changes: 5 additions & 4 deletions ydb/observability/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand All @@ -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):
Expand All @@ -769,4 +770,4 @@ def wrapper(callee, retry_settings=None, *args, **kwargs):
finally:
metrics.finish()

return wrapper
return cast(CallableT, wrapper)
Empty file added ydb/py.typed
Empty file.
10 changes: 7 additions & 3 deletions ydb/query/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,9 @@
Callable,
List,
DefaultDict,
TypeVar,
Union,
cast,
)

from .._grpc.grpcwrapper import ydb_query
Expand All @@ -31,6 +33,8 @@
from .transaction import BaseQueryTxContext
from .session import BaseQuerySession

CallableT = TypeVar("CallableT", bound=Callable[..., Any])


class QuerySyntax(enum.IntEnum):
UNSPECIFIED = 0
Expand Down Expand Up @@ -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:
Expand All @@ -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,
Expand Down
8 changes: 4 additions & 4 deletions ydb/query/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
*,
Expand Down Expand Up @@ -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]]:
Expand Down
8 changes: 6 additions & 2 deletions ydb/query/transaction.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,10 @@
Iterable,
Optional,
TYPE_CHECKING,
TypeVar,
Union,
Callable,
cast,
overload,
)

Expand All @@ -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):
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down
12 changes: 7 additions & 5 deletions ydb/retries.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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,
Expand All @@ -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
Loading
Loading