From 2b884c5240822d3127e64f6a120398babf7ab406 Mon Sep 17 00:00:00 2001 From: Adrian Chaves Date: Fri, 25 Sep 2026 23:36:11 +0200 Subject: [PATCH] Enable mypy strict mode --- pyproject.toml | 9 ++---- tests/conftest.py | 14 ++++++-- tests/mockserver.py | 73 ++++++++++++++++++++++++++++++------------ tests/test_apikey.py | 32 ++++++++++++------ tests/test_async.py | 49 ++++++++++++++++++++-------- tests/test_auth.py | 32 ++++++++++++------ tests/test_main.py | 53 ++++++++++++++++-------------- tests/test_retry.py | 56 +++++++++++++++++++------------- tests/test_sync.py | 10 +++--- tests/test_utils.py | 18 +++++++---- tests/test_x402.py | 34 ++++++++++++-------- zyte_api/_retry.py | 20 +++++++----- zyte_api/_x402.py | 5 +-- zyte_api/aio/client.py | 5 ++- zyte_api/aio/retry.py | 2 ++ zyte_api/errors.py | 6 ++-- 16 files changed, 270 insertions(+), 148 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index ec77d56..81c2d37 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -89,17 +89,14 @@ patch = [ ] [tool.mypy] -allow_untyped_defs = false -implicit_reexport = false +strict = true +# Unannotated third-party functions. +untyped_calls_exclude = ["twisted"] [[tool.mypy.overrides]] module = "runstats" ignore_missing_imports = true -[[tool.mypy.overrides]] -module = "tests.*" -allow_untyped_defs = true - [tool.pytest.ini_options] filterwarnings = [ "ignore:The zyte_api\\.aio module is deprecated:DeprecationWarning" diff --git a/tests/conftest.py b/tests/conftest.py index 9f37446..41d18ed 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,8 +1,18 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + import pytest +if TYPE_CHECKING: + from collections.abc import Generator + from pathlib import Path + + from .mockserver import MockServer + @pytest.fixture(autouse=True) -def isolated_apikey_env(tmp_path, monkeypatch): +def isolated_apikey_env(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: """Keep API-key resolution hermetic: drop ambient key env vars and run from an empty directory so ``find_dotenv()`` can't pick up a stray ``.env`` from the developer's working tree. Tests that need a ``.env`` create it in the @@ -13,7 +23,7 @@ def isolated_apikey_env(tmp_path, monkeypatch): @pytest.fixture(scope="session") -def mockserver(): +def mockserver() -> Generator[MockServer]: from .mockserver import MockServer # noqa: PLC0415 with MockServer() as server: diff --git a/tests/mockserver.py b/tests/mockserver.py index f6d1ec0..9c31f46 100644 --- a/tests/mockserver.py +++ b/tests/mockserver.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import argparse import json import socket @@ -7,12 +9,19 @@ from collections import defaultdict from importlib import import_module from subprocess import PIPE, Popen -from typing import Any +from typing import TYPE_CHECKING, Any from urllib.parse import urlparse from twisted.internet.task import deferLater from twisted.web.resource import Resource -from twisted.web.server import NOT_DONE_YET, Site +from twisted.web.server import NOT_DONE_YET, Request, Site + +if TYPE_CHECKING: + from collections.abc import Callable + from types import TracebackType + + from twisted.internet.defer import Deferred + from typing_extensions import Self SCREENSHOT = ( "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAACklEQVR4nGMAAQAABQABDQott" @@ -21,7 +30,13 @@ # https://github.com/scrapy/scrapy/blob/02b97f98e74a994ad3e4d74e7ed55207e508a576/tests/mockserver.py#L27C1-L33C19 -def getarg(request, name, default=None, type_=None): +def getarg( + request: Request, + name: bytes, + default: Any = None, + type_: Callable[[bytes], Any] | None = None, +) -> Any: + assert request.args is not None if name in request.args: value = request.args[name][0] if type_ is not None: @@ -30,36 +45,44 @@ def getarg(request, name, default=None, type_=None): return default -def get_ephemeral_port(): +def get_ephemeral_port() -> int: s = socket.socket() s.bind(("", 0)) - return s.getsockname()[1] + port: int = s.getsockname()[1] + return port class DropResource(Resource): isLeaf = True - def deferRequest(self, request, delay, f, *a, **kw): + def deferRequest( + self, + request: Request, + delay: float, + f: Callable[..., Any], + *a: Any, + **kw: Any, + ) -> Deferred[Any]: from twisted.internet import reactor - def _cancelrequest(_): + def _cancelrequest(_: Any) -> None: # silence CancelledError d.addErrback(lambda _: None) d.cancel() - d = deferLater(reactor, delay, f, *a, **kw) + d = deferLater(reactor, delay, f, *a, **kw) # type: ignore[arg-type] request.notifyFinish().addErrback(_cancelrequest) return d - def render_POST(self, request): + def render_POST(self, request: Request) -> int: request.setHeader(b"Content-Length", b"1024") self.deferRequest(request, 0, self._delayedRender, request) return NOT_DONE_YET - def _delayedRender(self, request): + def _delayedRender(self, request: Request) -> None: abort = getarg(request, b"abort", 0, type_=int) request.write(b"this connection will be dropped\n") - tr = request.channel.transport + tr: Any = request.channel.transport try: if abort and hasattr(tr, "abortConnection"): tr.abortConnection() @@ -94,10 +117,10 @@ def _delayedRender(self, request): class DefaultResource(Resource): request_count = 0 - def getChild(self, path, request): + def getChild(self, path: bytes, request: Request) -> Resource: return self - def render_POST(self, request): + def render_POST(self, request: Request) -> bytes: request.responseHeaders.setRawHeaders( b"Content-Type", [b"application/json"], @@ -107,6 +130,7 @@ def render_POST(self, request): [b"abcd1234"], ) + assert request.content is not None request_data = json.loads(request.content.read()) response_data: dict[str, Any] @@ -121,7 +145,7 @@ def render_POST(self, request): return b"" if domain == "e500.example": request.setResponseCode(500) - return "" + return b"" if domain == "e520.example": request.setResponseCode(520) response_data = {"status": 520, "type": "/download/temporary-error"} @@ -235,15 +259,17 @@ def render_POST(self, request): class MockServer: - def __init__(self, resource=None, port=None): + def __init__( + self, resource: type[Resource] | None = None, port: int | None = None + ) -> None: resource = resource or DefaultResource self.resource = f"{resource.__module__}.{resource.__name__}" - self.proc = None + self.proc: Popen[bytes] | None = None self.host = socket.gethostbyname(socket.gethostname()) self.port = port or get_ephemeral_port() self.root_url = f"http://{self.host}:{self.port}" - def __enter__(self): + def __enter__(self) -> Self: self.proc = Popen( [ sys.executable, @@ -260,17 +286,22 @@ def __enter__(self): self.proc.stdout.readline() return self - def __exit__(self, exc_type, exc_value, traceback): + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: assert self.proc is not None self.proc.kill() self.proc.wait() time.sleep(0.2) - def urljoin(self, path): + def urljoin(self, path: str) -> str: return self.root_url + path -def main(): +def main() -> None: from twisted.internet import reactor parser = argparse.ArgumentParser() @@ -283,7 +314,7 @@ def main(): # Typing issue: https://github.com/twisted/twisted/issues/9909 http_port = reactor.listenTCP(args.port, Site(resource)) # type: ignore[attr-defined] - def print_listening(): + def print_listening() -> None: host = http_port.getHost() print(f"Mock server {resource} running at http://{host.host}:{host.port}") diff --git a/tests/test_apikey.py b/tests/test_apikey.py index d09d4be..38fb62b 100644 --- a/tests/test_apikey.py +++ b/tests/test_apikey.py @@ -1,4 +1,7 @@ +from __future__ import annotations + import os +from typing import TYPE_CHECKING import pytest @@ -9,8 +12,11 @@ read_dotenv_auth, ) +if TYPE_CHECKING: + from pathlib import Path + -def test_get_apikey(monkeypatch): +def test_get_apikey(monkeypatch: pytest.MonkeyPatch) -> None: assert get_apikey("a") == "a" with pytest.raises(NoApiKey): get_apikey() @@ -22,7 +28,7 @@ def test_get_apikey(monkeypatch): assert get_apikey(None) == "b" -def test_get_apikey_from_dotenv(tmp_path): +def test_get_apikey_from_dotenv(tmp_path: Path) -> None: # The autouse fixture already chdir'd into the empty tmp_path. (tmp_path / ".env").write_text("ZYTE_API_KEY=fromdotenv\n") @@ -33,7 +39,9 @@ def test_get_apikey_from_dotenv(tmp_path): assert "ZYTE_API_KEY" not in os.environ -def test_get_apikey_from_dotenv_parent_dir(tmp_path, monkeypatch): +def test_get_apikey_from_dotenv_parent_dir( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: (tmp_path / ".env").write_text("ZYTE_API_KEY=fromparent\n") subdir = tmp_path / "project" / "subdir" subdir.mkdir(parents=True) @@ -42,14 +50,16 @@ def test_get_apikey_from_dotenv_parent_dir(tmp_path, monkeypatch): assert get_apikey() == "fromparent" -def test_get_apikey_env_takes_precedence_over_dotenv(tmp_path, monkeypatch): +def test_get_apikey_env_takes_precedence_over_dotenv( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: (tmp_path / ".env").write_text("ZYTE_API_KEY=fromdotenv\n") monkeypatch.setenv("ZYTE_API_KEY", "fromenv") assert get_apikey() == "fromenv" -def test_read_apikey_from_dotenv(tmp_path): +def test_read_apikey_from_dotenv(tmp_path: Path) -> None: (tmp_path / ".env").write_text("ZYTE_API_KEY=fromdotenv\nOTHER=ignored\n") assert read_apikey_from_dotenv() == "fromdotenv" @@ -57,19 +67,19 @@ def test_read_apikey_from_dotenv(tmp_path): assert "OTHER" not in os.environ -def test_read_apikey_from_dotenv_missing(tmp_path): +def test_read_apikey_from_dotenv_missing(tmp_path: Path) -> None: # Empty working directory, no .env anywhere relevant. assert read_apikey_from_dotenv() is None -def test_read_apikey_from_dotenv_custom_path(tmp_path): +def test_read_apikey_from_dotenv_custom_path(tmp_path: Path) -> None: env_file = tmp_path / "custom.env" env_file.write_text("ZYTE_API_KEY=fromcustom\n") assert read_apikey_from_dotenv(str(env_file)) == "fromcustom" -def test_read_dotenv_auth_reads_both_credentials(tmp_path): +def test_read_dotenv_auth_reads_both_credentials(tmp_path: Path) -> None: (tmp_path / ".env").write_text( "ZYTE_API_KEY=k\nZYTE_API_ETH_KEY=e\nOTHER=ignored\n" ) @@ -79,7 +89,9 @@ def test_read_dotenv_auth_reads_both_credentials(tmp_path): assert "ZYTE_API_ETH_KEY" not in os.environ -def test_read_dotenv_auth_eth_key_not_read_from_parent(tmp_path, monkeypatch): +def test_read_dotenv_auth_eth_key_not_read_from_parent( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: # Both credentials live in a parent .env, but only the API key is read from # there; the Ethereum private key is never looked up in parent directories. (tmp_path / ".env").write_text( @@ -92,7 +104,7 @@ def test_read_dotenv_auth_eth_key_not_read_from_parent(tmp_path, monkeypatch): assert read_dotenv_auth() == {"ZYTE_API_KEY": "fromparent"} -def test_read_dotenv_auth_explicit_path_reads_eth(tmp_path): +def test_read_dotenv_auth_explicit_path_reads_eth(tmp_path: Path) -> None: # An explicit path is honored for both credentials (no walking involved). env_file = tmp_path / "custom.env" env_file.write_text("ZYTE_API_ETH_KEY=e\n") diff --git a/tests/test_async.py b/tests/test_async.py index 3526e49..85d47ac 100644 --- a/tests/test_async.py +++ b/tests/test_async.py @@ -15,6 +15,8 @@ from zyte_api.utils import USER_AGENT if TYPE_CHECKING: + from pathlib import Path + from tests.mockserver import MockServer @@ -38,7 +40,9 @@ ), ), ) -def test_user_agent(client_cls, user_agent, expected): +def test_user_agent( + client_cls: type[AsyncZyteAPI], user_agent: str | None, expected: str +) -> None: client = client_cls(api_key="123", api_url="http:\\test", user_agent=user_agent) assert client.user_agent == expected @@ -50,7 +54,7 @@ def test_user_agent(client_cls, user_agent, expected): AsyncClient, ), ) -def test_api_key(client_cls): +def test_api_key(client_cls: type[AsyncZyteAPI]) -> None: client_cls(api_key="a") with pytest.raises(NoApiKey): client_cls() @@ -63,7 +67,7 @@ def test_api_key(client_cls): AsyncClient, ), ) -def test_api_key_from_dotenv(client_cls, tmp_path): +def test_api_key_from_dotenv(client_cls: type[AsyncZyteAPI], tmp_path: Path) -> None: # The autouse fixture already chdir'd into the empty tmp_path. (tmp_path / ".env").write_text("ZYTE_API_KEY=fromdotenv\n") @@ -75,14 +79,16 @@ def test_api_key_from_dotenv(client_cls, tmp_path): @pytest.mark.asyncio -async def test_session_inherits_client_trust_env(mockserver): +async def test_session_inherits_client_trust_env(mockserver: MockServer) -> None: client = AsyncZyteAPI(api_key="a", api_url=mockserver.urljoin("/"), trust_env=True) async with client.session() as session: assert session._session._trust_env is True @pytest.mark.asyncio -async def test_get_creates_session_with_client_trust_env(mockserver): +async def test_get_creates_session_with_client_trust_env( + mockserver: MockServer, +) -> None: client = AsyncZyteAPI(api_key="a", api_url=mockserver.urljoin("/"), trust_env=True) with patch( "zyte_api._async.create_session", wraps=create_session @@ -99,7 +105,9 @@ async def test_get_creates_session_with_client_trust_env(mockserver): ), ) @pytest.mark.asyncio -async def test_get(client_cls, get_method, mockserver): +async def test_get( + client_cls: type[AsyncZyteAPI], get_method: str, mockserver: MockServer +) -> None: client = client_cls(api_key="a", api_url=mockserver.urljoin("/")) expected_result = { "url": "https://a.example", @@ -119,7 +127,9 @@ async def test_get(client_cls, get_method, mockserver): ), ) @pytest.mark.asyncio -async def test_get_request_error(client_cls, get_method, mockserver): +async def test_get_request_error( + client_cls: type[AsyncZyteAPI], get_method: str, mockserver: MockServer +) -> None: client = client_cls(api_key="a", api_url=mockserver.urljoin("/")) with pytest.raises(RequestError) as request_error_info: await getattr(client, get_method)( @@ -143,7 +153,9 @@ async def test_get_request_error(client_cls, get_method, mockserver): ), ) @pytest.mark.asyncio -async def test_get_request_error_empty_body(client_cls, get_method, mockserver): +async def test_get_request_error_empty_body( + client_cls: type[AsyncZyteAPI], get_method: str, mockserver: MockServer +) -> None: client = client_cls(api_key="a", api_url=mockserver.urljoin("/")) with pytest.raises(RequestError) as request_error_info: await getattr(client, get_method)( @@ -162,7 +174,9 @@ async def test_get_request_error_empty_body(client_cls, get_method, mockserver): ), ) @pytest.mark.asyncio -async def test_get_request_error_non_json(client_cls, get_method, mockserver): +async def test_get_request_error_non_json( + client_cls: type[AsyncZyteAPI], get_method: str, mockserver: MockServer +) -> None: client = client_cls(api_key="a", api_url=mockserver.urljoin("/")) with pytest.raises(RequestError) as request_error_info: await getattr(client, get_method)( @@ -181,7 +195,9 @@ async def test_get_request_error_non_json(client_cls, get_method, mockserver): ), ) @pytest.mark.asyncio -async def test_get_request_error_unexpected_json(client_cls, get_method, mockserver): +async def test_get_request_error_unexpected_json( + client_cls: type[AsyncZyteAPI], get_method: str, mockserver: MockServer +) -> None: client = client_cls(api_key="a", api_url=mockserver.urljoin("/")) with pytest.raises(RequestError) as request_error_info: await getattr(client, get_method)( @@ -200,7 +216,9 @@ async def test_get_request_error_unexpected_json(client_cls, get_method, mockser ), ) @pytest.mark.asyncio -async def test_iter(client_cls, iter_method, mockserver): +async def test_iter( + client_cls: type[AsyncZyteAPI], iter_method: str, mockserver: MockServer +) -> None: client = client_cls(api_key="a", api_url=mockserver.urljoin("/")) queries = [ {"url": "https://a.example", "httpResponseBody": True}, @@ -241,7 +259,12 @@ async def test_iter(client_cls, iter_method, mockserver): ), ) @pytest.mark.asyncio -async def test_semaphore(client_cls, get_method, iter_method, mockserver): +async def test_semaphore( + client_cls: type[AsyncZyteAPI], + get_method: str, + iter_method: str, + mockserver: MockServer, +) -> None: client = client_cls(api_key="a", api_url=mockserver.urljoin("/")) client._semaphore = AsyncMock(wraps=client._semaphore) queries = [ @@ -361,7 +384,7 @@ async def test_session_no_context_manager(mockserver: MockServer) -> None: assert actual_result in expected_results -def test_retrying_class(): +def test_retrying_class() -> None: """A descriptive exception is raised when creating a client with an AsyncRetrying subclass or similar instead of an instance of it.""" with pytest.raises(ValueError, match="must be an instance of AsyncRetrying"): diff --git a/tests/test_auth.py b/tests/test_auth.py index 06eab2a..b8785f6 100644 --- a/tests/test_auth.py +++ b/tests/test_auth.py @@ -1,9 +1,12 @@ +from __future__ import annotations + from base64 import b64encode from contextlib import asynccontextmanager from os import environ from pathlib import Path -from subprocess import run +from subprocess import CompletedProcess, run from tempfile import NamedTemporaryFile +from typing import TYPE_CHECKING, Any from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -13,11 +16,18 @@ from .test_x402 import HAS_X402 from .test_x402 import KEY as ETH_KEY +if TYPE_CHECKING: + from collections.abc import AsyncIterator + + from .mockserver import MockServer + ETH_KEY_2 = ETH_KEY[-1] + ETH_KEY[:-1] assert ETH_KEY_2 != ETH_KEY -def run_zyte_api(args, env, mockserver): +def run_zyte_api( + args: list[str], env: dict[str, str], mockserver: MockServer +) -> CompletedProcess[bytes]: base_env = { key: value for key, value in environ.items() @@ -58,7 +68,9 @@ def run_zyte_api(args, env, mockserver): ), ), ) -def test(scenario, expected, mockserver): +def test( + scenario: dict[str, Any], expected: dict[str, str], mockserver: MockServer +) -> None: result = run_zyte_api( scenario.get("args", []), scenario.get("env", {}), @@ -71,7 +83,7 @@ def test(scenario, expected, mockserver): assert result.returncode == 0 -def test_dotenv_cli(mockserver, tmp_path): +def test_dotenv_cli(mockserver: MockServer, tmp_path: Path) -> None: # The autouse fixture chdir'd into the empty tmp_path, so there is no key # anywhere yet. result = run_zyte_api([], {}, mockserver) @@ -92,14 +104,14 @@ def test_dotenv_cli(mockserver, tmp_path): @pytest.mark.skipif(not HAS_X402, reason="x402 extra not installed") -def test_dotenv_cli_eth_key(mockserver, tmp_path): +def test_dotenv_cli_eth_key(mockserver: MockServer, tmp_path: Path) -> None: Path(".env").write_text(f"ZYTE_API_ETH_KEY={ETH_KEY}\n", encoding="utf8") result = run_zyte_api([], {}, mockserver) assert result.returncode == 0, result.stderr @pytest.mark.skipif(not HAS_X402, reason="x402 extra not installed") -def test_dotenv_eth_key(tmp_path): +def test_dotenv_eth_key(tmp_path: Path) -> None: (tmp_path / ".env").write_text(f"ZYTE_API_ETH_KEY={ETH_KEY}\n", encoding="utf8") client = AsyncZyteAPI() @@ -151,7 +163,9 @@ def test_dotenv_eth_key(tmp_path): ), ), ) -def test_precedence(scenario, expected, monkeypatch): +def test_precedence( + scenario: dict[str, Any], expected: dict[str, str], monkeypatch: pytest.MonkeyPatch +) -> None: for key, value in scenario.get("env", {}).items(): monkeypatch.setenv(key, value) if expected["key_type"] == "eth" and not HAS_X402: @@ -178,7 +192,7 @@ def test_precedence(scenario, expected, monkeypatch): @pytest.mark.asyncio -async def test_basic_auth_header(): +async def test_basic_auth_header() -> None: api_key = "testkey" captured_headers = {} @@ -190,7 +204,7 @@ async def test_basic_auth_header(): response_mock.json = AsyncMock(return_value={"url": "https://a.example"}) @asynccontextmanager - async def fake_post(**kwargs): + async def fake_post(**kwargs: Any) -> AsyncIterator[MagicMock]: captured_headers.update(kwargs.get("headers", {})) yield response_mock diff --git a/tests/test_main.py b/tests/test_main.py index 6db6479..b178efe 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -8,7 +8,7 @@ from json import JSONDecodeError from pathlib import Path from tempfile import NamedTemporaryFile -from typing import TYPE_CHECKING, Any +from typing import IO, TYPE_CHECKING, Any from unittest.mock import AsyncMock, Mock, patch import pytest @@ -17,13 +17,13 @@ from zyte_api.__main__ import _get_argument_parser, read_input, run if TYPE_CHECKING: - from collections.abc import Iterable + from collections.abc import Awaitable, Callable, Iterable from tests.mockserver import MockServer class MockRequestError(RequestError): - def __init__(self, *args, **kwargs): + def __init__(self, *args: Any, **kwargs: Any) -> None: super().__init__( *args, query={}, @@ -34,13 +34,13 @@ def __init__(self, *args, **kwargs): ) @property - def parsed(self): + def parsed(self) -> Mock: return Mock( response_body=Mock(decode=Mock(return_value=forbidden_domain_response())) ) -def get_json_content(file_object): +def get_json_content(file_object: IO[str] | None) -> Any: if not file_object: return None @@ -61,7 +61,7 @@ def forbidden_domain_response() -> dict[str, Any]: } -async def fake_exception(value=True): +async def fake_exception(value: bool = True) -> Any: # Simulating an error condition if value: raise MockRequestError @@ -104,7 +104,12 @@ async def fake_exception(value=True): ), ) @pytest.mark.asyncio -async def test_run(queries, expected_response, store_errors, exception): +async def test_run( + queries: list[dict[str, Any]], + expected_response: dict[str, Any] | None, + store_errors: bool, + exception: Callable[..., Awaitable[Any]], +) -> None: tmp_path = Path("temporary_file.jsonl") temporary_file = tmp_path.open("w") # noqa: ASYNC230 n_conn = 5 @@ -156,7 +161,7 @@ async def test_run(queries, expected_response, store_errors, exception): @pytest.mark.asyncio -async def test_run_stop_on_errors_false(mockserver): +async def test_run_stop_on_errors_false(mockserver: MockServer) -> None: queries = [{"url": "https://exception.example", "httpResponseBody": True}] with ( NamedTemporaryFile("w") as output_file, @@ -175,7 +180,7 @@ async def test_run_stop_on_errors_false(mockserver): @pytest.mark.asyncio -async def test_run_stop_on_errors_true(mockserver): +async def test_run_stop_on_errors_true(mockserver: MockServer) -> None: query = {"url": "https://exception.example", "httpResponseBody": True} queries = [query] with ( @@ -229,7 +234,7 @@ def _run( ) -def test_empty_input(mockserver): +def test_empty_input(mockserver: MockServer) -> None: result = _run(input_="", mockserver=mockserver) assert result.returncode assert result.stdout == b"" @@ -242,7 +247,7 @@ def test_trust_env_flag_parsing() -> None: assert args.trust_env is True -def test_intype_txt_implicit(mockserver): +def test_intype_txt_implicit(mockserver: MockServer) -> None: result = _run(input_="https://a.example", mockserver=mockserver) assert not result.returncode assert ( @@ -251,7 +256,7 @@ def test_intype_txt_implicit(mockserver): ) -def test_intype_txt_explicit(mockserver): +def test_intype_txt_explicit(mockserver: MockServer) -> None: result = _run( input_="https://a.example", mockserver=mockserver, @@ -264,7 +269,7 @@ def test_intype_txt_explicit(mockserver): ) -def test_intype_jsonl_implicit(mockserver): +def test_intype_jsonl_implicit(mockserver: MockServer) -> None: result = _run( input_='{"url": "https://a.example", "browserHtml": true}', mockserver=mockserver, @@ -276,7 +281,7 @@ def test_intype_jsonl_implicit(mockserver): ) -def test_intype_jsonl_explicit(mockserver): +def test_intype_jsonl_explicit(mockserver: MockServer) -> None: result = _run( input_='{"url": "https://a.example", "browserHtml": true}', mockserver=mockserver, @@ -289,7 +294,7 @@ def test_intype_jsonl_explicit(mockserver): ) -def test_stdin(mockserver): +def test_stdin(mockserver: MockServer) -> None: result = _run( input_="https://a.example", mockserver=mockserver, @@ -300,7 +305,7 @@ def test_stdin(mockserver): assert b'"httpResponseBody"' in result.stdout -def test_params_txt(mockserver): +def test_params_txt(mockserver: MockServer) -> None: result = _run( input_="https://a.example", mockserver=mockserver, @@ -311,7 +316,7 @@ def test_params_txt(mockserver): assert b'"httpResponseBody"' in result.stdout -def test_params_jsonl(mockserver): +def test_params_jsonl(mockserver: MockServer) -> None: result = _run( input_='{"url": "https://a.example", "browserHtml": true}', mockserver=mockserver, @@ -323,14 +328,14 @@ def test_params_jsonl(mockserver): @pytest.mark.parametrize("value", ("{", "[]")) -def test_params_invalid(value, capsys): +def test_params_invalid(value: str, capsys: pytest.CaptureFixture[str]) -> None: parser = _get_argument_parser() with pytest.raises(SystemExit): parser.parse_args(["--params", value, "README.rst"]) assert "--params/-p" in capsys.readouterr().err -def test_read_input_txt(): +def test_read_input_txt() -> None: assert read_input(StringIO("https://a.example\n\n"), "txt") == [ { "url": "https://a.example", @@ -340,7 +345,7 @@ def test_read_input_txt(): ] -def test_read_input_txt_params(): +def test_read_input_txt_params() -> None: parser = _get_argument_parser() args = parser.parse_args(["-p", '{"httpResponseBody": true}', "README.rst"]) assert read_input(StringIO("https://a.example\n"), "txt", args.params) == [ @@ -352,7 +357,7 @@ def test_read_input_txt_params(): ] -def test_read_input_jl_params(): +def test_read_input_jl_params() -> None: input_fp = StringIO('{"url": "https://a.example", "browserHtml": true}\n\n') params = {"browserHtml": False, "httpResponseBody": True} assert read_input(input_fp, "jl", params) == [ @@ -365,12 +370,12 @@ def test_read_input_jl_params(): ] -def test_read_input_empty(): +def test_read_input_empty() -> None: assert read_input(StringIO(""), "txt") == [] @pytest.mark.flaky(reruns=16) -def test_limit_and_shuffle(mockserver): +def test_limit_and_shuffle(mockserver: MockServer) -> None: result = _run( input_="https://a.example\nhttps://b.example", mockserver=mockserver, @@ -383,7 +388,7 @@ def test_limit_and_shuffle(mockserver): ) -def test_run_non_json_response(mockserver): +def test_run_non_json_response(mockserver: MockServer) -> None: result = _run( input_="https://nonjson.example", mockserver=mockserver, diff --git a/tests/test_retry.py b/tests/test_retry.py index 22ba690..42fb28e 100644 --- a/tests/test_retry.py +++ b/tests/test_retry.py @@ -1,6 +1,9 @@ +from __future__ import annotations + from collections import deque from copy import copy -from unittest.mock import patch +from typing import Any +from unittest.mock import Mock, patch import pytest from aiohttp.client_exceptions import ServerConnectionError @@ -18,7 +21,7 @@ from .mockserver import DropResource, MockServer -def test_deprecated_imports(): +def test_deprecated_imports() -> None: from zyte_api import RetryFactory, zyte_api_retrying # noqa: PLC0415 from zyte_api.aio.retry import ( # noqa: PLC0415 RetryFactory as DeprecatedRetryFactory, @@ -47,12 +50,14 @@ class OutlierException(RuntimeError): ), ) @pytest.mark.asyncio -async def test_get_handle_retries(value, exception, mockserver): - kwargs = {} +async def test_get_handle_retries( + value: object, exception: type[Exception], mockserver: MockServer +) -> None: + kwargs: dict[str, Any] = {} if value is not UNSET: kwargs["handle_retries"] = value - def broken_stop(_): + def broken_stop(_: RetryCallState) -> bool: raise OutlierException retrying = AsyncRetrying(stop=broken_stop) @@ -77,14 +82,15 @@ def broken_stop(_): ), ) @pytest.mark.asyncio -async def test_retry_wait(retry_factory, status, waiter, mockserver): - def broken_wait(self, retry_state): +async def test_retry_wait( + retry_factory: type[RetryFactory], status: int, waiter: str, mockserver: MockServer +) -> None: + def broken_wait(self: RetryFactory, retry_state: RetryCallState) -> float: raise OutlierException - class CustomRetryFactory(retry_factory): - pass - - setattr(CustomRetryFactory, f"{waiter}_wait", broken_wait) + CustomRetryFactory = type( + "CustomRetryFactory", (retry_factory,), {f"{waiter}_wait": broken_wait} + ) retrying = CustomRetryFactory().build() client = AsyncZyteAPI( api_key="a", api_url=mockserver.urljoin("/"), retrying=retrying @@ -103,16 +109,15 @@ class CustomRetryFactory(retry_factory): ), ) @pytest.mark.asyncio -async def test_retry_wait_network_error(retry_factory): +async def test_retry_wait_network_error(retry_factory: type[RetryFactory]) -> None: waiter = "network_error" - def broken_wait(self, retry_state): + def broken_wait(self: RetryFactory, retry_state: RetryCallState) -> float: raise OutlierException - class CustomRetryFactory(retry_factory): - pass - - setattr(CustomRetryFactory, f"{waiter}_wait", broken_wait) + CustomRetryFactory = type( + "CustomRetryFactory", (retry_factory,), {f"{waiter}_wait": broken_wait} + ) retrying = CustomRetryFactory().build() with MockServer(resource=DropResource) as mockserver: @@ -403,12 +408,17 @@ def __call__(self, number: float, add: int = 0) -> int: ) @pytest.mark.asyncio @patch("time.monotonic") -async def test_retry_stop(monotonic_mock, retrying, outcomes, exhausted): +async def test_retry_stop( + monotonic_mock: Mock, + retrying: AsyncRetrying, + outcomes: tuple[Any, ...], + exhausted: bool, +) -> None: monotonic_mock.return_value = 0 last_outcome = outcomes[-1] - outcomes = deque(outcomes) + queue = deque(outcomes) - def wait(retry_state): + def wait(retry_state: RetryCallState) -> float: return 0.0 retrying = copy(retrying) @@ -417,7 +427,7 @@ def wait(retry_state): async def run() -> None: while True: try: - outcome = outcomes.popleft() + outcome = queue.popleft() except IndexError: return else: @@ -437,7 +447,7 @@ async def run() -> None: @pytest.mark.asyncio -async def test_deprecated_temporary_download_error(): +async def test_deprecated_temporary_download_error() -> None: class CustomRetryFactory(RetryFactory): def wait(self, retry_state: RetryCallState) -> float: self.temporary_download_error_wait(retry_state=retry_state) @@ -451,7 +461,7 @@ def stop(self, retry_state: RetryCallState) -> bool: outcomes = deque((mock_request_error(status=520), None)) - async def run(): + async def run() -> Any: outcome = outcomes.popleft() if isinstance(outcome, Exception): raise outcome diff --git a/tests/test_sync.py b/tests/test_sync.py index 79c060b..7789e72 100644 --- a/tests/test_sync.py +++ b/tests/test_sync.py @@ -13,19 +13,19 @@ from tests.mockserver import MockServer -def test_api_key(): +def test_api_key() -> None: ZyteAPI(api_key="a") with pytest.raises(NoApiKey): ZyteAPI() -def test_trust_env_is_forwarded(): +def test_trust_env_is_forwarded() -> None: with patch("zyte_api._sync.AsyncZyteAPI") as async_client: ZyteAPI(api_key="a", trust_env=True) assert async_client.call_args.kwargs["trust_env"] is True -def test_get(mockserver): +def test_get(mockserver: MockServer) -> None: client = ZyteAPI(api_key="a", api_url=mockserver.urljoin("/")) expected_result = { "url": "https://a.example", @@ -35,7 +35,7 @@ def test_get(mockserver): assert actual_result == expected_result -def test_iter(mockserver): +def test_iter(mockserver: MockServer) -> None: client = ZyteAPI(api_key="a", api_url=mockserver.urljoin("/")) queries = [ {"url": "https://a.example", "httpResponseBody": True}, @@ -64,7 +64,7 @@ def test_iter(mockserver): assert actual_result in expected_results -def test_semaphore(mockserver): +def test_semaphore(mockserver: MockServer) -> None: client = ZyteAPI(api_key="a", api_url=mockserver.urljoin("/")) client._async_client._semaphore = AsyncMock(wraps=client._async_client._semaphore) queries = [ diff --git a/tests/test_utils.py b/tests/test_utils.py index d0c677f..43dd990 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -1,3 +1,7 @@ +from __future__ import annotations + +from typing import Any + import pytest from aiohttp import TCPConnector @@ -6,7 +10,7 @@ @pytest.mark.asyncio -async def test_create_session_custom_connector(): +async def test_create_session_custom_connector() -> None: # Declare a connector with a random parameter to avoid it matching the # default one. custom_connector = TCPConnector(limit=1850) @@ -16,14 +20,14 @@ async def test_create_session_custom_connector(): @pytest.mark.asyncio -async def test_create_session_trust_env_disabled_by_default(): +async def test_create_session_trust_env_disabled_by_default() -> None: session = create_session() assert session._trust_env is False await session.close() @pytest.mark.asyncio -async def test_create_session_trust_env_can_be_enabled(): +async def test_create_session_trust_env_can_be_enabled() -> None: session = create_session(trust_env=True) assert session._trust_env is True await session.close() @@ -79,7 +83,7 @@ async def test_create_session_trust_env_can_be_enabled(): ), ), ) -def test_guess_intype(file_name, first_line, expected): +def test_guess_intype(file_name: str, first_line: str, expected: str) -> None: assert _guess_intype(file_name, [first_line]) == expected @@ -119,17 +123,17 @@ def test_guess_intype(file_name, first_line, expected): # the URL escaping logic exist upstream. ), ) -def test_process_query(input_, output): +def test_process_query(input_: dict[str, Any], output: dict[str, Any]) -> None: assert _process_query(input_) == output -def test_process_query_bytes(): +def test_process_query_bytes() -> None: with pytest.raises(ValueError, match="Expected a str URL parameter"): _process_query({"url": b"https://example.com"}) @pytest.mark.asyncio # https://github.com/aio-libs/aiohttp/pull/1468 -async def test_deprecated_create_session(): +async def test_deprecated_create_session() -> None: from zyte_api.aio.client import create_session as _create_session # noqa: PLC0415 with pytest.warns( diff --git a/tests/test_x402.py b/tests/test_x402.py index fd3e6ed..bef6b20 100644 --- a/tests/test_x402.py +++ b/tests/test_x402.py @@ -1,6 +1,9 @@ +from __future__ import annotations + import contextlib import importlib.util from os import environ +from typing import TYPE_CHECKING, Any from unittest import mock import pytest @@ -8,7 +11,10 @@ from zyte_api import AsyncZyteAPI from zyte_api._errors import RequestError -from .mockserver import SCREENSHOT +from .mockserver import SCREENSHOT, MockServer + +if TYPE_CHECKING: + from collections.abc import Iterator BODY = "PGh0bWw+PGJvZHk+SGVsbG88aDE+V29ybGQhPC9oMT48L2JvZHk+PC9odG1sPg==" HAS_X402 = importlib.util.find_spec("x402") is not None @@ -16,7 +22,7 @@ KEY = "c85ef7d79691fe79573b1a7064c5232332f53bb1b44a08f1a737f57a68a4706e" -def test_eth_key_param(): +def test_eth_key_param() -> None: if HAS_X402: client = AsyncZyteAPI(eth_key=KEY) assert client.auth.key == KEY @@ -28,7 +34,7 @@ def test_eth_key_param(): @mock.patch.dict(environ, {"ZYTE_API_ETH_KEY": KEY}) -def test_eth_key_env_var(): +def test_eth_key_env_var() -> None: if HAS_X402: client = AsyncZyteAPI() assert client.auth.key == KEY @@ -39,7 +45,7 @@ def test_eth_key_env_var(): AsyncZyteAPI() -def test_eth_key_short(): +def test_eth_key_short() -> None: if HAS_X402: with pytest.raises(ValueError, match="must be exactly 32 bytes long"): AsyncZyteAPI(eth_key="a") @@ -49,7 +55,7 @@ def test_eth_key_short(): @contextlib.contextmanager -def reset_x402_cache(): +def reset_x402_cache() -> Iterator[dict[bytes, Any]]: from zyte_api import _x402 # noqa: PLC0415 try: @@ -351,7 +357,7 @@ def reset_x402_cache(): }, ), ) -async def test_cache(scenario, mockserver): +async def test_cache(scenario: dict[str, Any], mockserver: MockServer) -> None: """Requests that are expected to have the same cost (or cost modifiers) as a preceding request should hit the cache. @@ -378,7 +384,7 @@ async def test_cache(scenario, mockserver): @pytest.mark.skipif(not HAS_X402, reason="x402 not installed") @pytest.mark.asyncio @mock.patch("zyte_api._x402.MINIMIZE_REQUESTS", False) -async def test_no_cache(mockserver): +async def test_no_cache(mockserver: MockServer) -> None: client = AsyncZyteAPI(eth_key=KEY, api_url=mockserver.urljoin("/")) input_ = {"url": "https://a.example", "httpResponseBody": True} output = { @@ -416,7 +422,7 @@ async def test_no_cache(mockserver): @pytest.mark.skipif(not HAS_X402, reason="x402 not installed") @pytest.mark.asyncio -async def test_4xx(mockserver): +async def test_4xx(mockserver: MockServer) -> None: """An unexpected status code lower than 500 raises RequestError immediately.""" client = AsyncZyteAPI(eth_key=KEY, api_url=mockserver.urljoin("/")) @@ -433,7 +439,7 @@ async def test_4xx(mockserver): @pytest.mark.skipif(not HAS_X402, reason="x402 not installed") @pytest.mark.asyncio -async def test_5xx(mockserver): +async def test_5xx(mockserver: MockServer) -> None: """An unexpected status code ≥ 500 gets retried once.""" client = AsyncZyteAPI(eth_key=KEY, api_url=mockserver.urljoin("/")) input_ = {"url": "https://e500.example", "httpResponseBody": True} @@ -449,7 +455,7 @@ async def test_5xx(mockserver): @pytest.mark.skipif(not HAS_X402, reason="x402 not installed") @pytest.mark.asyncio -async def test_payment_retry(mockserver): +async def test_payment_retry(mockserver: MockServer) -> None: client = AsyncZyteAPI(eth_key=KEY, api_url=mockserver.urljoin("/")) input_ = { "url": "https://a.example", @@ -476,7 +482,7 @@ async def test_payment_retry(mockserver): @pytest.mark.skipif(not HAS_X402, reason="x402 not installed") @pytest.mark.asyncio -async def test_payment_retry_exceeded(mockserver): +async def test_payment_retry_exceeded(mockserver: MockServer) -> None: client = AsyncZyteAPI(eth_key=KEY, api_url=mockserver.urljoin("/")) input_ = { "url": "https://a.example", @@ -502,7 +508,7 @@ async def test_payment_retry_exceeded(mockserver): @pytest.mark.asyncio -async def test_no_payment_retry(mockserver): +async def test_no_payment_retry(mockserver: MockServer) -> None: """An HTTP 402 response received out of the context of the x402 protocol, as a response to a regular request using basic auth.""" client = AsyncZyteAPI(api_key="a", api_url=mockserver.urljoin("/")) @@ -530,7 +536,7 @@ async def test_no_payment_retry(mockserver): @pytest.mark.asyncio -async def test_no_payment_retry_exceeded(mockserver): +async def test_no_payment_retry_exceeded(mockserver: MockServer) -> None: client = AsyncZyteAPI(api_key="a", api_url=mockserver.urljoin("/")) input_ = { "url": "https://a.example", @@ -556,7 +562,7 @@ async def test_no_payment_retry_exceeded(mockserver): @pytest.mark.asyncio -async def test_long_error(mockserver): +async def test_long_error(mockserver: MockServer) -> None: client = AsyncZyteAPI(api_key="a", api_url=mockserver.urljoin("/")) input_ = { "url": "https://a.example", diff --git a/zyte_api/_retry.py b/zyte_api/_retry.py index 886f8ce..fde9c03 100644 --- a/zyte_api/_retry.py +++ b/zyte_api/_retry.py @@ -78,8 +78,9 @@ def __init__(self, max_count: int) -> None: def __call__(self, retry_state: RetryCallState) -> bool: if not hasattr(retry_state, "counter"): retry_state.counter = Counter() # type: ignore[attr-defined] - retry_state.counter[self._counter_id] += 1 # type: ignore[attr-defined] - return retry_state.counter[self._counter_id] >= self._max_count # type: ignore[attr-defined] + counter: Counter[Any] = retry_state.counter # type: ignore[attr-defined] + counter[self._counter_id] += 1 + return counter[self._counter_id] >= self._max_count time_unit_type = int | float | timedelta @@ -145,12 +146,13 @@ def __call__(self, retry_state: RetryCallState) -> bool: assert retry_state.outcome, "Unexpected empty outcome" exc = retry_state.outcome.exception() assert exc, "Unexpected empty exception" + counter: Counter[Any] = retry_state.counter # type: ignore[attr-defined] if exc.status == 521: # type: ignore[attr-defined] - retry_state.counter["permanent_download_error"] += 1 # type: ignore[attr-defined] - if retry_state.counter["permanent_download_error"] >= self._max_permanent: # type: ignore[attr-defined] + counter["permanent_download_error"] += 1 + if counter["permanent_download_error"] >= self._max_permanent: return True - retry_state.counter["download_error"] += 1 # type: ignore[attr-defined] - return retry_state.counter["download_error"] >= self._max_total # type: ignore[attr-defined] + counter["download_error"] += 1 + return counter["download_error"] >= self._max_total def _download_error(exc: BaseException) -> bool: @@ -169,8 +171,10 @@ def _402_error(exc: BaseException) -> bool: return isinstance(exc, RequestError) and exc.status == 402 -def _deprecated(message: str, callable_: Callable) -> Callable: - def wrapper(factory: Any, retry_state: RetryCallState) -> Callable: +def _deprecated( + message: str, callable_: Callable[..., Any] +) -> Callable[[Any, RetryCallState], Any]: + def wrapper(factory: Any, retry_state: RetryCallState) -> Any: warn(message, DeprecationWarning, stacklevel=3) return callable_(retry_state=retry_state) diff --git a/zyte_api/_x402.py b/zyte_api/_x402.py index 9fa1748..118bdcc 100644 --- a/zyte_api/_x402.py +++ b/zyte_api/_x402.py @@ -97,7 +97,7 @@ def get_max_cost_hash(query: dict[str, Any]) -> bytes: class X402RetryFactory(RetryFactory): # Disable ban response retries. - download_error_stop = stop_after_attempt(1) # type: ignore[assignment] + download_error_stop = stop_after_attempt(1) X402_RETRYING = X402RetryFactory().build() @@ -168,7 +168,8 @@ async def request() -> dict[str, Any]: self.stats.n_402_req += 1 async with self.semaphore, post_fn(**post_kwargs) as response: if response.status == 402: - return await response.json() + data: dict[str, Any] = await response.json() + return data content = await response.read() response.release() raise RequestError( diff --git a/zyte_api/aio/client.py b/zyte_api/aio/client.py index ca861b8..9c8880f 100644 --- a/zyte_api/aio/client.py +++ b/zyte_api/aio/client.py @@ -1,7 +1,10 @@ from .._async import AsyncZyteAPI # noqa: TID252 -from .._utils import deprecated_create_session as create_session # noqa: F401, TID252 +from .._utils import deprecated_create_session as create_session # noqa: TID252 class AsyncClient(AsyncZyteAPI): request_raw = AsyncZyteAPI.get request_parallel_as_completed = AsyncZyteAPI.iter + + +__all__ = ["AsyncClient", "create_session"] diff --git a/zyte_api/aio/retry.py b/zyte_api/aio/retry.py index 6a7b269..a0d7491 100644 --- a/zyte_api/aio/retry.py +++ b/zyte_api/aio/retry.py @@ -1 +1,3 @@ from .._retry import RetryFactory, zyte_api_retrying # noqa: TID252 + +__all__ = ["RetryFactory", "zyte_api_retrying"] diff --git a/zyte_api/errors.py b/zyte_api/errors.py index ff7d19d..863a5e8 100644 --- a/zyte_api/errors.py +++ b/zyte_api/errors.py @@ -1,7 +1,7 @@ from __future__ import annotations import json -from typing import Optional +from typing import Any, Optional, cast import attr @@ -20,7 +20,7 @@ class ParsedError: #: JSON-decoded response body. #: #: If ``None``, :data:`parse_error` indicates the reason. - data: Optional[dict] + data: Optional[dict[str, Any]] #: If :data:`data` is ``None``, this indicates whether the reason is that #: :data:`response_body` is not valid JSON (``"bad_json"``) or that it is @@ -51,7 +51,7 @@ def type(self) -> Optional[str]: ``"/download/temporary-error"``.""" data = self.data or {} if "type" in data: - return data["type"] + return cast("str", data["type"]) if "error" in data and isinstance(data["error"], str): # HTTP 402 try: prefix, _ = data["error"].split(":", 1)