Skip to content
Open
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
44 changes: 44 additions & 0 deletions stdlib/@tests/test_cases/check_hashlib.py
Original file line number Diff line number Diff line change
@@ -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"))
72 changes: 65 additions & 7 deletions stdlib/_hashlib.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -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]

Expand All @@ -22,17 +51,31 @@ 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
def block_size(self) -> int: ...
@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): ...
Expand Down Expand Up @@ -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: ...
Expand Down Expand Up @@ -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: ...
Expand Down
13 changes: 11 additions & 2 deletions stdlib/hashlib.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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__ = (
Expand Down Expand Up @@ -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]
Expand Down
Loading