diff --git a/.gitignore b/.gitignore index 8982aab..f8223b0 100644 --- a/.gitignore +++ b/.gitignore @@ -7,4 +7,6 @@ __pycache__ *.pyo *.egg-info docs/_build/ -*.lock \ No newline at end of file +*.lock +.venv +build diff --git a/src/srptools/client.py b/src/srptools/client.py index 891dfe0..1466848 100644 --- a/src/srptools/client.py +++ b/src/srptools/client.py @@ -1,5 +1,6 @@ from __future__ import annotations +from secrets import compare_digest from typing import TYPE_CHECKING from .common import SRPSessionBase @@ -12,7 +13,7 @@ class SRPClientSession(SRPSessionBase): role = 'client' - def __init__(self, srp_context: SRPContext, *, private: str = ''): + def __init__(self, srp_context: SRPContext, *, private: str | int | bytes = ''): super().__init__(srp_context, private) self._password_hash: int | None = None @@ -22,7 +23,7 @@ def __init__(self, srp_context: SRPContext, *, private: str = ''): self._client_public = srp_context.get_client_public(client_private=self._this_private) - def init_base(self, salt: str): + def init_base(self, salt: str | bytes): super().init_base(salt) self._password_hash = self._context.get_common_password_hash(self._salt) @@ -39,7 +40,10 @@ def init_session_key(self): self._key = self._context.get_common_session_key(premaster_secret) - def verify_proof(self, key_proof: str, *, base64: bool = False) -> bool: + def verify_proof(self, key_proof: str | bytes, *, base64: bool = False) -> bool: super().verify_proof(key_proof) + if isinstance(key_proof, bytes): + return compare_digest(key_proof, self._key_proof_hash) + return self._value_decode(key_proof, base64=base64) == self.key_proof_hash diff --git a/src/srptools/common.py b/src/srptools/common.py index aec8bdf..6172915 100644 --- a/src/srptools/common.py +++ b/src/srptools/common.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING from .exceptions import SRPException -from .utils import b64_from, hex_from, hex_from_b64, int_from_hex, value_encode +from .utils import b64_from, hex_from, hex_from_b64, int_from_bytes, int_from_hex, value_encode if TYPE_CHECKING: from .context import Salt, SRPContext @@ -15,7 +15,7 @@ class SRPSessionBase: role: str | None = None - def __init__(self, srp_context: SRPContext, private: str = '') -> None: + def __init__(self, srp_context: SRPContext, private: str | int | bytes = '') -> None: self._context = srp_context self._salt: Salt | None = None @@ -30,14 +30,19 @@ def __init__(self, srp_context: SRPContext, private: str = '') -> None: self._this_private: int | None = None if private: - self._this_private = int_from_hex(private) + if isinstance(private, int): + self._this_private = private + elif isinstance(private, bytes): + self._this_private = int_from_bytes(private) + else: + self._this_private = int_from_hex(private) @property def _this_public(self) -> int: return getattr(self, f'_{self.role}_public') def _other_public(self, val: int) -> None: - other = ('server' if self.role == 'client' else 'client') + other = 'server' if self.role == 'client' else 'client' setattr(self, f'_{other}_public', val) _other_public = property(None, _other_public) @@ -50,6 +55,10 @@ def private(self) -> str: def private_b64(self) -> str: return b64_from(self._this_private) + @property + def private_bin(self) -> bytes: + return self._context.pad(self._this_private) + @property def public(self) -> str: return hex_from(self._this_public) @@ -58,6 +67,10 @@ def public(self) -> str: def public_b64(self) -> str: return b64_from(self._this_public) + @property + def public_bin(self) -> bytes: + return self._context.pad(self._this_public) + @property def key(self) -> str: return hex_from(self._key) @@ -66,6 +79,10 @@ def key(self) -> str: def key_b64(self) -> str: return b64_from(self._key) + @property + def key_bin(self) -> bytes: + return self._key + @property def key_proof(self) -> str: return hex_from(self._key_proof) @@ -74,6 +91,10 @@ def key_proof(self) -> str: def key_proof_b64(self) -> str: return b64_from(self._key_proof) + @property + def key_proof_bin(self) -> bytes: + return self._key_proof + @property def key_proof_hash(self) -> str: return hex_from(self._key_proof_hash) @@ -82,18 +103,31 @@ def key_proof_hash(self) -> str: def key_proof_hash_b64(self) -> str: return b64_from(self._key_proof_hash) + @property + def key_proof_hash_bin(self) -> bytes: + return self._key_proof_hash + @classmethod - def _value_decode(cls, value: str, *, base64: bool = False) -> str: + def _value_decode(cls, value: str | bytes, *, base64: bool = False) -> str | bytes: """Decodes value into hex optionally from base64.""" - return hex_from_b64(value) if base64 else value + if base64: + if isinstance(value, bytes): + raise SRPException('Cannot decode base64 from bytes.') + return hex_from_b64(value) + return value def process( self, + other_public: str | bytes = '', + salt: str | bytes = '', *, - other_public: str, - salt: str, base64: bool = False, ) -> tuple[str, str, str]: + if base64 and (isinstance(other_public, bytes) or isinstance(salt, bytes)): + raise SRPException( + 'Cannot decode base64 from bytes. ' + 'If the value is bytes, it is already decoded and should not be treated as base64.' + ) salt = self._value_decode(salt, base64=base64) other_public = self._value_decode(other_public, base64=base64) @@ -108,18 +142,30 @@ def process( return key, key_proof, key_proof_hash - def init_base(self, salt: str) -> None: - salt = unhexlify(salt) - self._salt = salt + def init_base(self, salt: str | bytes) -> None: + if isinstance(salt, bytes): + self._salt = salt + else: + self._salt = unhexlify(salt) def init_session_key(self) -> None: pass - def verify_proof(self, key_prove: str, *, base64: bool = False) -> bool: + def verify_proof(self, key_prove: str | bytes, *, base64: bool = False) -> bool: pass - def init_common_secret(self, other_public: str) -> None: - other_public = int_from_hex(other_public) + def init_common_secret(self, other_public: str | int | bytes) -> None: + if isinstance(other_public, int): + pass + elif isinstance(other_public, bytes): + other_public = int_from_bytes(other_public) + else: + try: + other_public = int(other_public, 16) + except (ValueError, TypeError) as e: + raise SRPException( + f'Wrong public provided for {self.__class__.__name__}: cannot decode value: {e}', + ) from e if other_public % self._context._prime == 0: # A % N is zero | B % N is zero raise SRPException(f'Wrong public provided for {self.__class__.__name__}.') @@ -130,16 +176,11 @@ def init_common_secret(self, other_public: str) -> None: def init_session_key_proof(self) -> None: proof = self._context.get_common_session_key_proof( - session_key=self._key, - salt=self._salt, - server_public=self._server_public, - client_public=self._client_public + session_key=self._key, salt=self._salt, server_public=self._server_public, client_public=self._client_public ) self._key_proof = proof self._key_proof_hash = self._context.get_common_session_key_proof_hash( - session_key=self._key, - session_key_proof=proof, - client_public=self._client_public + session_key=self._key, session_key_proof=proof, client_public=self._client_public ) diff --git a/src/srptools/context.py b/src/srptools/context.py index a93fecb..9c764ab 100644 --- a/src/srptools/context.py +++ b/src/srptools/context.py @@ -201,11 +201,30 @@ def get_common_session_key_proof_hash( """H(A | M | K)""" return self.hash(client_public, session_key_proof, session_key, as_bytes=True) - def get_user_data_triplet(self, *, base64: bool = False) -> tuple[str, str, str]: - """( <_user>, <_password verifier>, )""" + def get_user_data_triplet( + self, + *, + base64: bool = False, + binary: bool = False, + ) -> tuple[str, str | bytes, str | bytes]: + """( <_user>, <_password verifier>, ) + + :param bool base64: Output verifier and salt as base64 strings. + :param bool binary: Output verifier and salt as raw bytes. Verifier + is prime-width (padded), salt is bits_salt-width (padded). + Mutually exclusive with ``base64``. + :raises SRPException: if both ``base64`` and ``binary`` are True. + """ + if binary and base64: + raise SRPException('binary and base64 are mutually exclusive') + salt = self.generate_salt() verifier = self.get_common_password_verifier(self.get_common_password_hash(salt)) + if binary: + salt_bytes = int_to_bytes(salt).rjust(self._bits_salt // 8, b'\x00') + return self._user, self.pad(verifier), salt_bytes + verifier = value_encode(verifier, base64=base64) salt = value_encode(salt, base64=base64) diff --git a/src/srptools/server.py b/src/srptools/server.py index 1d0ac8e..ac8a523 100644 --- a/src/srptools/server.py +++ b/src/srptools/server.py @@ -1,29 +1,39 @@ from __future__ import annotations +from secrets import compare_digest from typing import TYPE_CHECKING from .common import SRPSessionBase -from .utils import int_from_hex +from .utils import int_from_bytes, int_from_hex if TYPE_CHECKING: from .context import SRPContext class SRPServerSession(SRPSessionBase): - role = 'server' - def __init__(self, srp_context: SRPContext, *, password_verifier: str, private: str = ''): + def __init__( + self, + srp_context: SRPContext, + *, + password_verifier: str | int | bytes, + private: str | int | bytes = '', + ): super().__init__(srp_context, private) - self._password_verifier = int_from_hex(password_verifier) + if isinstance(password_verifier, int): + self._password_verifier = password_verifier + elif isinstance(password_verifier, bytes): + self._password_verifier = int_from_bytes(password_verifier) + else: + self._password_verifier = int_from_hex(password_verifier) if not private: self._this_private = srp_context.generate_server_private() self._server_public = srp_context.get_server_public( - password_verifier=self._password_verifier, - server_private=self._this_private + password_verifier=self._password_verifier, server_private=self._this_private ) def init_session_key(self) -> None: @@ -33,12 +43,15 @@ def init_session_key(self) -> None: password_verifier=self._password_verifier, server_private=self._this_private, client_public=self._client_public, - common_secret=self._common_secret + common_secret=self._common_secret, ) self._key = self._context.get_common_session_key(premaster_secret) - def verify_proof(self, key_proof: str, *, base64: bool = False) -> bool: + def verify_proof(self, key_proof: str | bytes, *, base64: bool = False) -> bool: super().verify_proof(key_proof) + if isinstance(key_proof, bytes): + return compare_digest(key_proof, self._key_proof) + return self._value_decode(key_proof, base64=base64) == self.key_proof diff --git a/src/srptools/utils.py b/src/srptools/utils.py index 2e5e29d..1b84bc9 100644 --- a/src/srptools/utils.py +++ b/src/srptools/utils.py @@ -6,11 +6,18 @@ def value_encode(val: int | bytes, *, base64: bool = False) -> str: return b64_from(val) if base64 else hex_from(val) -def hex_from_b64(val: str) -> str: +def hex_from_b64(val: str | bytes) -> str: """Returns hex string representation for a base64 encoded value.""" + if isinstance(val, bytes): + val = val.decode('ascii') return b64decode(val).hex() +def int_from_bytes(val: bytes) -> int: + """Returns int representation for a given bytes value (big-endian).""" + return int.from_bytes(val, 'big') + + def hex_from(val: int | bytes) -> str: """Returns hex string representation for a given value.""" if isinstance(val, int): diff --git a/tests/test_binary.py b/tests/test_binary.py new file mode 100644 index 0000000..cc2c05f --- /dev/null +++ b/tests/test_binary.py @@ -0,0 +1,177 @@ +from binascii import unhexlify + +import pytest + +from srptools import SRPClientSession, SRPContext, SRPException, SRPServerSession +from srptools.utils import int_from_hex + + +def test_full_handshake_binary(): + """Full SRP handshake with bytes input at every IO boundary.""" + context = SRPContext('alice', 'password123') + username, password_verifier, salt = context.get_user_data_triplet() + prime, gen = context.prime, context.generator + + # Convert hex outputs to bytes for binary path. + verifier_bin = unhexlify(password_verifier) + salt_bin = unhexlify(salt) + + # Server accepts bytes verifier. + server_session = SRPServerSession(SRPContext(username, prime=prime, generator=gen), password_verifier=verifier_bin) + server_public_bin = server_session.public_bin + + # Client processes bytes public + bytes salt. + client_session = SRPClientSession(SRPContext(username, 'password123', prime=prime, generator=gen)) + client_session.process(server_public_bin, salt_bin) + client_public_bin = client_session.public_bin + + # Server processes bytes public + bytes salt. + server_session.process(client_public_bin, salt_bin) + + # Session keys and proofs match. + assert client_session.key_bin == server_session.key_bin + assert client_session.key_proof_bin == server_session.key_proof_bin + assert client_session.key_proof_hash_bin == server_session.key_proof_hash_bin + + +def test_verify_proof_binary(): + """verify_proof accepts bytes proofs on both sides.""" + context = SRPContext('alice', 'password123') + username, password_verifier, salt = context.get_user_data_triplet() + prime, gen = context.prime, context.generator + + salt_bin = unhexlify(salt) + verifier_bin = unhexlify(password_verifier) + + server_session = SRPServerSession(SRPContext(username, prime=prime, generator=gen), password_verifier=verifier_bin) + client_session = SRPClientSession(SRPContext(username, 'password123', prime=prime, generator=gen)) + + client_session.process(server_session.public_bin, salt_bin) + server_session.process(client_session.public_bin, salt_bin) + + # Server verifies client's M (bytes). + assert server_session.verify_proof(client_session.key_proof_bin) + # Client verifies server's H(A|M|K) (bytes). + assert client_session.verify_proof(server_session.key_proof_hash_bin) + + +def test_session_restore_via_bytes_private(): + """Session restored with private= reproduces the same public.""" + context = SRPContext('alice', 'password123') + username, password_verifier, _ = context.get_user_data_triplet() + prime, gen = context.prime, context.generator + + verifier_bin = unhexlify(password_verifier) + + original_server = SRPServerSession( + SRPContext(username, prime=prime, generator=gen), password_verifier=password_verifier + ) + server_private_bin = original_server.private_bin + + restored_server = SRPServerSession( + SRPContext(username, prime=prime, generator=gen), password_verifier=verifier_bin, private=server_private_bin + ) + assert restored_server.public_bin == original_server.public_bin + + original_client = SRPClientSession(SRPContext(username, 'password123', prime=prime, generator=gen)) + client_private_bin = original_client.private_bin + + restored_client = SRPClientSession( + SRPContext(username, 'password123', prime=prime, generator=gen), private=client_private_bin + ) + assert restored_client.public_bin == original_client.public_bin + + +def test_server_accepts_int_verifier(): + """SRPServerSession accepts int password_verifier directly.""" + context = SRPContext('alice', 'password123') + username, password_verifier, _ = context.get_user_data_triplet() + prime, gen = context.prime, context.generator + + verifier_int = int_from_hex(password_verifier) + + server_session = SRPServerSession(SRPContext(username, prime=prime, generator=gen), password_verifier=verifier_int) + assert server_session.public_bin + + +def test_get_user_data_triplet_binary(): + """get_user_data_triplet(binary=True) emits bytes verifier and salt.""" + context = SRPContext('alice', 'password123', bits_salt=64) + + username, verifier, salt = context.get_user_data_triplet(binary=True) + + assert username == 'alice' + assert isinstance(verifier, bytes) + assert isinstance(salt, bytes) + # Salt is bits_salt-width (64 bits = 8 bytes). + assert len(salt) == 8 + # Verifier is prime-width (1024-bit prime = 128 bytes). + assert len(verifier) == 128 + + +def test_get_user_data_triplet_binary_mutual_exclusion(): + """binary=True and base64=True together must raise.""" + context = SRPContext('alice', 'password123') + with pytest.raises(SRPException): + context.get_user_data_triplet(base64=True, binary=True) + + +def test_process_bytes_base64_mutual_exclusion(): + """process() must reject bytes input with base64=True.""" + context = SRPContext('alice', 'password123') + username, password_verifier, salt = context.get_user_data_triplet() + prime, gen = context.prime, context.generator + + salt_bin = unhexlify(salt) + server_session = SRPServerSession( + SRPContext(username, prime=prime, generator=gen), password_verifier=password_verifier + ) + client_session = SRPClientSession(SRPContext(username, 'password123', prime=prime, generator=gen)) + + with pytest.raises(SRPException): + client_session.process(server_session.public_bin, salt_bin, base64=True) + + +def test_init_common_secret_rejects_garbage_str(): + """init_common_secret raises SRPException on non-hex str.""" + context = SRPContext('alice', 'password123') + server_session = SRPServerSession(context, password_verifier='1') + with pytest.raises(SRPException): + server_session.init_common_secret('not-hex-at-all') + + +def test_binary_path_matches_hex_path(): + """Binary handshake produces the same session key as hex handshake, + given the same private values and salt. + """ + context = SRPContext('alice', 'password123') + username, password_verifier, salt = context.get_user_data_triplet() + prime, gen = context.prime, context.generator + + salt_bin = unhexlify(salt) + verifier_bin = unhexlify(password_verifier) + + # Hex path: capture privates to reuse in binary path. + server_hex = SRPServerSession(SRPContext(username, prime=prime, generator=gen), password_verifier=password_verifier) + server_private_bin = server_hex.private_bin + + client_hex = SRPClientSession(SRPContext(username, 'password123', prime=prime, generator=gen)) + client_private_bin = client_hex.private_bin + + client_hex.process(server_hex.public, salt) + server_hex.process(client_hex.public, salt) + hex_key = client_hex.key_bin + + # Binary path: restore sessions with the same privates, feed bytes. + server_bin = SRPServerSession( + SRPContext(username, prime=prime, generator=gen), password_verifier=verifier_bin, private=server_private_bin + ) + client_bin = SRPClientSession( + SRPContext(username, 'password123', prime=prime, generator=gen), private=client_private_bin + ) + + client_bin.process(server_bin.public_bin, salt_bin) + server_bin.process(client_bin.public_bin, salt_bin) + bin_key = client_bin.key_bin + + assert hex_key == bin_key