diff --git a/src/pyrobusta/bindings/http_connection.py b/src/pyrobusta/bindings/http_connection.py index 4b7ad6b..826fb3c 100644 --- a/src/pyrobusta/bindings/http_connection.py +++ b/src/pyrobusta/bindings/http_connection.py @@ -82,8 +82,11 @@ async def _run_state_machine(self): # [2] process request by state machine while True: - self._engine.run(self._recv_buf) - if self._prev_state == self._engine.state: + promise = self._engine.run(self._recv_buf) + if promise: + while not promise.done: + await sleep_ms(self.STATE_MACHINE_SLEEP_MS) + elif self._prev_state == self._engine.state: # No state transition occurred, read more data break self._prev_state = self._engine.state diff --git a/src/pyrobusta/protocol/http.py b/src/pyrobusta/protocol/http.py index 9321d59..415479c 100644 --- a/src/pyrobusta/protocol/http.py +++ b/src/pyrobusta/protocol/http.py @@ -591,9 +591,9 @@ def run(self, rx): This method returns on every state transition. """ if self.is_terminated(): - return + return None try: - self.state(rx) + return self.state(rx) except BufferOverflowError: self.abort(500) self.set_response_body(b"Buffer full") @@ -613,6 +613,8 @@ def run(self, rx): self.abort(500) self.set_response_body(b"Internal Server Error") + return None + # ======================================== # Helpers for routing, state machine logic # ======================================== diff --git a/src/pyrobusta/protocol/http_basic_auth.py b/src/pyrobusta/protocol/http_basic_auth.py index a5b0eb0..2acc616 100644 --- a/src/pyrobusta/protocol/http_basic_auth.py +++ b/src/pyrobusta/protocol/http_basic_auth.py @@ -7,6 +7,7 @@ # pylint: disable=W0212,R0401 +import asyncio import binascii import os @@ -16,11 +17,8 @@ create_csrf_cookie, verify_csrf_cookie, ) -from pyrobusta.utils.patch import add_method -from pyrobusta.utils.crypto import ( - constant_time_equal, - pbkdf2_sha256, -) +from pyrobusta.utils.patch import add_method, patch_extra_property +from pyrobusta.utils.crypto import constant_time_equal, pbkdf2_sha256, a_pbkdf2_sha256 from pyrobusta.utils.iam import ( NO_POLICY, IAMDatabase, @@ -34,56 +32,42 @@ _DUMMY_ITER = 5000 _DUMMY_SALT = os.urandom(16) -_DUMMY_HASH = pbkdf2_sha256(os.urandom(20), _DUMMY_SALT, _DUMMY_ITER) +_DUMMY_HASH = pbkdf2_sha256(os.urandom(20), _DUMMY_SALT, 100) _BROWSER_SECURITY = None -_SESSIONS = None +_SESSIONS_ENABLED = None _SESSION_TTL_SEC = None +_AUTH_PROVIDER = None -def _auth_user(self, auth_provider: IAMDatabase, sessions=False): - # Session validation - if sessions and (session_cookie := self.get_cookie("session")): - if credentials := verify_session_cookie( - session_cookie.encode("ascii"), auth_provider - ): - username, user_info = credentials - is_session = True - return username, user_info, is_session +class AuthPromise: + # pylint: disable=R0903 + """ + Helper class for the asynchronous computation of PBKDF2 password hashing. + """ - is_session = False + __slots__ = ("done", "result") - # Protocol validation - auth_header = self.headers.get("authorization") - if not auth_header or auth_header[:6].lower() != "basic ": - return None + def __init__(self, authenticator, username, password): + self.done = None + self.result = None - # Decoding - auth_header = auth_header[6:].strip() - try: - auth_header = binascii.a2b_base64(auth_header).decode() - except binascii.Error: - return None + asyncio.create_task(authenticator(self, username, password)) - # Authentication - user_sep = auth_header.find(":") - if user_sep < 0: - return None - username = auth_header[:user_sep].lower() - password = auth_header[user_sep + 1 :].strip().encode("ascii") - user_info = auth_provider.get_user_info(username) +async def _auth_user(promise, username, password): + user_info = _AUTH_PROVIDER.get_user_info(username) stored_hash = user_info[PASS_HASH] if user_info else _DUMMY_HASH if user_info: - password_hash = pbkdf2_sha256( + password_hash = await a_pbkdf2_sha256( password, user_info[PASS_SALT], user_info[PASS_ITER], len(user_info[PASS_HASH]), ) else: - password_hash = pbkdf2_sha256( + password_hash = await a_pbkdf2_sha256( password, _DUMMY_SALT, _DUMMY_ITER, len(_DUMMY_HASH) ) @@ -92,9 +76,9 @@ def _auth_user(self, auth_provider: IAMDatabase, sessions=False): if not (user_ok and hash_ok): logging.info("authentication failed for user=[%s]", username) - return None - - return username, user_info, is_session + else: + promise.result = (username, user_info) + promise.done = True def _handle_auth_st(self, _): @@ -103,7 +87,7 @@ def _handle_auth_st(self, _): method = self.method.decode("ascii") url = self.url.decode("ascii") - policy = self.get_policy(url) + policy = _AUTH_PROVIDER.get_access_policies(url) if not policy: self.state = self._handle_auth_header_st return @@ -120,16 +104,69 @@ def _handle_auth_st(self, _): self.state = self._handle_auth_header_st +def parse_auth_headers(http_ctx): + """ + Parse authorization headers and return + a username and password. + """ + # Protocol validation + auth_header = http_ctx.headers.get("authorization") + if not auth_header or auth_header[:6].lower() != "basic ": + return None + + # Decoding + auth_header = auth_header[6:].strip() + try: + auth_header = binascii.a2b_base64(auth_header).decode("ascii") + except binascii.Error: + return None + + # Authentication + user_sep = auth_header.find(":") + if user_sep < 0: + return None + + username = auth_header[:user_sep].lower() + password = auth_header[user_sep + 1 :].strip().encode("ascii") + return username, password + + def _handle_auth_header_st(self, _): method = self.method.decode("ascii") url = self.url.decode("ascii") + is_session = False + user_data = None # Authentication - if not (credentials := self._authenticate()): + is_session = ( + _SESSIONS_ENABLED + and (session_cookie := self.get_cookie("session")) + and ( + user_data := verify_session_cookie( + session_cookie.encode("ascii"), _AUTH_PROVIDER + ) + ) + ) + + if not is_session: + if not self.auth_promise: + credentials = parse_auth_headers(self) + if not credentials: + self.set_response_header(b"WWW-Authenticate", b'Basic realm="Device"') + self.terminate(401) + return None + username, password = credentials + self.auth_promise = AuthPromise(_auth_user, username, password) + if not self.auth_promise.done: + return self.auth_promise + user_data = self.auth_promise.result + + if not user_data: self.set_response_header(b"WWW-Authenticate", b'Basic realm="Device"') self.terminate(401) - return - username, user_info, is_session = credentials + return None + + username, user_info = user_data # CSRF validation, cookie setting if _BROWSER_SECURITY and not is_session: @@ -144,21 +181,21 @@ def _handle_auth_header_st(self, _): user_info[USER_SECRET], ): self.terminate(403) - return + return None elif self.method in (self.GET, self.HEAD): if self.get_cookie("csrf-token") is None: cookie = create_csrf_cookie(user_info[USER_SECRET], self.TLS) self.set_response_header(b"set-cookie", cookie, override=False) # Session creation - if not is_session and _SESSIONS: + if not is_session and _SESSIONS_ENABLED: session_cookie = create_session_cookie( username, user_info[USER_SECRET], _SESSION_TTL_SEC, self.TLS ) self.set_response_header(b"set-cookie", session_cookie, override=False) # Authorization - policy = self.get_policy(url) + policy = _AUTH_PROVIDER.get_access_policies(url) if not policy: allowed_roles = 0 @@ -169,9 +206,10 @@ def _handle_auth_header_st(self, _): if (allowed_roles & user_info[ROLE_MASK]) == 0: self.terminate(403) - return + return None self.state = self._route_request_st + return None def apply_patches(cls, config, auth_provider: IAMDatabase): @@ -194,21 +232,13 @@ def apply_patches(cls, config, auth_provider: IAMDatabase): "authenticated clients are vulnerable to CSRF attacks" ) - def get_policy(route: str): - return auth_provider.get_access_policies(route) - - allow_sessions = config.http_sessions - - def _authenticate(self): - return _auth_user(self, auth_provider, allow_sessions) - - add_method(cls, _handle_auth_st) - add_method(cls, _handle_auth_header_st) - add_method(cls, get_policy, "static") - add_method(cls, _authenticate) - # pylint: disable=W0603 - global _BROWSER_SECURITY, _SESSIONS, _SESSION_TTL_SEC + global _AUTH_PROVIDER, _BROWSER_SECURITY, _SESSIONS_ENABLED, _SESSION_TTL_SEC + _AUTH_PROVIDER = auth_provider _BROWSER_SECURITY = config.http_browser_security - _SESSIONS = config.http_sessions + _SESSIONS_ENABLED = config.http_sessions _SESSION_TTL_SEC = config.http_session_ttl_sec + + add_method(cls, _handle_auth_st) + add_method(cls, _handle_auth_header_st) + patch_extra_property(cls, "auth_promise") diff --git a/src/pyrobusta/utils/crypto.py b/src/pyrobusta/utils/crypto.py index e0b904b..b2b6c2c 100644 --- a/src/pyrobusta/utils/crypto.py +++ b/src/pyrobusta/utils/crypto.py @@ -2,6 +2,7 @@ Utility functions for cryptography operations. """ +import asyncio import binascii import hashlib import math @@ -76,37 +77,6 @@ def verify_signed_token(secret: bytes, token: bytes, payload_size: int): return constant_time_equal(request_signature, expected_signature) -def pbkdf2_sha256(password: bytes, salt: bytes, iterations: int, dklen: int = 32): - """ - Compute PBKDF2-SHA256-based password hash, - based on RFC8018. - """ - hmac = HmacSha256(password) - output = bytearray() - block_number = 1 - if iterations <= 0: - raise ValueError("iterations must be positive") - - if dklen <= 0: - raise ValueError("dklen must be positive") - - if dklen > (2**dklen - 1) * dklen: - raise ValueError("derived key too long") - - while len(output) < dklen: - # U1 = PRF(password, salt || INT(block)) - u = hmac.digest(salt + block_number.to_bytes(4, "big")) - t = bytearray(u) - # U2 ... Uc - for _ in range(iterations - 1): - u = hmac.digest(u) - for i, _ in enumerate(t): - t[i] ^= u[i] - output.extend(t) - block_number += 1 - return bytes(output[:dklen]) - - def validate_password(password: str, min_length: int = 16, min_entropy: float = 80.0): """ Validate user password complexity. @@ -135,3 +105,69 @@ def validate_password(password: str, min_length: int = 16, min_entropy: float = f"Password entropy too low ({entropy:.1f} bits, " f"minimum {min_entropy:.1f} bits)" ) + + +def pbkdf2_validate_arguments(iterations, dklen): + """ + Validate PBKDF2 arguments. + """ + if iterations <= 0: + raise ValueError("iterations must be positive") + + if dklen <= 0: + raise ValueError("dklen must be positive") + + if dklen > (2**dklen - 1) * dklen: + raise ValueError("derived key too long") + + +def pbkdf2_sha256(password: bytes, salt: bytes, iterations: int, dklen: int = 32): + """ + Compute PBKDF2-SHA256-based password hash, + based on RFC8018. + """ + hmac = HmacSha256(password) + output = bytearray() + block_number = 1 + pbkdf2_validate_arguments(iterations, dklen) + + while len(output) < dklen: + # U1 = PRF(password, salt || INT(block)) + u = hmac.digest(salt + block_number.to_bytes(4, "big")) + t = bytearray(u) + # U2 ... Uc + for _ in range(iterations - 1): + u = hmac.digest(u) + for i, _ in enumerate(t): + t[i] ^= u[i] + output.extend(t) + block_number += 1 + return bytes(output[:dklen]) + + +async def a_pbkdf2_sha256( + password: bytes, salt: bytes, iterations: int, dklen: int = 32 +): + """ + Compute PBKDF2-SHA256-based password hash asynchronously, + based on RFC8018. + """ + hmac = HmacSha256(password) + output = bytearray() + block_number = 1 + pbkdf2_validate_arguments(iterations, dklen) + + while len(output) < dklen: + # U1 = PRF(password, salt || INT(block)) + u = hmac.digest(salt + block_number.to_bytes(4, "big")) + t = bytearray(u) + # U2 ... Uc + for i in range(iterations - 1): + u = hmac.digest(u) + for j, _ in enumerate(t): + t[j] ^= u[j] + if i % 100 == 0: + await asyncio.sleep(0.001) # pylint: disable=E1101 + output.extend(t) + block_number += 1 + return bytes(output[:dklen]) diff --git a/tests/unit/http_base.py b/tests/unit/http_base.py index ca4a07a..f098c7f 100644 --- a/tests/unit/http_base.py +++ b/tests/unit/http_base.py @@ -1,9 +1,11 @@ +import asyncio import os import unittest import sys from pathlib import Path from unittest.mock import patch, mock_open + from tests.unit.utils import load_module, patch_time patch_time() @@ -25,6 +27,18 @@ def setUpClass(cls): cls.base_config = {} cls.cwd = os.getcwd() + @staticmethod + def run_coroutine(coro): + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = None + + if loop is not None: + raise RuntimeError("Unexpected running event loop in synchronous test") + + asyncio.run(coro) + def setUp(self): # Patch current working directory and config import pyrobusta @@ -99,6 +113,14 @@ def setUp(self): self.http_module.HttpEngine, config, self.iam_db ) + self.coroutine_patcher = patch( + "pyrobusta.protocol.http_basic_auth.asyncio.create_task" + ) + create_task = self.coroutine_patcher.start() + create_task.side_effect = self.run_coroutine + + self.addCleanup(self.coroutine_patcher.stop) + self.engine = self.http_module.HttpEngine() # --------------------