From 033a91a0e5bfabbbf3910f5b6c0a53addc6fc9b2 Mon Sep 17 00:00:00 2001 From: subotac <73706465+subotac@users.noreply.github.com> Date: Mon, 5 Oct 2026 23:14:16 +0300 Subject: [PATCH] Fix SHAKE return types for hashlib.new --- stdlib/@tests/test_cases/check_hashlib.py | 44 ++++++++++++++ stdlib/_hashlib.pyi | 72 ++++++++++++++++++++--- stdlib/hashlib.pyi | 13 +++- 3 files changed, 120 insertions(+), 9 deletions(-) create mode 100644 stdlib/@tests/test_cases/check_hashlib.py diff --git a/stdlib/@tests/test_cases/check_hashlib.py b/stdlib/@tests/test_cases/check_hashlib.py new file mode 100644 index 000000000000..32e8451b615e --- /dev/null +++ b/stdlib/@tests/test_cases/check_hashlib.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +import _hashlib +import hashlib +import sys +from _hashlib import HASH, HASHXOF +from io import BytesIO +from typing_extensions import assert_type + + +def check_new(name: str) -> None: + for shake in ( + hashlib.new("shake_128"), + hashlib.new("shake_256"), + hashlib.new("SHAKE-128", b"data", usedforsecurity=False), + hashlib.new("SHAKE-256"), + ): + assert_type(shake, HASHXOF) + assert_type(shake.copy(), HASHXOF) + assert_type(shake.digest(16), bytes) + assert_type(shake.hexdigest(length=16), str) + shake.digest() # type: ignore + shake.hexdigest() # type: ignore + + fixed = hashlib.new("sha256") + assert_type(fixed, HASH) + assert_type(fixed.digest(), bytes) + fixed.digest(16) # type: ignore + + dynamic = hashlib.new(name) + assert_type(dynamic.digest(), bytes) + assert_type(dynamic.digest(16), bytes) + assert_type(dynamic.copy().hexdigest(length=16), str) + dynamic.digest("16") # type: ignore + + assert_type(_hashlib.new("shake_128", string=b"data").digest(16), bytes) + assert_type(_hashlib.new("SHAKE-256").hexdigest(16), str) + assert_type(_hashlib.new("MD5"), HASH) + assert_type(_hashlib.new(name).hexdigest(16), str) + if sys.version_info >= (3, 13): + assert_type(_hashlib.new("shake_256", data=b"data").copy(), HASHXOF) + + if sys.version_info >= (3, 11): + hashlib.file_digest(BytesIO(b"data"), lambda: hashlib.new("sha256")) diff --git a/stdlib/_hashlib.pyi b/stdlib/_hashlib.pyi index b98edc5757c3..5b8050ebc8d4 100644 --- a/stdlib/_hashlib.pyi +++ b/stdlib/_hashlib.pyi @@ -2,10 +2,39 @@ import sys from _typeshed import ReadableBuffer from collections.abc import Callable from types import ModuleType -from typing import AnyStr, Protocol, TypeAlias, final, overload, type_check_only +from typing import AnyStr, Literal, Protocol, TypeAlias, final, overload, type_check_only from typing_extensions import Self, disjoint_base _DigestMod: TypeAlias = str | Callable[[], _HashObject] | ModuleType | None +_FixedDigestName: TypeAlias = Literal[ + "md5", + "MD5", + "sha1", + "SHA1", + "sha224", + "SHA224", + "sha256", + "SHA256", + "sha384", + "SHA384", + "sha512", + "SHA512", + "sha3_224", + "sha3-224", + "SHA3-224", + "sha3_256", + "sha3-256", + "SHA3-256", + "sha3_384", + "sha3-384", + "SHA3-384", + "sha3_512", + "sha3-512", + "SHA3-512", +] +_XofDigestName: TypeAlias = Literal[ + "shake_128", "shake128", "SHAKE128", "shake-128", "SHAKE-128", "shake_256", "shake256", "SHAKE256", "shake-256", "SHAKE-256" +] openssl_md_meth_names: frozenset[str] @@ -22,8 +51,22 @@ class _HashObject(Protocol): def hexdigest(self) -> str: ... def update(self, obj: ReadableBuffer, /) -> None: ... +@type_check_only +class _HashObjectWithOptionalLength(Protocol): + @property + def digest_size(self) -> int: ... + @property + def block_size(self) -> int: ... + @property + def name(self) -> str: ... + def copy(self) -> Self: ... + # Depending on the algorithm, length is either required or not accepted. + def digest(self, length: int = ...) -> bytes: ... + def hexdigest(self, length: int = ...) -> str: ... + def update(self, obj: ReadableBuffer, /) -> None: ... + @disjoint_base -class HASH: +class HASH(_HashObjectWithOptionalLength): @property def digest_size(self) -> int: ... @property @@ -31,8 +74,8 @@ class HASH: @property def name(self) -> str: ... def copy(self) -> Self: ... - def digest(self) -> bytes: ... - def hexdigest(self) -> str: ... + def digest(self) -> bytes: ... # type: ignore[override] + def hexdigest(self) -> str: ... # type: ignore[override] def update(self, obj: ReadableBuffer, /) -> None: ... class UnsupportedDigestmodError(ValueError): ... @@ -63,9 +106,19 @@ def get_fips_mode() -> int: ... def hmac_new(key: ReadableBuffer, msg: ReadableBuffer = b"", digestmod: _DigestMod = None) -> HMAC: ... if sys.version_info >= (3, 13): + @overload def new( - name: str, data: ReadableBuffer = b"", *, usedforsecurity: bool = True, string: ReadableBuffer | None = None + name: _XofDigestName, data: ReadableBuffer = b"", *, usedforsecurity: bool = True, string: ReadableBuffer | None = None + ) -> HASHXOF: ... + @overload + def new( + name: _FixedDigestName, data: ReadableBuffer = b"", *, usedforsecurity: bool = True, string: ReadableBuffer | None = None ) -> HASH: ... + @overload + def new( + name: str, data: ReadableBuffer = b"", *, usedforsecurity: bool = True, string: ReadableBuffer | None = None + ) -> _HashObjectWithOptionalLength: ... + def openssl_md5( data: ReadableBuffer = b"", *, usedforsecurity: bool = True, string: ReadableBuffer | None = None ) -> HASH: ... @@ -102,9 +155,14 @@ if sys.version_info >= (3, 13): def openssl_shake_256( data: ReadableBuffer = b"", *, usedforsecurity: bool = True, string: ReadableBuffer | None = None ) -> HASHXOF: ... - else: - def new(name: str, string: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASH: ... + @overload + def new(name: _XofDigestName, string: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASHXOF: ... + @overload + def new(name: _FixedDigestName, string: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASH: ... + @overload + def new(name: str, string: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> _HashObjectWithOptionalLength: ... + def openssl_md5(string: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASH: ... def openssl_sha1(string: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASH: ... def openssl_sha224(string: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASH: ... diff --git a/stdlib/hashlib.pyi b/stdlib/hashlib.pyi index 50bc8e21f1d5..8cabacd1a380 100644 --- a/stdlib/hashlib.pyi +++ b/stdlib/hashlib.pyi @@ -2,7 +2,11 @@ import sys from _blake2 import blake2b as blake2b, blake2s as blake2s from _hashlib import ( HASH, + HASHXOF, + _FixedDigestName, _HashObject, + _HashObjectWithOptionalLength, + _XofDigestName, openssl_md5 as md5, openssl_sha1 as sha1, openssl_sha3_224 as sha3_224, @@ -20,7 +24,7 @@ from _hashlib import ( ) from _typeshed import ReadableBuffer from collections.abc import Callable, Set as AbstractSet -from typing import Protocol, type_check_only +from typing import Protocol, overload, type_check_only if sys.version_info >= (3, 15): __all__ = ( @@ -89,7 +93,12 @@ else: "pbkdf2_hmac", ) -def new(name: str, data: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASH: ... +@overload +def new(name: _XofDigestName, data: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASHXOF: ... +@overload +def new(name: _FixedDigestName, data: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> HASH: ... +@overload +def new(name: str, data: ReadableBuffer = b"", *, usedforsecurity: bool = True) -> _HashObjectWithOptionalLength: ... algorithms_guaranteed: AbstractSet[str] algorithms_available: AbstractSet[str]