diff --git a/lambda/pyproject.toml b/lambda/pyproject.toml index 068c8ccf..0fce5283 100644 --- a/lambda/pyproject.toml +++ b/lambda/pyproject.toml @@ -3,17 +3,17 @@ name = "ffsync-lambda-handler" version = "0.1.0" requires-python = ">=3.14" dependencies = [ - "aws-lambda-powertools==3.34.0", - "boto3==1.43.89", + "aws-lambda-powertools==3.35.0", + "boto3==1.43.96", "cryptography==50.0.1", "pydantic==2.13.5", "mohawk==1.1.0", - "PyJWT==2.13.0", + "PyJWT==2.14.0", "requests==2.34.2" ] -[project.optional-dependencies] +[dependency-groups] dev = [ "pytest==9.1.1", "pytest-cov==7.1.0", @@ -23,10 +23,10 @@ dev = [ "flake8==7.3.0", "Flake8-pyproject==1.2.4", "mypy==2.3.1", - "types-boto3==1.43.89", + "types-boto3[apigatewaymanagementapi,dynamodb,kms]==1.43.96", "hypothesis>=6.0.0", - "types-requests==2.33.0.20260712", - "datamodel-code-generator==0.76.2", + "types-requests==v2.33.0.20260906", + "datamodel-code-generator==0.82.0", ] [tool.pytest.ini_options] @@ -118,9 +118,12 @@ select = [ [tool.mypy] check_untyped_defs = true -show_error_codes = true +disallow_untyped_defs = true pretty = true -ignore_missing_imports = true +# Deliberately NOT ignore_missing_imports: it turns an uninstalled stub package into a silent +# Any, so an annotation like `table: "Table"` type-checks while checking nothing. Scope the +# fallback per-module below instead, so a missing stub surfaces as an error. +warn_unused_ignores = true # keeps the remaining `type: ignore`s from rotting files = ["src/", "tests/"] exclude = [ "__pycache__", @@ -129,3 +132,12 @@ exclude = [ "htmlcov", ] plugins = ["pydantic.mypy"] + +# mohawk ships no stubs and has no types-* package on PyPI. +[[tool.mypy.overrides]] +module = ["mohawk", "mohawk.*"] +ignore_missing_imports = true + +[tool.pydantic-mypy] +init_typed = true # Use field names rather than aliases for constructors + diff --git a/lambda/src/environment/service_provider.py b/lambda/src/environment/service_provider.py index 526237d0..65f075e1 100644 --- a/lambda/src/environment/service_provider.py +++ b/lambda/src/environment/service_provider.py @@ -2,11 +2,13 @@ import json import os from functools import cached_property +from typing import TYPE_CHECKING, Any, Callable, Optional import boto3 from aws_lambda_powertools.event_handler import CORSConfig, Response from aws_lambda_powertools.logging import Logger from aws_lambda_powertools.metrics import Metrics +from aws_lambda_powertools.utilities.typing import LambdaContext from src.middlewares.hawk_auth import HawkAuthenticationError, HawkAuthMiddleware, UidMismatchError from src.middlewares.request_logging import RequestLoggingMiddleware @@ -59,6 +61,9 @@ from src.services.token_generator import TokenGenerator from src.services.user_manager import UserManager +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import DynamoDBServiceResource, Table + @functools.lru_cache(maxsize=1) def create_service_provider() -> "ServiceProvider": # pragma: nocover @@ -70,7 +75,7 @@ def create_service_provider() -> "ServiceProvider": # pragma: nocover return ServiceProvider() -def lambda_entrypoint(fn): +def lambda_entrypoint(fn: Callable[..., Any]) -> Callable[..., Any]: """Decorator that injects a cached ServiceProvider when none is provided. In production, creates/reuses a cached ServiceProvider via lru_cache. @@ -78,7 +83,9 @@ def lambda_entrypoint(fn): """ @functools.wraps(fn) - def wrapper(event, context, service_provider=None): + def wrapper( + event: dict, context: LambdaContext, service_provider: Optional["ServiceProvider"] = None + ) -> Any: if service_provider is None: # pragma: nocover service_provider = create_service_provider() try: @@ -102,24 +109,24 @@ def user_agent(self) -> str: return "layertwo-ffsync/1.0" @cached_property - def aws_region(self): # pragma: nocover + def aws_region(self) -> Optional[str]: # pragma: nocover return os.environ.get("AWS_REGION") @cached_property - def session(self): # pragma: nocover + def session(self) -> boto3.Session: # pragma: nocover return boto3.Session(region_name=self.aws_region) @cached_property - def table_name(self): - return os.environ.get("STORAGE_TABLE_NAME") + def table_name(self) -> str: + return os.environ["STORAGE_TABLE_NAME"] @cached_property - def dynamodb_resource(self): # pragma: nocover + def dynamodb_resource(self) -> "DynamoDBServiceResource": # pragma: nocover """Shared DynamoDB resource — reuses a single connection pool.""" return self.session.resource("dynamodb") @cached_property - def dynamodb_table(self): + def dynamodb_table(self) -> "Table": """Create DynamoDB Table resource""" return self.dynamodb_resource.Table(self.table_name) @@ -128,11 +135,11 @@ def storage_manager(self) -> StorageManager: return StorageManager(table=self.dynamodb_table) @cached_property - def token_users_table_name(self): - return os.environ.get("TOKEN_USERS_TABLE_NAME") + def token_users_table_name(self) -> str: + return os.environ["TOKEN_USERS_TABLE_NAME"] @cached_property - def token_users_table(self): + def token_users_table(self) -> "Table": """Create DynamoDB Table resource for token users""" return self.dynamodb_resource.Table(self.token_users_table_name) @@ -142,14 +149,14 @@ def user_manager(self) -> UserManager: @cached_property def _storage_exception_handlers(self) -> dict: - def handle_hawk_auth(ex): + def handle_hawk_auth(ex: Exception) -> Response: return Response( status_code=401, content_type="application/json", body='{"error": "Unauthorized"}', ) - def handle_uid_mismatch(ex): + def handle_uid_mismatch(ex: Exception) -> Response: return Response( status_code=403, content_type="application/json", @@ -163,7 +170,7 @@ def handle_uid_mismatch(ex): @cached_property def _auth_exception_handlers(self) -> dict: - def handle_hawk_auth(ex): + def handle_hawk_auth(ex: Exception) -> Response: return Response( status_code=401, content_type="application/json", @@ -175,7 +182,7 @@ def handle_hawk_auth(ex): } @cached_property - def storage_api_router(self): + def storage_api_router(self) -> ApiRouter: return ApiRouter( routes=[ DeleteAllRootRoute(self.storage_manager), @@ -214,7 +221,7 @@ def oidc_client_id(self) -> str: return os.environ["OIDC_CLIENT_ID"] @cached_property - def base_domain(self): + def base_domain(self) -> Optional[str]: return os.environ.get("BASE_DOMAIN") @cached_property @@ -270,11 +277,11 @@ def token_generator(self) -> TokenGenerator: # Auth API properties @cached_property - def auth_table_name(self): - return os.environ.get("AUTH_TABLE_NAME") + def auth_table_name(self) -> str: + return os.environ["AUTH_TABLE_NAME"] @cached_property - def auth_table(self): + def auth_table(self) -> "Table": """DynamoDB Table for auth accounts, sessions, and OAuth codes""" return self.dynamodb_resource.Table(self.auth_table_name) @@ -283,7 +290,7 @@ def auth_signing_key_id(self) -> str: return os.environ["AUTH_SIGNING_KEY_ID"] @cached_property - def kms_client(self): # pragma: nocover + def kms_client(self) -> Any: # pragma: nocover return self.session.client("kms") @cached_property @@ -327,7 +334,7 @@ def cors_config(self) -> CORSConfig: ) @cached_property - def auth_api_router(self): + def auth_api_router(self) -> ApiRouter: """Create API router for Auth API with all FxA-compatible routes""" return ApiRouter( routes=[ @@ -404,7 +411,7 @@ def auth_api_router(self): ) @cached_property - def token_api_router(self): + def token_api_router(self) -> ApiRouter: """Create API router for Token API (sync token issuance)""" return ApiRouter( routes=[ @@ -422,7 +429,7 @@ def token_api_router(self): ) @cached_property - def profile_api_router(self): + def profile_api_router(self) -> ApiRouter: """Create API router for Profile API (OAuth Bearer auth)""" return ApiRouter( routes=[ @@ -440,15 +447,15 @@ def profile_api_router(self): # HAWK Authorizer properties @cached_property - def token_cache_table_name(self): - return os.environ.get("TOKEN_CACHE_TABLE_NAME") + def token_cache_table_name(self) -> str: + return os.environ["TOKEN_CACHE_TABLE_NAME"] @cached_property def token_duration(self) -> int: return int(os.environ["TOKEN_DURATION"]) @cached_property - def token_cache_table(self): + def token_cache_table(self) -> "Table": """Create DynamoDB Table resource for token cache""" return self.dynamodb_resource.Table(self.token_cache_table_name) @@ -464,11 +471,11 @@ def hawk_service(self) -> HawkService: # Channel Service properties @cached_property - def channel_table_name(self): - return os.environ.get("CHANNEL_TABLE_NAME") + def channel_table_name(self) -> str: + return os.environ["CHANNEL_TABLE_NAME"] @cached_property - def channel_table(self): + def channel_table(self) -> "Table": """DynamoDB Table for pairing channel state""" resource = self.session.resource("dynamodb") return resource.Table(self.channel_table_name) diff --git a/lambda/src/middlewares/hawk_auth.py b/lambda/src/middlewares/hawk_auth.py index 335c4802..5ba40237 100644 --- a/lambda/src/middlewares/hawk_auth.py +++ b/lambda/src/middlewares/hawk_auth.py @@ -11,6 +11,7 @@ from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler, NextMiddleware from aws_lambda_powertools.metrics import Metrics, MetricUnit +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.fxa_token_manager import FxATokenManager from src.services.hawk_service import HawkService @@ -65,7 +66,15 @@ def handler(self, app: APIGatewayRestResolver, next_middleware: NextMiddleware) self._metrics.add_metric("HawkAuthSuccess", MetricUnit.Count, 1) return next_middleware(app) - def _validate_storage_hawk(self, event, auth_header, method, path, host, port): + def _validate_storage_hawk( + self, + event: APIGatewayProxyEvent, + auth_header: str, + method: str, + path: str, + host: str, + port: int, + ) -> None: """Validate storage Hawk token and check URL uid matches authenticated user.""" assert self._hawk_service is not None try: @@ -86,7 +95,15 @@ def _validate_storage_hawk(self, event, auth_header, method, path, host, port): event["requestContext"]["hawk_uid"] = creds.user_id - def _validate_session_hawk(self, event, auth_header, method, path, host, port): + def _validate_session_hawk( + self, + event: APIGatewayProxyEvent, + auth_header: str, + method: str, + path: str, + host: str, + port: int, + ) -> None: """Validate FxA session Hawk token.""" assert self._token_manager is not None uid = self._token_manager.verify_session_hawk(auth_header, method, path, host, port) diff --git a/lambda/src/middlewares/request_logging.py b/lambda/src/middlewares/request_logging.py index 1956534a..65750e40 100644 --- a/lambda/src/middlewares/request_logging.py +++ b/lambda/src/middlewares/request_logging.py @@ -19,7 +19,7 @@ def handler(self, app: APIGatewayRestResolver, next_middleware: NextMiddleware) method = event.get("httpMethod", "UNKNOWN") path = event.get("path", "UNKNOWN") - user_id = event.get("requestContext", {}).get("hawk_uid", "anonymous") # type: ignore + user_id = (event.get("requestContext") or {}).get("hawk_uid", "anonymous") logger.info( "Request received", diff --git a/lambda/src/routes/auth/account_attached_clients.py b/lambda/src/routes/auth/account_attached_clients.py index 9f8f2415..8bc4b7ba 100644 --- a/lambda/src/routes/auth/account_attached_clients.py +++ b/lambda/src/routes/auth/account_attached_clients.py @@ -1,9 +1,10 @@ """AccountAttachedClients route — GET /v1/account/attached_clients""" -from typing import Sequence +from typing import Any, Sequence from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.device_manager import DeviceManager from src.shared.base_route import BaseRoute @@ -21,12 +22,12 @@ def __init__( self._device_manager = device_manager self.middlewares = middlewares - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/v1/account/attached_clients", middlewares=list(self.middlewares)) - def handle_account_attached_clients(): + def handle_account_attached_clients() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: uid = event["requestContext"]["hawk_uid"] session_token_id = event["requestContext"].get("hawk_token_id", "") diff --git a/lambda/src/routes/auth/account_create.py b/lambda/src/routes/auth/account_create.py index 89117c10..c3cd0cd2 100644 --- a/lambda/src/routes/auth/account_create.py +++ b/lambda/src/routes/auth/account_create.py @@ -3,8 +3,10 @@ import json import re import uuid +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.auth_account_manager import AuthAccountManager from src.services.fxa_crypto import derive_verify_hash, generate_random_bytes @@ -30,14 +32,14 @@ def __init__( self._token_manager = token_manager self._oidc_validator = oidc_validator - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/account/create") - def handle_account_create(): + def handle_account_create() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: # Validate OIDC Bearer token - headers = event.headers or {} + headers = event.headers auth_header = headers.get("authorization") if not auth_header: return self._error(401, 110, "Missing Authorization header") diff --git a/lambda/src/routes/auth/account_device.py b/lambda/src/routes/auth/account_device.py index 6928bd27..82c13c60 100644 --- a/lambda/src/routes/auth/account_device.py +++ b/lambda/src/routes/auth/account_device.py @@ -1,10 +1,11 @@ """AccountDevice route — POST /v1/account/device""" import json -from typing import Sequence +from typing import Any, Sequence from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.device_manager import DeviceManager from src.shared.base_route import BaseRoute @@ -22,12 +23,12 @@ def __init__( self._device_manager = device_manager self.middlewares = middlewares - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/account/device", middlewares=list(self.middlewares)) - def handle_account_device(): + def handle_account_device() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: uid = event["requestContext"]["hawk_uid"] session_token_id = event["requestContext"].get("hawk_token_id", "") body = json.loads(event.body or "{}") diff --git a/lambda/src/routes/auth/account_devices.py b/lambda/src/routes/auth/account_devices.py index c69671d3..53ae2055 100644 --- a/lambda/src/routes/auth/account_devices.py +++ b/lambda/src/routes/auth/account_devices.py @@ -1,9 +1,10 @@ """AccountDevices route — GET /v1/account/devices""" -from typing import Sequence +from typing import Any, Sequence from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.device_manager import DeviceManager from src.shared.base_route import BaseRoute @@ -21,12 +22,12 @@ def __init__( self._device_manager = device_manager self.middlewares = middlewares - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/v1/account/devices", middlewares=list(self.middlewares)) - def handle_account_devices(): + def handle_account_devices() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: uid = event["requestContext"]["hawk_uid"] session_token_id = event["requestContext"].get("hawk_token_id", "") diff --git a/lambda/src/routes/auth/account_devices_notify.py b/lambda/src/routes/auth/account_devices_notify.py index 7a37abd4..3376ef5c 100644 --- a/lambda/src/routes/auth/account_devices_notify.py +++ b/lambda/src/routes/auth/account_devices_notify.py @@ -1,10 +1,11 @@ """AccountDevicesNotify route — POST /v1/account/devices/notify (no-op)""" import json -from typing import Sequence +from typing import Any, Sequence from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.shared.base_route import BaseRoute @@ -15,12 +16,12 @@ class AccountDevicesNotifyRoute(BaseRoute): def __init__(self, middlewares: Sequence[BaseMiddlewareHandler] = ()): self.middlewares = middlewares - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/account/devices/notify", middlewares=list(self.middlewares)) - def handle_account_devices_notify(): + def handle_account_devices_notify() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: return Response( status_code=200, content_type="application/json", diff --git a/lambda/src/routes/auth/account_keys.py b/lambda/src/routes/auth/account_keys.py index 214a615e..aefa2926 100644 --- a/lambda/src/routes/auth/account_keys.py +++ b/lambda/src/routes/auth/account_keys.py @@ -1,8 +1,10 @@ """AccountKeys route — GET /v1/account/keys""" import json +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.auth_account_manager import AuthAccountManager from src.services.fxa_crypto import derive_key_request_key, encrypt_key_bundle @@ -23,14 +25,14 @@ def __init__( self._account_manager = account_manager self._token_manager = token_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/v1/account/keys") - def handle_account_keys(): + def handle_account_keys() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: # Authenticate via key-fetch token with Hawk HMAC verification - headers = event.headers or {} + headers = event.headers auth_header = headers.get("authorization", "") if not auth_header: return self._error(401, 110, "Missing or invalid authorization") diff --git a/lambda/src/routes/auth/account_login.py b/lambda/src/routes/auth/account_login.py index d654dc64..8601e0b5 100644 --- a/lambda/src/routes/auth/account_login.py +++ b/lambda/src/routes/auth/account_login.py @@ -2,8 +2,10 @@ import json import re +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.auth_account_manager import AuthAccountManager from src.services.fxa_crypto import constant_time_compare, derive_verify_hash @@ -25,12 +27,12 @@ def __init__( self._account_manager = account_manager self._token_manager = token_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/account/login") - def handle_account_login(): + def handle_account_login() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: # Parse body body_str = event.body if not body_str: diff --git a/lambda/src/routes/auth/account_status.py b/lambda/src/routes/auth/account_status.py index aaa0585f..01672d66 100644 --- a/lambda/src/routes/auth/account_status.py +++ b/lambda/src/routes/auth/account_status.py @@ -1,8 +1,10 @@ """AccountStatus route — GET /v1/account/status""" import json +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.auth_account_manager import AuthAccountManager from src.shared.base_route import BaseRoute @@ -15,12 +17,12 @@ class AccountStatusRoute(BaseRoute): def __init__(self, account_manager: AuthAccountManager): self._account_manager = account_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/v1/account/status") - def handle_account_status(): + def handle_account_status() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: params = event.query_string_parameters or {} email = params.get("email") if not email: diff --git a/lambda/src/routes/auth/jwks.py b/lambda/src/routes/auth/jwks.py index 916900f0..9780dab2 100644 --- a/lambda/src/routes/auth/jwks.py +++ b/lambda/src/routes/auth/jwks.py @@ -1,8 +1,10 @@ """JWKS route — GET /v1/jwks""" import json +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.jwt_service import JWTService from src.shared.base_route import BaseRoute @@ -14,12 +16,12 @@ class JWKSRoute(BaseRoute): def __init__(self, jwt_service: JWTService): self._jwt_service = jwt_service - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/v1/jwks") - def handle_jwks(): + def handle_jwks() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: jwk = self._jwt_service.get_public_key_jwk() return Response( diff --git a/lambda/src/routes/auth/oauth_authorization.py b/lambda/src/routes/auth/oauth_authorization.py index eea4921c..26b61fd6 100644 --- a/lambda/src/routes/auth/oauth_authorization.py +++ b/lambda/src/routes/auth/oauth_authorization.py @@ -1,10 +1,11 @@ """OAuthAuthorization route — POST /v1/oauth/authorization""" import json -from typing import Sequence +from typing import Any, Sequence from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.oauth_code_manager import OAuthCodeManager from src.shared.base_route import BaseRoute @@ -29,12 +30,12 @@ def __init__( self._oauth_code_manager = oauth_code_manager self.middlewares = middlewares - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/oauth/authorization", middlewares=list(self.middlewares)) - def handle_oauth_authorization(): + def handle_oauth_authorization() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: uid = event["requestContext"]["hawk_uid"] # Parse body diff --git a/lambda/src/routes/auth/oauth_destroy.py b/lambda/src/routes/auth/oauth_destroy.py index 5a32fbe5..a49346e2 100644 --- a/lambda/src/routes/auth/oauth_destroy.py +++ b/lambda/src/routes/auth/oauth_destroy.py @@ -2,8 +2,10 @@ import hashlib import json +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.oauth_code_manager import OAuthCodeManager from src.shared.base_route import BaseRoute @@ -15,12 +17,12 @@ class OAuthDestroyRoute(BaseRoute): def __init__(self, oauth_code_manager: OAuthCodeManager): self._oauth_code_manager = oauth_code_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/oauth/destroy") - def handle_oauth_destroy(): + def handle_oauth_destroy() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: body_str = event.body if not body_str: return Response( diff --git a/lambda/src/routes/auth/oauth_token.py b/lambda/src/routes/auth/oauth_token.py index 2be82ac3..36a3c426 100644 --- a/lambda/src/routes/auth/oauth_token.py +++ b/lambda/src/routes/auth/oauth_token.py @@ -3,9 +3,11 @@ import hashlib import json import time +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.metrics import Metrics, MetricUnit +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.auth_account_manager import AuthAccountManager from src.services.fxa_token_manager import FxATokenManager @@ -36,12 +38,12 @@ def __init__( self._metrics = metrics self._token_manager = token_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/oauth/token") - def handle_oauth_token(): + def handle_oauth_token() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: body_str = event.body if not body_str: return self._error(400, 107, "Missing request body") @@ -191,12 +193,12 @@ def _handle_refresh_token(self, body: dict) -> Response: body=result.model_dump_json(exclude_none=True), ) - def _handle_fxa_credentials(self, event, body: dict) -> Response: + def _handle_fxa_credentials(self, event: APIGatewayProxyEvent, body: dict) -> Response: """Issue an access token using Hawk-authenticated session credentials.""" if self._token_manager is None: return self._error(400, 107, "fxa-credentials grant not supported") - headers = event.headers or {} + headers = event.headers auth_header = headers.get("authorization", "") if not auth_header: return self._error(401, 110, "Missing or invalid authorization") diff --git a/lambda/src/routes/auth/oidc_discovery.py b/lambda/src/routes/auth/oidc_discovery.py index c145d5fa..8971044f 100644 --- a/lambda/src/routes/auth/oidc_discovery.py +++ b/lambda/src/routes/auth/oidc_discovery.py @@ -1,8 +1,10 @@ """OIDCDiscovery route — GET /.well-known/openid-configuration""" import json +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.jwt_service import JWTService from src.shared.base_route import BaseRoute @@ -14,12 +16,12 @@ class OIDCDiscoveryRoute(BaseRoute): def __init__(self, jwt_service: JWTService): self._jwt_service = jwt_service - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/.well-known/openid-configuration") - def handle_oidc_discovery(): + def handle_oidc_discovery() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: issuer = self._jwt_service.issuer return Response( diff --git a/lambda/src/routes/auth/oidc_exchange.py b/lambda/src/routes/auth/oidc_exchange.py index 53267e5b..941b553a 100644 --- a/lambda/src/routes/auth/oidc_exchange.py +++ b/lambda/src/routes/auth/oidc_exchange.py @@ -2,11 +2,13 @@ import json import time +from typing import Any import requests from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.metrics import Metrics, MetricUnit +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.auth_account_manager import AuthAccountManager from src.services.oidc_validator import OIDCValidator @@ -22,12 +24,12 @@ class OIDCProviderConfigRoute(BaseRoute): def __init__(self, oidc_validator: OIDCValidator): self._oidc_validator = oidc_validator - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/v1/oidc/config") - def handle_oidc_provider_config(): + def handle_oidc_provider_config() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: try: config = self._oidc_validator.discover_provider_config() except Exception: @@ -64,12 +66,12 @@ def __init__( def _default_headers(self) -> dict[str, str]: return {"User-Agent": self._user_agent} - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/oidc/exchange") - def handle_oidc_code_exchange(): + def handle_oidc_code_exchange() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: # Parse request body try: body = json.loads(event.body or "{}") diff --git a/lambda/src/routes/auth/scoped_key_data.py b/lambda/src/routes/auth/scoped_key_data.py index 54896c61..b0127abe 100644 --- a/lambda/src/routes/auth/scoped_key_data.py +++ b/lambda/src/routes/auth/scoped_key_data.py @@ -1,10 +1,11 @@ """ScopedKeyData route — POST /v1/account/scoped-key-data""" import json -from typing import Sequence +from typing import Any, Sequence from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from pydantic import ValidationError as PydanticValidationError from src.services.auth_account_manager import AuthAccountManager @@ -23,12 +24,12 @@ def __init__( self._account_manager = account_manager self.middlewares = middlewares - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/account/scoped-key-data", middlewares=list(self.middlewares)) - def handle_scoped_key_data(): + def handle_scoped_key_data() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: uid = event["requestContext"]["hawk_uid"] # Parse and validate body diff --git a/lambda/src/routes/auth/session_destroy.py b/lambda/src/routes/auth/session_destroy.py index 57fceed2..c9cfaef0 100644 --- a/lambda/src/routes/auth/session_destroy.py +++ b/lambda/src/routes/auth/session_destroy.py @@ -2,10 +2,11 @@ import json import re -from typing import Sequence +from typing import Any, Sequence from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.fxa_token_manager import FxATokenManager from src.shared.base_route import BaseRoute @@ -24,14 +25,14 @@ def __init__( self._token_manager = token_manager self.middlewares = middlewares - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/v1/session/destroy", middlewares=list(self.middlewares)) - def handle_session_destroy(): + def handle_session_destroy() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: # Extract token id from Hawk header for deletion - headers = event.headers or {} + headers = event.headers auth_header = headers.get("authorization", "") match = HAWK_ID_PATTERN.search(auth_header) if match: # pragma: no branch diff --git a/lambda/src/routes/auth/session_status.py b/lambda/src/routes/auth/session_status.py index 9d1a64db..724cdcc8 100644 --- a/lambda/src/routes/auth/session_status.py +++ b/lambda/src/routes/auth/session_status.py @@ -1,9 +1,10 @@ """SessionStatus route — GET /v1/session/status""" -from typing import Sequence +from typing import Any, Sequence from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.shared.base_route import BaseRoute from src.shared.models import SessionStatusOutput @@ -15,12 +16,12 @@ class SessionStatusRoute(BaseRoute): def __init__(self, middlewares: Sequence[BaseMiddlewareHandler] = ()): self.middlewares = middlewares - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/v1/session/status", middlewares=list(self.middlewares)) - def handle_session_status(): + def handle_session_status() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: uid = event["requestContext"]["hawk_uid"] result = SessionStatusOutput(state="verified", uid=uid) diff --git a/lambda/src/routes/bso/delete.py b/lambda/src/routes/bso/delete.py index fa7c71e9..4ae5cc77 100644 --- a/lambda/src/routes/bso/delete.py +++ b/lambda/src/routes/bso/delete.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -24,22 +26,17 @@ class DeleteBSORoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.delete("/1.5//storage//") - def handle_request(uid: str, collectionName: str, objectId: str): + def handle_request(uid: str, collectionName: str, objectId: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Delete a specific storage object""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() path_params = event.path_parameters or {} collection_name = path_params["collectionName"] diff --git a/lambda/src/routes/bso/read.py b/lambda/src/routes/bso/read.py index 97dfa1b5..32fed484 100644 --- a/lambda/src/routes/bso/read.py +++ b/lambda/src/routes/bso/read.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -23,22 +25,17 @@ class ReadBSORoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.5//storage//") - def handle_request(uid: str, collectionName: str, objectId: str): + def handle_request(uid: str, collectionName: str, objectId: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Get a specific storage object""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() path_params = event.path_parameters or {} collection_name = path_params["collectionName"] @@ -50,7 +47,7 @@ def handle(self, event) -> Response: raise ValidationException(str(e)) # Handle conditional GET headers (Requirements 6.1-6.4) - headers = event.get("headers", {}) + headers = event.headers if_modified_since_header = headers.get("x-if-modified-since") if_unmodified_since_header = headers.get("x-if-unmodified-since") diff --git a/lambda/src/routes/bso/update.py b/lambda/src/routes/bso/update.py index 6e77b099..c5faa46c 100644 --- a/lambda/src/routes/bso/update.py +++ b/lambda/src/routes/bso/update.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from pydantic import ValidationError as PydanticValidationError from src.services.storage_manager import StorageManager @@ -28,25 +30,20 @@ class UpdateBSORoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.put("/1.5//storage//") - def handle_request(uid: str, collectionName: str, objectId: str): + def handle_request(uid: str, collectionName: str, objectId: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Update a storage object""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() path_params = event.path_parameters or {} - body = event.body + body = event.body or "" collection_name = path_params["collectionName"] object_id = path_params["objectId"] try: diff --git a/lambda/src/routes/collections/create.py b/lambda/src/routes/collections/create.py index 8aefd12f..16263090 100644 --- a/lambda/src/routes/collections/create.py +++ b/lambda/src/routes/collections/create.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -26,25 +28,20 @@ class CreateCollectionRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.post("/1.5//storage/") - def handle_request(uid: str, collectionName: str): + def handle_request(uid: str, collectionName: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Create a new collection or batch create/update objects""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() path_params = event.path_parameters or {} - headers = event.headers or {} + headers = event.headers body = event.body collection_name = path_params["collectionName"] @@ -182,7 +179,9 @@ def handle(self, event) -> Response: body=json.dumps({"error": "Internal server error"}), ) - def _check_precondition(self, user_id, collection_name, if_unmodified_since): + def _check_precondition( + self, user_id: str, collection_name: str, if_unmodified_since: str | float + ) -> bool: """Check if collection was modified since given timestamp""" try: timestamp = float(if_unmodified_since) diff --git a/lambda/src/routes/collections/delete.py b/lambda/src/routes/collections/delete.py index ad84a58f..54527c13 100644 --- a/lambda/src/routes/collections/delete.py +++ b/lambda/src/routes/collections/delete.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -15,22 +17,17 @@ class DeleteCollectionRoute(BaseRoute): def __init__(self, dynamodb_service: StorageManager): self.dynamodb_service = dynamodb_service - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.delete("/1.5//storage/") - def handle_request(uid: str, collectionName: str): + def handle_request(uid: str, collectionName: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Delete an entire collection or specific BSOs""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() path_params = event.path_parameters or {} query_params = event.query_string_parameters or {} diff --git a/lambda/src/routes/collections/list.py b/lambda/src/routes/collections/list.py index 6b017b01..4949dbec 100644 --- a/lambda/src/routes/collections/list.py +++ b/lambda/src/routes/collections/list.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -14,22 +16,17 @@ class ListCollectionsRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.5//storage") - def handle_request(uid: str): + def handle_request(uid: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """List all collections with their metadata""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() # Get collections using storage manager collections = self.storage_manager.list_collections(user_id) diff --git a/lambda/src/routes/collections/read.py b/lambda/src/routes/collections/read.py index cea21eab..b9fba37d 100644 --- a/lambda/src/routes/collections/read.py +++ b/lambda/src/routes/collections/read.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -19,26 +21,21 @@ class ReadCollectionRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.5//storage/") - def handle_request(uid: str, collectionName: str): + def handle_request(uid: str, collectionName: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Get collection metadata or retrieve objects with filtering""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() path_params = event.path_parameters or {} query_params = event.query_string_parameters or {} - headers = event.headers or {} + headers = event.headers collection_name = path_params["collectionName"] # Validate collection name before any storage call @@ -151,7 +148,7 @@ def handle(self, event) -> Response: body=json.dumps({"error": "Internal server error"}), ) - def _parse_timestamp(self, value): + def _parse_timestamp(self, value: str | None) -> float | None: """Parse timestamp from string""" if value is None: return None @@ -160,7 +157,7 @@ def _parse_timestamp(self, value): except ValueError, TypeError: # pragma: nocover return None - def _parse_int(self, value, default): + def _parse_int(self, value: str | None, default: int) -> int: """Parse integer with default""" if value is None: return default @@ -169,7 +166,7 @@ def _parse_int(self, value, default): except ValueError, TypeError: # pragma: nocover return default - def _parse_bool(self, value): + def _parse_bool(self, value: str | None) -> bool: """Parse boolean from string""" if value is None: return True # pragma: nocover diff --git a/lambda/src/routes/collections/update.py b/lambda/src/routes/collections/update.py index f5456979..a340c09c 100644 --- a/lambda/src/routes/collections/update.py +++ b/lambda/src/routes/collections/update.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -24,25 +26,20 @@ class UpdateCollectionRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.put("/1.5//storage/") - def handle_request(uid: str, collectionName: str): + def handle_request(uid: str, collectionName: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Update collection with batch objects""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() path_params = event.path_parameters or {} - body = event.body + body = event.body or "" collection_name = path_params["collectionName"] try: validate_collection_name(collection_name) diff --git a/lambda/src/routes/info/read_collections.py b/lambda/src/routes/info/read_collections.py index 7d6e7da8..bddae96e 100644 --- a/lambda/src/routes/info/read_collections.py +++ b/lambda/src/routes/info/read_collections.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -13,12 +15,12 @@ class ReadCollectionsInfoRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.5//info/collections") - def handle_request(uid: str): + def handle_request(uid: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """ Get metadata for all collections. @@ -26,14 +28,9 @@ def handle(self, event) -> Response: Example: {"bookmarks": 1234567890.12, "tabs": 1234567880.00} """ try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() # Get collections using storage manager collections = self.storage_manager.list_collections(user_id) diff --git a/lambda/src/routes/info/read_configuration.py b/lambda/src/routes/info/read_configuration.py index f10581a0..27bef1ec 100644 --- a/lambda/src/routes/info/read_configuration.py +++ b/lambda/src/routes/info/read_configuration.py @@ -1,5 +1,8 @@ +from typing import Any + from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.shared.base_route import BaseRoute from src.shared.models import ConfigurationOutput @@ -33,12 +36,12 @@ def __init__( self.max_total_records = max_total_records self.max_total_bytes = max_total_bytes - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.5//info/configuration") - def handle_request(uid: str): + def handle_request(uid: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """ Get server configuration limits. diff --git a/lambda/src/routes/info/read_counts.py b/lambda/src/routes/info/read_counts.py index b0f3fa1d..0f0d2704 100644 --- a/lambda/src/routes/info/read_counts.py +++ b/lambda/src/routes/info/read_counts.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -13,12 +15,12 @@ class ReadCollectionCountsRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.5//info/collection_counts") - def handle_request(uid: str): + def handle_request(uid: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """ Get count information for all collections. @@ -26,14 +28,9 @@ def handle(self, event) -> Response: Example: {"bookmarks": 15, "tabs": 7} """ try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() # Get collections using storage manager collections = self.storage_manager.list_collections(user_id) diff --git a/lambda/src/routes/info/read_quota.py b/lambda/src/routes/info/read_quota.py index b8192245..528b02e8 100644 --- a/lambda/src/routes/info/read_quota.py +++ b/lambda/src/routes/info/read_quota.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -17,12 +19,12 @@ def __init__(self, storage_manager: StorageManager, quota_kb: int | None = DEFAU self.storage_manager = storage_manager self.quota_kb = quota_kb - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.5//info/quota") - def handle_request(uid: str): + def handle_request(uid: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """ Get quota information for the authenticated user. @@ -31,14 +33,9 @@ def handle(self, event) -> Response: - quota_kb: Storage quota in KB, or null if not enforced """ try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() # Get collections using storage manager to calculate current usage collections = self.storage_manager.list_collections(user_id) diff --git a/lambda/src/routes/info/read_usage.py b/lambda/src/routes/info/read_usage.py index 58cdfdce..699f4bde 100644 --- a/lambda/src/routes/info/read_usage.py +++ b/lambda/src/routes/info/read_usage.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -13,12 +15,12 @@ class ReadCollectionUsageRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.5//info/collection_usage") - def handle_request(uid: str): + def handle_request(uid: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """ Get usage information for all collections. @@ -26,14 +28,9 @@ def handle(self, event) -> Response: Example: {"bookmarks": 1.5, "tabs": 0.5} """ try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() # Get collections using storage manager collections = self.storage_manager.list_collections(user_id) diff --git a/lambda/src/routes/profile/get_profile.py b/lambda/src/routes/profile/get_profile.py index d63db6b4..2e883801 100644 --- a/lambda/src/routes/profile/get_profile.py +++ b/lambda/src/routes/profile/get_profile.py @@ -1,9 +1,11 @@ """GetProfile route — GET /v1/profile (OAuth Bearer auth)""" import json +from typing import Any from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.metrics import Metrics, MetricUnit +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.auth_account_manager import AuthAccountManager from src.services.jwt_verifier import JWTVerifier @@ -25,13 +27,13 @@ def __init__( self._auth_account_manager = auth_account_manager self._metrics = metrics - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/v1/profile") - def handle_get_profile(): + def handle_get_profile() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: - headers = event.headers or {} + def handle(self, event: APIGatewayProxyEvent) -> Response: + headers = event.headers auth_header = headers.get("authorization", "") if not auth_header: diff --git a/lambda/src/routes/storage/delete_all.py b/lambda/src/routes/storage/delete_all.py index 14d28a0c..1c330252 100644 --- a/lambda/src/routes/storage/delete_all.py +++ b/lambda/src/routes/storage/delete_all.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -14,22 +16,17 @@ class DeleteAllStorageRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.delete("/1.5//storage") - def handle_request(uid: str): + def handle_request(uid: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Delete all storage data for the authenticated user""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() # Delete all collections and BSOs for the authenticated user modified_timestamp = self.storage_manager.delete_all_storage(user_id) diff --git a/lambda/src/routes/storage/delete_root.py b/lambda/src/routes/storage/delete_root.py index 6eaebf61..82407ea6 100644 --- a/lambda/src/routes/storage/delete_root.py +++ b/lambda/src/routes/storage/delete_root.py @@ -1,7 +1,9 @@ import json +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.services.storage_manager import StorageManager from src.shared.base_route import BaseRoute @@ -14,22 +16,17 @@ class DeleteAllRootRoute(BaseRoute): def __init__(self, storage_manager: StorageManager): self.storage_manager = storage_manager - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.delete("/1.5/") - def handle_request(uid: str): + def handle_request(uid: str) -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Delete all storage data for the authenticated user (root endpoint alias)""" try: - # Extract user_id from authorizer context - user_id = event.get("requestContext", {}).get("hawk_uid") + user_id = self.hawk_uid(event) if not user_id: - return Response( - status_code=401, - content_type="application/json", - body=json.dumps({"error": "Unauthorized"}), - ) + return self.unauthorized() # Delete all collections and BSOs for the authenticated user modified_timestamp = self.storage_manager.delete_all_storage(user_id) diff --git a/lambda/src/routes/token/request.py b/lambda/src/routes/token/request.py index f274282c..6e30590c 100644 --- a/lambda/src/routes/token/request.py +++ b/lambda/src/routes/token/request.py @@ -2,11 +2,14 @@ import re from dataclasses import asdict +from typing import Any from aws_lambda_powertools import Logger from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.metrics import Metrics, MetricUnit +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent +from src.services.jwt_verifier import JWTVerifier from src.services.token_generator import TokenGenerator from src.services.user_manager import UserManager from src.shared.base_route import BaseRoute @@ -35,24 +38,24 @@ class GetTokenRoute(BaseRoute): def __init__( self, - oidc_validator, + oidc_validator: JWTVerifier, user_manager: UserManager, token_generator: TokenGenerator, metrics: Metrics, retry_after_seconds: int = 30, ): - self.oidc_validator = oidc_validator # OIDCValidator or JWTVerifier + self.oidc_validator = oidc_validator self.user_manager = user_manager self.token_generator = token_generator self._metrics = metrics self.retry_after_seconds = retry_after_seconds - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: @app.get("/1.0/sync/1.5") - def handle_request(): + def handle_request() -> Response[Any]: return self.handle(app.current_event) - def handle(self, event) -> Response: + def handle(self, event: APIGatewayProxyEvent) -> Response: """Handle token creation request.""" try: body = event.body @@ -69,7 +72,7 @@ def handle(self, event) -> Response: source_ip = identity.source_ip if identity else None except KeyError, AttributeError: source_ip = None - headers = event.headers or {} + headers = event.headers user_agent = headers.get("user-agent") auth_header = headers.get("authorization") @@ -265,10 +268,12 @@ def handle(self, event) -> Response: exception_type=type(e).__name__, ) - def _validate_content_type(self, body, event) -> Response | None: + def _validate_content_type( + self, body: str | None, event: APIGatewayProxyEvent + ) -> Response | None: if not body: return None - headers = event.headers or {} + headers = event.headers content_type = headers.get("content-type") if content_type and not self._is_valid_content_type(content_type): logger.warning("Invalid Content-Type", extra={"content_type": content_type}) @@ -303,7 +308,7 @@ def _error_response( name: str, description: str, log_level: str = "warning", - **extra_log_fields, + **extra_log_fields: Any, ) -> Response: log_extra = { "status_code": status_code, diff --git a/lambda/src/services/api_router.py b/lambda/src/services/api_router.py index a034dd42..b9d9a8a0 100644 --- a/lambda/src/services/api_router.py +++ b/lambda/src/services/api_router.py @@ -29,12 +29,12 @@ def __init__( self._register_middleware() self._register_routes() - def _register_exception_handlers(self, handlers: dict): + def _register_exception_handlers(self, handlers: dict) -> None: for exc_type, handler_fn in handlers.items(): self.app.exception_handler(exc_type)(handler_fn) @self.app.exception_handler(RequestValidationError) - def _handle_request_validation(ex: RequestValidationError): # pragma: nocover + def _handle_request_validation(ex: RequestValidationError) -> Response: # pragma: nocover return Response( status_code=422, content_type="application/json", @@ -42,22 +42,22 @@ def _handle_request_validation(ex: RequestValidationError): # pragma: nocover ) @self.app.exception_handler(ResponseValidationError) - def _handle_response_validation(ex: ResponseValidationError): # pragma: nocover + def _handle_response_validation(ex: ResponseValidationError) -> Response: # pragma: nocover return Response( status_code=500, content_type="application/json", body=json.dumps({"error": "Internal server error"}), ) - def _register_middleware(self): + def _register_middleware(self) -> None: """Register middleware handlers""" - self.app.use(middlewares=self._middlewares) # type: ignore + self.app.use(middlewares=list(self._middlewares)) - def _register_routes(self): + def _register_routes(self) -> None: """Register routes by calling each route's bind method""" for route in self._routes: route.bind(self.app) - def handler(self, event: dict, context: LambdaContext): + def handler(self, event: dict, context: LambdaContext) -> dict: """Main Lambda handler entry point""" return self.app.resolve(event=event, context=context) diff --git a/lambda/src/services/auth_account_manager.py b/lambda/src/services/auth_account_manager.py index a9fe398b..ea34dbef 100644 --- a/lambda/src/services/auth_account_manager.py +++ b/lambda/src/services/auth_account_manager.py @@ -2,10 +2,13 @@ import logging import time -from typing import Optional +from typing import TYPE_CHECKING, Any, Optional, cast from botocore.exceptions import ClientError +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + _PK = "PK" ACCOUNT_PREFIX = "ACCOUNT" EMAIL_PREFIX = "EMAIL" @@ -17,7 +20,7 @@ class AuthAccountManager: """Manages FxA account operations with DynamoDB""" - def __init__(self, table): + def __init__(self, table: "Table"): """Initialize AuthAccountManager Args: @@ -153,7 +156,7 @@ def get_account_by_email(self, email: str) -> Optional[dict]: if "Item" not in response: return None - uid = response["Item"]["uid"] + uid = cast(dict[str, Any], response["Item"])["uid"] # Look up ACCOUNT# record return self.get_account_by_uid(uid) @@ -197,7 +200,7 @@ def get_account_by_oidc_sub(self, oidc_sub: str) -> Optional[dict]: if "Item" not in response: return None - uid = response["Item"]["uid"] + uid = cast(dict[str, Any], response["Item"])["uid"] return self.get_account_by_uid(uid) def get_account_by_uid(self, uid: str) -> Optional[dict]: diff --git a/lambda/src/services/channel_service.py b/lambda/src/services/channel_service.py index 72e4258f..c4aaec7d 100644 --- a/lambda/src/services/channel_service.py +++ b/lambda/src/services/channel_service.py @@ -4,9 +4,14 @@ import logging import time import uuid +from typing import TYPE_CHECKING, Any, cast +from boto3.session import Session from botocore.exceptions import ClientError +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + logger = logging.getLogger(__name__) MAX_CONNECTIONS_PER_CHANNEL = 3 @@ -22,12 +27,12 @@ class ChannelService: - CONN#{connectionId} — channelId, expiry """ - def __init__(self, table, session): + def __init__(self, table: "Table", session: Session): self._table = table self._session = session - self._apigw_clients = {} + self._apigw_clients: dict[str, Any] = {} - def handle(self, event, context): + def handle(self, event: dict, context: Any) -> dict[str, Any]: """Dispatch on WebSocket route key.""" route_key = event["requestContext"]["routeKey"] connection_id = event["requestContext"]["connectionId"] @@ -43,7 +48,7 @@ def handle(self, event, context): else: return {"statusCode": 400, "body": "Unknown route"} - def _handle_connect(self, event, connection_id): + def _handle_connect(self, event: dict, connection_id: str) -> dict[str, Any]: """Handle $connect — create or join a channel.""" params = event.get("queryStringParameters") or {} channel_id = params.get("channelId") @@ -54,7 +59,7 @@ def _handle_connect(self, event, connection_id): else: return self._create_channel(event, connection_id, expiry) - def _create_channel(self, event, connection_id, expiry): + def _create_channel(self, event: dict, connection_id: str, expiry: int) -> dict[str, Any]: """Create a new channel with this connection as the first member.""" channel_id = str(uuid.uuid4()) logger.info("Creating channel=%s for connection=%s", channel_id, connection_id) @@ -87,7 +92,7 @@ def _create_channel(self, event, connection_id, expiry): return {"statusCode": 200} - def _join_channel(self, channel_id, connection_id, expiry): + def _join_channel(self, channel_id: str, connection_id: str, expiry: int) -> dict[str, Any]: """Join an existing channel atomically.""" logger.info("Joining channel=%s connection=%s", channel_id, connection_id) try: @@ -121,14 +126,14 @@ def _join_channel(self, channel_id, connection_id, expiry): return {"statusCode": 200} - def _handle_disconnect(self, connection_id): + def _handle_disconnect(self, connection_id: str) -> None: """Handle disconnect — remove connection from channel.""" # Delete reverse lookup first (idempotent guard against double-disconnect) result = self._table.get_item(Key={"PK": f"CONN#{connection_id}"}) if "Item" not in result: return - channel_id = result["Item"]["channelId"] + channel_id = cast(dict[str, Any], result["Item"])["channelId"] self._table.delete_item(Key={"PK": f"CONN#{connection_id}"}) # Get channel to find connection index @@ -136,7 +141,7 @@ def _handle_disconnect(self, connection_id): if "Item" not in channel_result: return - connections = channel_result["Item"]["connections"] + connections = cast(dict[str, Any], channel_result["Item"])["connections"] if connection_id in connections: index = connections.index(connection_id) self._table.update_item( @@ -144,7 +149,7 @@ def _handle_disconnect(self, connection_id): UpdateExpression=f"REMOVE connections[{index}]", ) - def _handle_message(self, event, connection_id): + def _handle_message(self, event: dict, connection_id: str) -> dict[str, Any]: """Handle incoming message — relay to other connections.""" logger.info("Message from connection=%s", connection_id) # Look up channel for this connection @@ -152,7 +157,7 @@ def _handle_message(self, event, connection_id): if "Item" not in result: return {"statusCode": 404, "body": "Connection not found"} - channel_id = result["Item"]["channelId"] + channel_id = cast(dict[str, Any], result["Item"])["channelId"] # Atomic message count increment with limit check try: @@ -180,14 +185,16 @@ def _handle_message(self, event, connection_id): if "Item" not in channel_result: return {"statusCode": 404, "body": "Channel not found"} - connections = channel_result["Item"]["connections"] + connections = cast(dict[str, Any], channel_result["Item"])["connections"] message_body = event.get("body", "") self._relay_message(event, connection_id, connections, message_body) return {"statusCode": 200} - def _relay_message(self, event, sender_connection_id, connections, message_body): + def _relay_message( + self, event: dict, sender_connection_id: str, connections: list, message_body: str + ) -> None: """Relay message to all connections except sender.""" data = json.dumps( { @@ -199,7 +206,7 @@ def _relay_message(self, event, sender_connection_id, connections, message_body) if conn_id != sender_connection_id: self._post_to_connection(event, conn_id, data) - def _post_to_connection(self, event, connection_id, data): + def _post_to_connection(self, event: dict, connection_id: str, data: str) -> None: """Post data to a WebSocket connection via API Gateway Management API.""" client = self._get_apigw_client(event) try: @@ -211,7 +218,7 @@ def _post_to_connection(self, event, connection_id, data): logger.warning("Connection %s is gone, cleaning up", connection_id) self._handle_disconnect(connection_id) - def _get_apigw_client(self, event): + def _get_apigw_client(self, event: dict) -> Any: """Lazy API Gateway Management API client, cached by endpoint. Uses the execute-api domain (not the custom domain) because the diff --git a/lambda/src/services/device_manager.py b/lambda/src/services/device_manager.py index cb0369ab..d9c11467 100644 --- a/lambda/src/services/device_manager.py +++ b/lambda/src/services/device_manager.py @@ -3,10 +3,13 @@ import re import time import uuid -from typing import Optional +from typing import TYPE_CHECKING, Any, Optional, cast from boto3.dynamodb.conditions import Attr +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + DEVICE_PREFIX = "DEVICE" _HAWK_ID_PATTERN = re.compile(r'id="([^"]+)"') @@ -14,7 +17,7 @@ class DeviceManager: """Manages FxA device records in DynamoDB.""" - def __init__(self, table): + def __init__(self, table: "Table"): self.table = table def _device_pk(self, uid: str, device_id: str) -> str: @@ -60,7 +63,7 @@ def get_devices(self, uid: str, filter_idle_timestamp: Optional[int] = None) -> FilterExpression=Attr("PK").begins_with(f"{DEVICE_PREFIX}#{uid}#") ) devices = [] - for item in response.get("Items", []): + for item in cast(list[dict[str, Any]], response.get("Items", [])): item.pop("PK", None) if filter_idle_timestamp and item.get("lastAccessTime", 0) < filter_idle_timestamp: continue diff --git a/lambda/src/services/fxa_token_manager.py b/lambda/src/services/fxa_token_manager.py index 06c6ef21..c54d2a65 100644 --- a/lambda/src/services/fxa_token_manager.py +++ b/lambda/src/services/fxa_token_manager.py @@ -2,16 +2,19 @@ import re import time -from typing import Optional +from typing import TYPE_CHECKING, Any, Callable, Optional, cast import mohawk import mohawk.exc from aws_lambda_powertools import Logger -from aws_lambda_powertools.metrics import MetricUnit +from aws_lambda_powertools.metrics import Metrics, MetricUnit from botocore.exceptions import ClientError from src.services import fxa_crypto +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + logger = Logger(child=True) _PK = "PK" @@ -37,8 +40,8 @@ def extract_token_id_from_hawk_header(authorization_header: str) -> str | None: def __init__( self, - table, - metrics, + table: "Table", + metrics: Metrics, session_ttl_seconds: int = 2592000, keyfetch_ttl_seconds: int = 300, ): @@ -109,7 +112,7 @@ def verify_session_token_id(self, token_id_hex: str) -> Optional[str]: if "Item" not in response: return None - item = response["Item"] + item = cast(dict[str, Any], response["Item"]) # Check expiry server-side if item.get("expiry", 0) < int(time.time()): @@ -117,7 +120,7 @@ def verify_session_token_id(self, token_id_hex: str) -> Optional[str]: return item["uid"] - def _seen_nonce(self, sender_id, nonce, timestamp): + def _seen_nonce(self, sender_id: str, nonce: str, timestamp: str) -> bool: """Check if a nonce has been seen before (replay protection). Uses DynamoDB conditional write: if the nonce record already exists, @@ -137,7 +140,15 @@ def _seen_nonce(self, sender_id, nonce, timestamp): return True # Replay detected raise - def _verify_hawk(self, authorization_header, method, path, host, port, credentials_map): + def _verify_hawk( + self, + authorization_header: str, + method: str, + path: str, + host: str, + port: int, + credentials_map: Callable[[str], dict], + ) -> bool: """Verify Hawk signature using mohawk.Receiver. Returns True on success, False on any authentication failure. @@ -185,11 +196,11 @@ def verify_session_hawk( """Verify Hawk HMAC signature for session-authenticated routes.""" uid_holder = {} - def credentials_map(sender_id): + def credentials_map(sender_id: str) -> dict: response = self.table.get_item(Key={_PK: f"{SESSION_PREFIX}#{sender_id}"}) if "Item" not in response: raise mohawk.exc.CredentialsLookupError("Session not found") - item = response["Item"] + item = cast(dict[str, Any], response["Item"]) if item.get("expiry", 0) < int(time.time()): raise mohawk.exc.CredentialsLookupError("Session expired") key = item.get("reqHMACkey") @@ -217,7 +228,7 @@ def verify_keyfetch_hawk( """ result_holder = {} - def credentials_map(sender_id): + def credentials_map(sender_id: str) -> dict: try: response = self.table.delete_item( Key={_PK: f"{KEYFETCH_PREFIX}#{sender_id}"}, @@ -228,7 +239,7 @@ def credentials_map(sender_id): if e.response["Error"]["Code"] == "ConditionalCheckFailedException": raise mohawk.exc.CredentialsLookupError("Token not found") raise - item = response.get("Attributes") + item = cast(Optional[dict[str, Any]], response.get("Attributes")) if not item: raise mohawk.exc.CredentialsLookupError("Token not found") if item.get("expiry", 0) < int(time.time()): @@ -306,7 +317,7 @@ def consume_key_fetch_token(self, token_id_hex: str) -> Optional[dict]: return None raise - item = response.get("Attributes") + item = cast(Optional[dict[str, Any]], response.get("Attributes")) if not item: return None diff --git a/lambda/src/services/hawk_service.py b/lambda/src/services/hawk_service.py index 2e3d0ff2..b131cf5e 100644 --- a/lambda/src/services/hawk_service.py +++ b/lambda/src/services/hawk_service.py @@ -14,7 +14,7 @@ import time from dataclasses import dataclass from itertools import permutations -from typing import Optional, Tuple +from typing import TYPE_CHECKING, Any, Optional, Tuple, cast import mohawk import mohawk.exc @@ -29,6 +29,9 @@ InvalidHawkSignatureException, ) +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + logger = Logger(child=True) # Pattern to extract hawk id from header without full parse @@ -60,7 +63,10 @@ class HawkService: """ def __init__( - self, token_cache_table, timestamp_skew_tolerance: int = 60, token_duration: int = 300 + self, + token_cache_table: "Table", + timestamp_skew_tolerance: int = 60, + token_duration: int = 300, ): self.token_cache_table = token_cache_table self.timestamp_skew_tolerance = timestamp_skew_tolerance @@ -103,7 +109,7 @@ def validate( path = corrected # Credentials lookup called by mohawk during MAC verification - def credentials_map(sender_id): + def credentials_map(sender_id: str) -> dict: hawk_key, cached_user_id, cached_generation = self.get_hawk_key_from_cache(sender_id) if cached_generation != generation: raise InvalidGenerationException( @@ -149,7 +155,7 @@ def credentials_map(sender_id): user_id=user_id, generation=generation, expiry=expiry, hawk_id=hawk_id ) - def _seen_nonce(self, sender_id, nonce, timestamp): + def _seen_nonce(self, sender_id: str, nonce: str, timestamp: str) -> bool: """Check if a nonce has been seen before (replay protection). Uses DynamoDB conditional write: if the nonce record already exists, @@ -264,7 +270,7 @@ def get_hawk_key_from_cache(self, hawk_id: str) -> Tuple[str, str, int]: response = self.token_cache_table.get_item(Key={"PK": f"TOKEN#{hawk_id}"}) if "Item" not in response: raise AuthenticationException(f"HAWK token not found: {hawk_id}") - item = response["Item"] + item = cast(dict[str, Any], response["Item"]) return (item["hawk_key"], item["user_id"], int(item["generation"])) except ClientError as e: logger.error(f"Failed to retrieve HAWK token from cache: {e}") diff --git a/lambda/src/services/jwt_verifier.py b/lambda/src/services/jwt_verifier.py index 8b98298e..6a705d0c 100644 --- a/lambda/src/services/jwt_verifier.py +++ b/lambda/src/services/jwt_verifier.py @@ -6,7 +6,7 @@ from functools import cached_property from cryptography.hazmat.primitives.asymmetric.padding import PKCS1v15 -from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicNumbers +from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicKey, RSAPublicNumbers from cryptography.hazmat.primitives.hashes import SHA256 from src.services.jwt_service import JWTService @@ -82,7 +82,7 @@ def validate_token(self, token: str) -> OIDCTokenClaims: ) @cached_property - def _public_key(self): + def _public_key(self) -> RSAPublicKey: """Construct and cache the RSA public key from JWK.""" jwk = self._jwt_service.get_public_key_jwk() n_bytes = self._b64url_decode(jwk["n"]) diff --git a/lambda/src/services/storage_manager.py b/lambda/src/services/storage_manager.py index 9401b267..03dd457c 100644 --- a/lambda/src/services/storage_manager.py +++ b/lambda/src/services/storage_manager.py @@ -1,7 +1,7 @@ """Storage manager for DynamoDB operations""" from decimal import Decimal -from typing import Dict, List, Optional +from typing import TYPE_CHECKING, Any, Dict, List, Optional from botocore.exceptions import ClientError @@ -20,6 +20,9 @@ get_current_timestamp, ) +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + _PK = "PK" _SK = "SK" @@ -32,7 +35,7 @@ class StorageManager: """Manages storage operations with DynamoDB""" - def __init__(self, table): + def __init__(self, table: "Table"): """Initialize StorageManager Args: @@ -647,7 +650,7 @@ def update_storage_object( collection_name: str, object_id: str, if_unmodified_since: Optional[float] = None, - **kwargs, + **kwargs: Any, ) -> BasicStorageObject: """Update a storage object diff --git a/lambda/src/services/user_manager.py b/lambda/src/services/user_manager.py index 90b753d7..07b0b3fc 100644 --- a/lambda/src/services/user_manager.py +++ b/lambda/src/services/user_manager.py @@ -2,13 +2,16 @@ import time from decimal import Decimal -from typing import List, Optional +from typing import TYPE_CHECKING, Any, List, Optional, cast from botocore.exceptions import ClientError from src.shared.exceptions import InvalidClientStateError, ServiceUnavailableError from src.shared.user import UserRecord +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + _PK = "PK" PK_PREFIX = "USER" MAX_CLIENT_STATE_HISTORY = 50 @@ -17,7 +20,7 @@ class UserManager: """Manages user operations with DynamoDB for the Token Server""" - def __init__(self, table): + def __init__(self, table: "Table"): """Initialize UserManager Args: @@ -298,7 +301,7 @@ def increment_generation(self, user_id: str) -> int: ReturnValues="ALL_NEW", ) - updated_item = response["Attributes"] + updated_item = cast(dict[str, Any], response["Attributes"]) return updated_item["generation"] except ClientError as e: diff --git a/lambda/src/shared/base_route.py b/lambda/src/shared/base_route.py index 7d0c420c..01f97ce6 100644 --- a/lambda/src/shared/base_route.py +++ b/lambda/src/shared/base_route.py @@ -1,8 +1,10 @@ +import json from abc import ABC, abstractmethod from typing import Sequence -from aws_lambda_powertools.event_handler import APIGatewayRestResolver +from aws_lambda_powertools.event_handler import APIGatewayRestResolver, Response from aws_lambda_powertools.event_handler.middlewares import BaseMiddlewareHandler +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent class BaseRoute(ABC): @@ -11,11 +13,25 @@ class BaseRoute(ABC): middlewares: Sequence[BaseMiddlewareHandler] = () @abstractmethod - def bind(self, app: APIGatewayRestResolver): + def bind(self, app: APIGatewayRestResolver) -> None: """Bind this route to the API with appropriate decorators""" pass # pragma: nocover @abstractmethod - def handle(self, event): + def handle(self, event: APIGatewayProxyEvent) -> Response: """Handle the route request""" pass # pragma: nocover + + @staticmethod + def hawk_uid(event: APIGatewayProxyEvent) -> str | None: + """Authenticated user id injected by HawkAuthMiddleware; None if absent.""" + return (event.get("requestContext") or {}).get("hawk_uid") + + @staticmethod + def unauthorized() -> Response: + """The 401 every storage route returns when no authenticated user is present.""" + return Response( + status_code=401, + content_type="application/json", + body=json.dumps({"error": "Unauthorized"}), + ) diff --git a/lambda/src/shared/exceptions.py b/lambda/src/shared/exceptions.py index 742539d0..6d315c99 100644 --- a/lambda/src/shared/exceptions.py +++ b/lambda/src/shared/exceptions.py @@ -19,15 +19,16 @@ class SyncStorageException(Exception): status_code = HTTPStatus.INTERNAL_SERVER_ERROR error_code = "InternalServerError" mozilla_code: Optional[int] = None # Mozilla response code (integer) for specific errors + default_message = "Internal server error" # Subclasses override; used when message is omitted def __init__( self, - message: str = "Internal server error", + message: Optional[str] = None, retry_after: Optional[int] = None, backoff: Optional[int] = None, alert: Optional[str] = None, ): - self.message = message + self.message = message if message is not None else self.default_message self.retry_after = retry_after # Retry-After header (seconds) self.backoff = backoff # X-Weave-Backoff header (seconds) self.alert = alert # X-Weave-Alert header (message or JSON) @@ -75,9 +76,7 @@ class ValidationException(SyncStorageException): status_code = HTTPStatus.BAD_REQUEST error_code = "ValidationException" - - def __init__(self, message: str = "Invalid request parameters", **kwargs): - super().__init__(message, **kwargs) + default_message = "Invalid request parameters" class ConflictException(SyncStorageException): @@ -85,9 +84,7 @@ class ConflictException(SyncStorageException): status_code = HTTPStatus.CONFLICT error_code = "ConflictException" - - def __init__(self, message: str = "Resource conflict", **kwargs): - super().__init__(message, **kwargs) + default_message = "Resource conflict" class PreconditionFailedException(SyncStorageException): @@ -95,9 +92,7 @@ class PreconditionFailedException(SyncStorageException): status_code = HTTPStatus.PRECONDITION_FAILED error_code = "PreconditionFailedException" - - def __init__(self, message: str = "Precondition failed", **kwargs): - super().__init__(message, **kwargs) + default_message = "Precondition failed" class QuotaExceededException(SyncStorageException): @@ -106,9 +101,7 @@ class QuotaExceededException(SyncStorageException): status_code = HTTPStatus.INSUFFICIENT_STORAGE error_code = "QuotaExceededException" mozilla_code = CODE_QUOTA_EXCEEDED - - def __init__(self, message: str = "Storage quota exceeded", **kwargs): - super().__init__(message, **kwargs) + default_message = "Storage quota exceeded" class CollectionNotFoundException(SyncStorageException): @@ -116,9 +109,7 @@ class CollectionNotFoundException(SyncStorageException): status_code = HTTPStatus.NOT_FOUND error_code = "CollectionNotFoundException" - - def __init__(self, message: str = "Collection not found", **kwargs): - super().__init__(message, **kwargs) + default_message = "Collection not found" class StorageObjectNotFoundException(SyncStorageException): @@ -126,9 +117,7 @@ class StorageObjectNotFoundException(SyncStorageException): status_code = HTTPStatus.NOT_FOUND error_code = "StorageObjectNotFoundException" - - def __init__(self, message: str = "Storage object not found", **kwargs): - super().__init__(message, **kwargs) + default_message = "Storage object not found" class AuthenticationException(SyncStorageException): @@ -136,9 +125,7 @@ class AuthenticationException(SyncStorageException): status_code = HTTPStatus.UNAUTHORIZED error_code = "AuthenticationException" - - def __init__(self, message: str = "Authentication required", **kwargs): - super().__init__(message, **kwargs) + default_message = "Authentication required" class RequestTooLargeException(SyncStorageException): @@ -146,9 +133,7 @@ class RequestTooLargeException(SyncStorageException): status_code = HTTPStatus.REQUEST_ENTITY_TOO_LARGE error_code = "RequestTooLargeException" - - def __init__(self, message: str = "Request entity too large", **kwargs): - super().__init__(message, **kwargs) + default_message = "Request entity too large" class MethodNotAllowedException(SyncStorageException): @@ -156,9 +141,7 @@ class MethodNotAllowedException(SyncStorageException): status_code = HTTPStatus.METHOD_NOT_ALLOWED error_code = "MethodNotAllowedException" - - def __init__(self, message: str = "Method not allowed", **kwargs): - super().__init__(message, **kwargs) + default_message = "Method not allowed" class UnsupportedMediaTypeException(SyncStorageException): @@ -166,9 +149,7 @@ class UnsupportedMediaTypeException(SyncStorageException): status_code = HTTPStatus.UNSUPPORTED_MEDIA_TYPE error_code = "UnsupportedMediaTypeException" - - def __init__(self, message: str = "Unsupported media type", **kwargs): - super().__init__(message, **kwargs) + default_message = "Unsupported media type" class ServerLimitExceededException(SyncStorageException): @@ -177,9 +158,7 @@ class ServerLimitExceededException(SyncStorageException): status_code = HTTPStatus.BAD_REQUEST error_code = "ServerLimitExceededException" mozilla_code = CODE_SERVER_LIMIT_EXCEEDED - - def __init__(self, message: str = "Server limit exceeded", **kwargs): - super().__init__(message, **kwargs) + default_message = "Server limit exceeded" class InvalidBSOException(SyncStorageException): @@ -188,9 +167,7 @@ class InvalidBSOException(SyncStorageException): status_code = HTTPStatus.BAD_REQUEST error_code = "InvalidBSOException" mozilla_code = CODE_INVALID_BSO - - def __init__(self, message: str = "Invalid BSO", **kwargs): - super().__init__(message, **kwargs) + default_message = "Invalid BSO" class InvalidCollectionException(SyncStorageException): @@ -199,9 +176,7 @@ class InvalidCollectionException(SyncStorageException): status_code = HTTPStatus.BAD_REQUEST error_code = "InvalidCollectionException" mozilla_code = CODE_INVALID_COLLECTION - - def __init__(self, message: str = "Invalid collection name", **kwargs): - super().__init__(message, **kwargs) + default_message = "Invalid collection name" class JSONParseException(SyncStorageException): @@ -210,9 +185,7 @@ class JSONParseException(SyncStorageException): status_code = HTTPStatus.BAD_REQUEST error_code = "JSONParseException" mozilla_code = CODE_JSON_PARSE_FAILURE - - def __init__(self, message: str = "JSON parse failure", **kwargs): - super().__init__(message, **kwargs) + default_message = "JSON parse failure" class IncompatibleClientException(SyncStorageException): @@ -221,9 +194,7 @@ class IncompatibleClientException(SyncStorageException): status_code = HTTPStatus.BAD_REQUEST error_code = "IncompatibleClientException" mozilla_code = CODE_INCOMPATIBLE_CLIENT - - def __init__(self, message: str = "Incompatible client", **kwargs): - super().__init__(message, **kwargs) + default_message = "Incompatible client" # Token Server specific exceptions @@ -234,9 +205,7 @@ class InvalidTokenError(SyncStorageException): status_code = HTTPStatus.UNAUTHORIZED error_code = "InvalidTokenError" - - def __init__(self, message: str = "Invalid or expired token", **kwargs): - super().__init__(message, **kwargs) + default_message = "Invalid or expired token" class InvalidCredentialsError(SyncStorageException): @@ -244,16 +213,13 @@ class InvalidCredentialsError(SyncStorageException): status_code = HTTPStatus.UNAUTHORIZED error_code = "InvalidCredentialsError" - - def __init__(self, message: str = "Invalid credentials", **kwargs): - super().__init__(message, **kwargs) + default_message = "Invalid credentials" class TokenValidationError(ValidationException): """Raised when token validation fails""" - def __init__(self, message: str = "Token validation failed", **kwargs): - super().__init__(message, **kwargs) + default_message = "Token validation failed" class ServiceUnavailableError(SyncStorageException): @@ -261,9 +227,7 @@ class ServiceUnavailableError(SyncStorageException): status_code = HTTPStatus.SERVICE_UNAVAILABLE error_code = "ServiceUnavailableError" - - def __init__(self, message: str = "Service temporarily unavailable", **kwargs): - super().__init__(message, **kwargs) + default_message = "Service temporarily unavailable" class InvalidTimestampError(SyncStorageException): @@ -272,11 +236,7 @@ class InvalidTimestampError(SyncStorageException): status_code = HTTPStatus.UNAUTHORIZED error_code = "InvalidTimestampError" status_field = "invalid-timestamp" - - def __init__( - self, message: str = "Token timestamp differs significantly from server time", **kwargs - ): - super().__init__(message, **kwargs) + default_message = "Token timestamp differs significantly from server time" class InvalidGenerationError(SyncStorageException): @@ -285,9 +245,7 @@ class InvalidGenerationError(SyncStorageException): status_code = HTTPStatus.UNAUTHORIZED error_code = "InvalidGenerationError" status_field = "invalid-generation" - - def __init__(self, message: str = "Token generation number is outdated", **kwargs): - super().__init__(message, **kwargs) + default_message = "Token generation number is outdated" class InvalidClientStateError(SyncStorageException): @@ -296,9 +254,7 @@ class InvalidClientStateError(SyncStorageException): status_code = HTTPStatus.UNAUTHORIZED error_code = "InvalidClientStateError" status_field = "invalid-client-state" - - def __init__(self, message: str = "Invalid client state transition", **kwargs): - super().__init__(message, **kwargs) + default_message = "Invalid client state transition" class NewUsersDisabledError(SyncStorageException): @@ -307,9 +263,7 @@ class NewUsersDisabledError(SyncStorageException): status_code = HTTPStatus.UNAUTHORIZED error_code = "NewUsersDisabledError" status_field = "new-users-disabled" - - def __init__(self, message: str = "New user registration is disabled", **kwargs): - super().__init__(message, **kwargs) + default_message = "New user registration is disabled" # HAWK Authentication specific exceptions @@ -318,26 +272,22 @@ def __init__(self, message: str = "New user registration is disabled", **kwargs) class InvalidHawkHeaderException(AuthenticationException): """Raised when HAWK Authorization header is malformed""" - def __init__(self, message: str = "Malformed HAWK Authorization header", **kwargs): - super().__init__(message, **kwargs) + default_message = "Malformed HAWK Authorization header" class InvalidHawkSignatureException(AuthenticationException): """Raised when HAWK signature verification fails""" - def __init__(self, message: str = "HAWK signature verification failed", **kwargs): - super().__init__(message, **kwargs) + default_message = "HAWK signature verification failed" class ExpiredHawkTokenException(AuthenticationException): """Raised when HAWK token has expired""" - def __init__(self, message: str = "HAWK token has expired", **kwargs): - super().__init__(message, **kwargs) + default_message = "HAWK token has expired" class InvalidGenerationException(AuthenticationException): """Raised when HAWK token has an outdated generation number""" - def __init__(self, message: str = "HAWK token generation number is outdated", **kwargs): - super().__init__(message, **kwargs) + default_message = "HAWK token generation number is outdated" diff --git a/lambda/src/shared/models.py b/lambda/src/shared/models.py index be2ad0aa..c6e8e440 100644 --- a/lambda/src/shared/models.py +++ b/lambda/src/shared/models.py @@ -10,7 +10,7 @@ import re import time from decimal import Decimal -from typing import Annotated +from typing import Annotated, Any from pydantic import BaseModel, ConfigDict, Field, StringConstraints, TypeAdapter from pydantic import ValidationError as PydanticValidationError @@ -76,7 +76,7 @@ def to_dynamo_dict(model: BaseModel) -> dict: return {k: _to_dynamo(v) for k, v in model.model_dump().items()} -def _to_dynamo(v): +def _to_dynamo(v: Any) -> Any: if isinstance(v, float): return Decimal(str(v)) if isinstance(v, dict): diff --git a/lambda/src/shared/utils.py b/lambda/src/shared/utils.py index 909f0eb2..848fde02 100644 --- a/lambda/src/shared/utils.py +++ b/lambda/src/shared/utils.py @@ -1,5 +1,7 @@ from datetime import datetime, timezone +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent + def get_weave_timestamp() -> str: """ @@ -14,7 +16,7 @@ def get_weave_timestamp() -> str: return f"{datetime.now(timezone.utc).timestamp():.2f}" -def extract_hawk_request_params(event) -> tuple[str, str, str, int]: +def extract_hawk_request_params(event: APIGatewayProxyEvent) -> tuple[str, str, str, int]: """Extract (method, path, host, port) for Hawk MAC verification. Uses request_context.domain_name (the custom domain) rather than diff --git a/lambda/tests/conftest.py b/lambda/tests/conftest.py index 28bfb569..6305bba7 100644 --- a/lambda/tests/conftest.py +++ b/lambda/tests/conftest.py @@ -1,12 +1,14 @@ """Shared test fixtures and configuration""" import json +from typing import Any, Generator from unittest.mock import MagicMock, Mock, patch +import boto3 import pytest +from aws_lambda_powertools.event_handler import Response from src.environment.service_provider import ServiceProvider -from src.services.token_generator import TokenGenerator from src.shared.models import ( BasicStorageObject, BatchResult, @@ -16,49 +18,49 @@ @pytest.fixture -def storage_table_name(): +def storage_table_name() -> str: return "test-storage-table" @pytest.fixture -def token_users_table_name(): +def token_users_table_name() -> str: return "test-token-users-table" @pytest.fixture -def oidc_provider_url(): +def oidc_provider_url() -> str: return "https://auth.example.com" @pytest.fixture -def oidc_client_id(): +def oidc_client_id() -> str: return "test-client-id" @pytest.fixture -def base_domain(): +def base_domain() -> str: return "sync.example.com" @pytest.fixture -def token_cache_table_name(): +def token_cache_table_name() -> str: return "test-token-cache-table" @pytest.fixture(autouse=True) def setup_environment( - monkeypatch, - aws_region_name, - aws_access_key_id, - aws_secret_access_key, - aws_session_token, - storage_table_name, - token_users_table_name, - token_cache_table_name, - oidc_provider_url, - oidc_client_id, - base_domain, -): + monkeypatch: pytest.MonkeyPatch, + aws_region_name: str, + aws_access_key_id: str, + aws_secret_access_key: str, + aws_session_token: str, + storage_table_name: str, + token_users_table_name: str, + token_cache_table_name: str, + oidc_provider_url: str, + oidc_client_id: str, + base_domain: str, +) -> None: """Mock environment variables""" monkeypatch.setenv("AWS_REGION", aws_region_name) monkeypatch.setenv("AWS_ACCESS_KEY_ID", aws_access_key_id) @@ -80,12 +82,12 @@ def setup_environment( @pytest.fixture -def mock_service_provider(boto_session): +def mock_service_provider(boto_session: boto3.Session) -> ServiceProvider: return ServiceProvider() @pytest.fixture -def mock_storage_manager(): +def mock_storage_manager() -> MagicMock: """Mock StorageManager for testing route handlers""" manager = MagicMock() @@ -123,44 +125,25 @@ def mock_storage_manager(): @pytest.fixture -def test_user_id(): +def test_user_id() -> str: """Test user ID for authenticated requests""" return "test-user-123" -def make_event_with_auth(event_dict: dict, user_id: str = "test-user-123") -> dict: - """Helper to add hawk_uid to an event dict""" - if "requestContext" not in event_dict: - event_dict["requestContext"] = {} - event_dict["requestContext"]["hawk_uid"] = user_id - return event_dict +def json_body(response: Response[Any]) -> Any: + """Decode a route Response body as JSON. powertools types it Optional; assert once here.""" + assert response.body is not None, f"expected a JSON body, got {response.status_code}" + return json.loads(response.body) -@pytest.fixture -def sample_lambda_event(test_user_id): - """Sample Lambda event structure""" - uid = str(TokenGenerator.generate_uid(test_user_id, 0)) - return { - "httpMethod": "GET", - "path": f"/1.5/{uid}/storage/test_collection/test_object", - "pathParameters": { - "uid": uid, - "collectionName": "test_collection", - "objectId": "test_object", - }, - "headers": {"Content-Type": "application/json"}, - "body": None, - "queryStringParameters": None, - "requestContext": { - "requestId": "test-request-id", - "accountId": "123456789012", - "hawk_uid": test_user_id, - }, - } +def header(response: Response[Any], name: str) -> str: + """Single response header value (powertools types them str | list[str]).""" + value = response.headers[name] + return value if isinstance(value, str) else value[0] @pytest.fixture -def sample_lambda_context(): +def sample_lambda_context() -> Mock: """Sample Lambda context object""" context = Mock() context.function_name = "test-function" @@ -174,7 +157,7 @@ def sample_lambda_context(): @pytest.fixture -def sample_bso(): +def sample_bso() -> BasicStorageObject: """Sample BasicStorageObject""" return BasicStorageObject( id="test_bso", @@ -186,7 +169,7 @@ def sample_bso(): @pytest.fixture -def sample_collection(): +def sample_collection() -> CollectionData: """Sample CollectionData""" return CollectionData( name="bookmarks", @@ -197,7 +180,7 @@ def sample_collection(): @pytest.fixture -def sample_batch_result(): +def sample_batch_result() -> BatchResult: """Sample BatchResult""" return BatchResult( success=["obj1", "obj2", "obj3"], @@ -206,53 +189,15 @@ def sample_batch_result(): ) -@pytest.fixture -def post_event_with_body(test_user_id): - """Sample POST event with body""" - uid = str(TokenGenerator.generate_uid(test_user_id, 0)) - return { - "httpMethod": "POST", - "path": f"/1.5/{uid}/storage/test_collection", - "pathParameters": {"uid": uid, "collectionName": "test_collection"}, - "headers": {"Content-Type": "application/json"}, - "body": json.dumps({"objects": [{"id": "obj1", "payload": "data1", "sortindex": 100}]}), - "queryStringParameters": None, - } - - -@pytest.fixture -def delete_event(test_user_id): - """Sample DELETE event""" - uid = str(TokenGenerator.generate_uid(test_user_id, 0)) - return { - "httpMethod": "DELETE", - "path": f"/1.5/{uid}/storage/test_collection/test_object", - "pathParameters": { - "uid": uid, - "collectionName": "test_collection", - "objectId": "test_object", - }, - "headers": {}, - "body": None, - "queryStringParameters": None, - } - - # Timestamp fixtures for testing @pytest.fixture -def mock_timestamp(): +def mock_timestamp() -> float: """Mock timestamp value used across tests""" return 1234567890.00 @pytest.fixture -def mock_timestamp_datetime(mock_timestamp): - """Mock timestamp (kept for backwards-compat with existing test sigs).""" - return mock_timestamp - - -@pytest.fixture -def mock_datetime_now(mock_timestamp): +def mock_datetime_now(mock_timestamp: float) -> Generator[None, None, None]: """Mock time.time() for user_manager tests""" with patch("src.services.user_manager.time") as mock: mock.time.return_value = mock_timestamp @@ -260,22 +205,22 @@ def mock_datetime_now(mock_timestamp): @pytest.fixture -def mock_get_current_timestamp(mock_timestamp): +def mock_get_current_timestamp(mock_timestamp: float) -> Generator[None, None, None]: """Mock get_current_timestamp() for storage_manager tests""" with patch("src.services.storage_manager.get_current_timestamp", return_value=mock_timestamp): yield @pytest.fixture -def base_url(): +def base_url() -> str: return "sync.example.com" @pytest.fixture -def storage_domain(base_url): +def storage_domain(base_url: str) -> str: return f"storage.{base_url}" @pytest.fixture -def storage_url(storage_domain): +def storage_url(storage_domain: str) -> str: return f"https://{storage_domain}" diff --git a/lambda/tests/entrypoint/test_auth_api.py b/lambda/tests/entrypoint/test_auth_api.py index 70b929f4..4891bb3a 100644 --- a/lambda/tests/entrypoint/test_auth_api.py +++ b/lambda/tests/entrypoint/test_auth_api.py @@ -1,15 +1,19 @@ """Tests for Auth API lambda entrypoint""" import json +from unittest.mock import Mock from src.entrypoint import auth_api_handler +from src.environment.service_provider import ServiceProvider from src.services.api_router import ApiRouter class TestAuthApiErrors: """Tests for auth API error handling""" - def test_unknown_route_returns_404(self, mock_service_provider, sample_lambda_context): + def test_unknown_route_returns_404( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test request to unknown path returns 404""" event = { "httpMethod": "GET", @@ -23,8 +27,8 @@ def test_unknown_route_returns_404(self, mock_service_provider, sample_lambda_co assert result["statusCode"] == 404 def test_session_route_without_auth_returns_401( - self, mock_service_provider, sample_lambda_context - ): + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test session-protected route without auth header returns 401 via exception handler. The HawkAuthMiddleware raises HawkAuthenticationError which is caught @@ -50,7 +54,9 @@ def test_session_route_without_auth_returns_401( class TestServiceProviderAuthApiProperties: """Tests for ServiceProvider auth API property initialization""" - def test_auth_api_router_creates_router_with_routes(self, mock_service_provider): + def test_auth_api_router_creates_router_with_routes( + self, mock_service_provider: ServiceProvider + ) -> None: """Test auth_api_router creates ApiRouter with auth routes""" router = mock_service_provider.auth_api_router diff --git a/lambda/tests/entrypoint/test_channel_api.py b/lambda/tests/entrypoint/test_channel_api.py index c9724725..d4de6fdd 100644 --- a/lambda/tests/entrypoint/test_channel_api.py +++ b/lambda/tests/entrypoint/test_channel_api.py @@ -1,13 +1,14 @@ """Tests for Channel API lambda entrypoint""" -from unittest.mock import MagicMock +from unittest.mock import MagicMock, Mock from src.entrypoint import channel_api_handler +from src.environment.service_provider import ServiceProvider from src.services.channel_service import ChannelService class TestChannelApiHandler: - def test_delegates_to_channel_service(self, sample_lambda_context): + def test_delegates_to_channel_service(self, sample_lambda_context: Mock) -> None: """Handler delegates to channel_service.handle.""" mock_channel = MagicMock(spec=ChannelService) mock_channel.handle.return_value = {"statusCode": 200} @@ -33,16 +34,18 @@ def test_delegates_to_channel_service(self, sample_lambda_context): class TestServiceProviderChannelProperties: """Tests for ServiceProvider channel property initialization""" - def test_channel_table_name_from_env(self, mock_service_provider): + def test_channel_table_name_from_env(self, mock_service_provider: ServiceProvider) -> None: """Test channel_table_name reads from environment.""" assert mock_service_provider.channel_table_name == "test-channel-table" - def test_channel_table_creates_table_resource(self, mock_service_provider): + def test_channel_table_creates_table_resource( + self, mock_service_provider: ServiceProvider + ) -> None: """Test channel_table returns a DynamoDB Table resource.""" table = mock_service_provider.channel_table assert table is not None - def test_channel_service_creates_instance(self, mock_service_provider): + def test_channel_service_creates_instance(self, mock_service_provider: ServiceProvider) -> None: """Test channel_service property creates ChannelService.""" service = mock_service_provider.channel_service assert isinstance(service, ChannelService) diff --git a/lambda/tests/entrypoint/test_profile_api.py b/lambda/tests/entrypoint/test_profile_api.py index 2dd82e94..9767c2d4 100644 --- a/lambda/tests/entrypoint/test_profile_api.py +++ b/lambda/tests/entrypoint/test_profile_api.py @@ -1,15 +1,19 @@ """Tests for Profile API lambda entrypoint""" import json +from unittest.mock import Mock from src.entrypoint import profile_api_handler +from src.environment.service_provider import ServiceProvider from src.services.api_router import ApiRouter class TestProfileApiAuthErrors: """Tests for authentication error handling""" - def test_missing_auth_header_returns_401(self, mock_service_provider, sample_lambda_context): + def test_missing_auth_header_returns_401( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test request without Authorization header returns 401""" event = { "httpMethod": "GET", @@ -22,7 +26,9 @@ def test_missing_auth_header_returns_401(self, mock_service_provider, sample_lam result = profile_api_handler(event, sample_lambda_context, mock_service_provider) assert result["statusCode"] == 401 - def test_missing_auth_header_error_format(self, mock_service_provider, sample_lambda_context): + def test_missing_auth_header_error_format( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test missing auth header returns proper error body""" event = { "httpMethod": "GET", @@ -40,7 +46,9 @@ def test_missing_auth_header_error_format(self, mock_service_provider, sample_la class TestServiceProviderProfileApiProperties: """Tests for ServiceProvider profile API property initialization""" - def test_profile_api_router_creates_router_with_routes(self, mock_service_provider): + def test_profile_api_router_creates_router_with_routes( + self, mock_service_provider: ServiceProvider + ) -> None: """Test profile_api_router creates ApiRouter with profile route""" router = mock_service_provider.profile_api_router diff --git a/lambda/tests/entrypoint/test_storage_api.py b/lambda/tests/entrypoint/test_storage_api.py index 48353012..ccee2c10 100644 --- a/lambda/tests/entrypoint/test_storage_api.py +++ b/lambda/tests/entrypoint/test_storage_api.py @@ -1,8 +1,11 @@ """Tests for lambda entrypoint""" -from unittest.mock import patch +from unittest.mock import Mock, patch + +from botocore.stub import Stubber from src.entrypoint import storage_api_handler +from src.environment.service_provider import ServiceProvider from src.services.hawk_service import HawkCredentials from src.services.token_generator import TokenGenerator @@ -11,7 +14,11 @@ TEST_UID = str(TokenGenerator.generate_uid(TEST_USER_ID, TEST_GENERATION)) -def test_storage_api_happ_path(mock_service_provider, dynamodb_stubber, sample_lambda_context): +def test_storage_api_happ_path( + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, +) -> None: """ Integration test: storage_api with MockServiceProvider and stubbed DynamoDB. diff --git a/lambda/tests/entrypoint/test_token_api.py b/lambda/tests/entrypoint/test_token_api.py index 5381d6b0..2eddff8c 100644 --- a/lambda/tests/entrypoint/test_token_api.py +++ b/lambda/tests/entrypoint/test_token_api.py @@ -1,16 +1,18 @@ """Tests for Token API lambda entrypoint""" import json +from unittest.mock import Mock import pytest from src.entrypoint import token_api_handler +from src.environment.service_provider import ServiceProvider from src.services.api_router import ApiRouter from src.services.oidc_validator import OIDCValidator @pytest.fixture -def token_request_event(): +def token_request_event() -> dict: """Sample token request event""" return { "httpMethod": "GET", @@ -28,7 +30,9 @@ def token_request_event(): class TestTokenApiAuthErrors: """Tests for authentication error handling - these don't need OIDC validation""" - def test_missing_auth_header_returns_401(self, mock_service_provider, sample_lambda_context): + def test_missing_auth_header_returns_401( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test request without Authorization header returns 401""" event = { "httpMethod": "GET", @@ -41,7 +45,9 @@ def test_missing_auth_header_returns_401(self, mock_service_provider, sample_lam result = token_api_handler(event, sample_lambda_context, mock_service_provider) assert result["statusCode"] == 401 - def test_missing_auth_header_error_format(self, mock_service_provider, sample_lambda_context): + def test_missing_auth_header_error_format( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test missing auth header returns proper error format""" event = { "httpMethod": "GET", @@ -58,7 +64,9 @@ def test_missing_auth_header_error_format(self, mock_service_provider, sample_la assert body["errors"][0]["name"] == "Authorization" assert body["errors"][0]["location"] == "header" - def test_malformed_auth_header_returns_400(self, mock_service_provider, sample_lambda_context): + def test_malformed_auth_header_returns_400( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test request with malformed Authorization header returns 400""" event = { "httpMethod": "GET", @@ -71,7 +79,9 @@ def test_malformed_auth_header_returns_400(self, mock_service_provider, sample_l result = token_api_handler(event, sample_lambda_context, mock_service_provider) assert result["statusCode"] == 400 - def test_malformed_auth_header_error_format(self, mock_service_provider, sample_lambda_context): + def test_malformed_auth_header_error_format( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test malformed auth header returns proper error format""" event = { "httpMethod": "GET", @@ -89,7 +99,9 @@ def test_malformed_auth_header_error_format(self, mock_service_provider, sample_ class TestTokenApiValidationErrors: """Tests for request validation error handling""" - def test_invalid_content_type_returns_415(self, mock_service_provider, sample_lambda_context): + def test_invalid_content_type_returns_415( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test request with invalid Content-Type returns 415""" event = { "httpMethod": "GET", @@ -105,7 +117,9 @@ def test_invalid_content_type_returns_415(self, mock_service_provider, sample_la result = token_api_handler(event, sample_lambda_context, mock_service_provider) assert result["statusCode"] == 415 - def test_invalid_content_type_error_format(self, mock_service_provider, sample_lambda_context): + def test_invalid_content_type_error_format( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """Test invalid content type returns proper error format""" event = { "httpMethod": "GET", @@ -126,7 +140,7 @@ def test_invalid_content_type_error_format(self, mock_service_provider, sample_l class TestServiceProviderTokenApiProperties: """Tests for ServiceProvider token API property initialization""" - def test_oidc_validator_uses_env_vars(self, mock_service_provider): + def test_oidc_validator_uses_env_vars(self, mock_service_provider: ServiceProvider) -> None: """Test oidc_validator is initialized with config from env vars""" validator = mock_service_provider.oidc_validator @@ -134,7 +148,9 @@ def test_oidc_validator_uses_env_vars(self, mock_service_provider): assert validator.provider_url == "https://auth.example.com" assert validator.client_id == "test-client-id" - def test_token_api_router_creates_router_with_routes(self, mock_service_provider): + def test_token_api_router_creates_router_with_routes( + self, mock_service_provider: ServiceProvider + ) -> None: """Test token_api_router creates ApiRouter with token route""" router = mock_service_provider.token_api_router diff --git a/lambda/tests/fixtures/boto.py b/lambda/tests/fixtures/boto.py index b24248b7..a45d3ea8 100644 --- a/lambda/tests/fixtures/boto.py +++ b/lambda/tests/fixtures/boto.py @@ -1,56 +1,66 @@ """AWS service fixtures with botocore stubbing""" -from typing import Generator -from unittest.mock import patch +from typing import TYPE_CHECKING, Any, Generator, cast +from unittest.mock import MagicMock, patch import boto3 import pytest from botocore.stub import Stubber +if TYPE_CHECKING: + from types_boto3_apigatewaymanagementapi.client import ApiGatewayManagementApiClient + from types_boto3_dynamodb.client import DynamoDBClient + from types_boto3_dynamodb.service_resource import DynamoDBServiceResource, Table + from types_boto3_kms.client import KMSClient + @pytest.fixture(scope="session") -def aws_region_name(): +def aws_region_name() -> str: return "us-east-1" @pytest.fixture(scope="session") -def aws_account_id(): +def aws_account_id() -> str: return "00000000000" @pytest.fixture(scope="session") -def aws_access_key_id(): +def aws_access_key_id() -> str: return "fake-access-key-id" @pytest.fixture(scope="session") -def aws_secret_access_key(): +def aws_secret_access_key() -> str: return "fake-secret-access-key" @pytest.fixture(scope="session") -def aws_session_token(): +def aws_session_token() -> str: return "fake-session-token" @pytest.fixture -def dynamodb_client(boto_session): +def dynamodb_client(boto_session: boto3.session.Session) -> DynamoDBClient: return boto_session.client("dynamodb") @pytest.fixture -def dynamodb_resource(boto_session): +def dynamodb_resource(boto_session: boto3.session.Session) -> "DynamoDBServiceResource": return boto_session.resource("dynamodb") @pytest.fixture -def dynamodb_stubber(dynamodb_resource): +def dynamodb_stubber( + dynamodb_resource: "DynamoDBServiceResource", +) -> Generator[Stubber, None, None]: with Stubber(dynamodb_resource.meta.client) as stubber: yield stubber @pytest.fixture -def dynamodb_table(boto_session, dynamodb_stubber, storage_table_name): +def dynamodb_table( + boto_session: boto3.session.Session, dynamodb_stubber: Stubber, storage_table_name: str +) -> "Table": """ Provides a DynamoDB Table resource with stubbed client. @@ -77,28 +87,26 @@ def test_something(dynamodb_table, dynamodb_stubber): table = resource.Table(storage_table_name) # Replace the Table's internal client with the stubbed one - table.meta.client = dynamodb_stubber.client + table.meta.client = cast("DynamoDBClient", dynamodb_stubber.client) return table @pytest.fixture -def kms_client(boto_session): +def kms_client(boto_session: boto3.session.Session) -> KMSClient: """KMS client from the test boto session.""" return boto_session.client("kms") @pytest.fixture -def kms_stubber(kms_client): +def kms_stubber(kms_client: KMSClient) -> Generator[Stubber, None, None]: """Botocore Stubber for KMS. Tests that call KMS add their own stubs.""" - stubber = Stubber(kms_client) - stubber.activate() - yield stubber - stubber.deactivate() + with Stubber(kms_client) as stubber: + yield stubber @pytest.fixture -def apigw_client(boto_session): +def apigw_client(boto_session: boto3.session.Session) -> ApiGatewayManagementApiClient: """API Gateway Management API client for WebSocket connection posting.""" return boto_session.client( "apigatewaymanagementapi", @@ -107,16 +115,19 @@ def apigw_client(boto_session): @pytest.fixture -def apigw_stubber(apigw_client): +def apigw_stubber(apigw_client: ApiGatewayManagementApiClient) -> Generator[Stubber, None, None]: """Botocore Stubber for API Gateway Management API.""" - stubber = Stubber(apigw_client) - stubber.activate() - yield stubber - stubber.deactivate() + with Stubber(apigw_client) as stubber: + yield stubber @pytest.fixture(autouse=True) -def boto_session(aws_region_name, aws_access_key_id, aws_secret_access_key, aws_session_token): +def boto_session( + aws_region_name: str, + aws_access_key_id: str, + aws_secret_access_key: str, + aws_session_token: str, +) -> boto3.session.Session: return boto3.session.Session( aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, @@ -126,7 +137,7 @@ def boto_session(aws_region_name, aws_access_key_id, aws_secret_access_key, aws_ @pytest.fixture -def boto_session_patch(boto_session): +def boto_session_patch(boto_session: boto3.session.Session) -> Generator[MagicMock, None, None]: with ( patch("boto3.Session", autospec=True) as m, patch("boto3.session.Session", autospec=True) as m2, @@ -138,9 +149,14 @@ def boto_session_patch(boto_session): @pytest.fixture(autouse=True) def boto_resource_patch( - boto_session, boto_session_patch, dynamodb_client, dynamodb_resource, kms_client, apigw_client + boto_session: boto3.session.Session, + boto_session_patch: MagicMock, + dynamodb_client: DynamoDBClient, + dynamodb_resource: "DynamoDBServiceResource", + kms_client: KMSClient, + apigw_client: ApiGatewayManagementApiClient, ) -> Generator: - def client(service, *args, **kwargs): + def client(service: str, *args: Any, **kwargs: Any) -> Any: if service == "dynamodb": return dynamodb_client if service == "kms": @@ -150,7 +166,7 @@ def client(service, *args, **kwargs): raise ValueError(f"client for {service} not recognized") - def resource(service, *args, **kwargs): + def resource(service: str, *args: Any, **kwargs: Any) -> Any: if service == "dynamodb": return dynamodb_resource diff --git a/lambda/tests/fixtures/integration.py b/lambda/tests/fixtures/integration.py index 9a32092a..e05bf79b 100644 --- a/lambda/tests/fixtures/integration.py +++ b/lambda/tests/fixtures/integration.py @@ -2,10 +2,11 @@ import json import time -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Union import mohawk import pytest +from botocore.stub import Stubber from src.services.token_generator import TokenGenerator @@ -14,7 +15,15 @@ # ============================================================================ -def build_hawk_auth_header(hawk_id, hawk_key, method, path, host, port, **kwargs): +def build_hawk_auth_header( + hawk_id: str, + hawk_key: Union[str, bytes], + method: str, + path: str, + host: str, + port: Union[str, int], + **kwargs: Any, +) -> str: """Build a Hawk Authorization header using mohawk.Sender. Args: @@ -42,7 +51,7 @@ def build_hawk_auth_header(hawk_id, hawk_key, method, path, host, port, **kwargs @pytest.fixture -def valid_hawk_credentials(): +def valid_hawk_credentials() -> Dict[str, Any]: """Valid HAWK credentials for integration tests""" return { "user_id": "test-user-123", @@ -54,7 +63,7 @@ def valid_hawk_credentials(): @pytest.fixture -def expired_hawk_credentials(): +def expired_hawk_credentials() -> Dict[str, Any]: """Expired HAWK credentials for testing authentication failures""" return { "user_id": "test-user-123", @@ -66,7 +75,7 @@ def expired_hawk_credentials(): @pytest.fixture -def hawk_authorization_header(valid_hawk_credentials): +def hawk_authorization_header(valid_hawk_credentials: Dict[str, Any]) -> str: """Generate a valid HAWK Authorization header""" timestamp = int(time.time()) nonce = "test-nonce-123" @@ -86,7 +95,7 @@ def hawk_authorization_header(valid_hawk_credentials): @pytest.fixture -def valid_bso_data(): +def valid_bso_data() -> Dict[str, Any]: """Valid BSO data for creation/update""" return { "id": "test-bso-001", @@ -97,7 +106,7 @@ def valid_bso_data(): @pytest.fixture -def batch_bso_data(): +def batch_bso_data() -> List[Dict[str, Any]]: """Batch of valid BSOs for batch operations""" return [ { @@ -110,7 +119,7 @@ def batch_bso_data(): @pytest.fixture -def large_batch_bso_data(): +def large_batch_bso_data() -> List[Dict[str, Any]]: """Large batch of BSOs for testing limits (100 items)""" return [ { @@ -122,7 +131,7 @@ def large_batch_bso_data(): @pytest.fixture -def oversized_bso_data(): +def oversized_bso_data() -> Dict[str, Any]: """BSO with payload exceeding max size (256 KB)""" return { "id": "oversized-bso", @@ -137,7 +146,7 @@ def oversized_bso_data(): @pytest.fixture -def valid_collection_names(): +def valid_collection_names() -> List[str]: """List of valid collection names""" return [ "bookmarks", @@ -154,7 +163,7 @@ def valid_collection_names(): @pytest.fixture -def invalid_collection_names(): +def invalid_collection_names() -> List[str]: """List of invalid collection names for validation testing""" return [ "a" * 33, # Too long (>32 chars) @@ -268,14 +277,14 @@ def build_authorizer_event( def stub_get_bso( - stubber, + stubber: Stubber, table_name: str, user_id: str, collection_name: str, object_id: str, bso_data: Optional[Dict[str, Any]] = None, exists: bool = True, -): +) -> None: """ Stub a DynamoDB get_item call for retrieving a BSO. @@ -316,12 +325,12 @@ def stub_get_bso( def stub_put_bso( - stubber, + stubber: Stubber, table_name: str, user_id: str, collection_name: str, object_id: str, -): +) -> None: """ Stub a DynamoDB put_item call for creating/updating a BSO. @@ -345,12 +354,12 @@ def stub_put_bso( def stub_query_collection( - stubber, + stubber: Stubber, table_name: str, user_id: str, collection_name: str, items: List[Dict[str, Any]], -): +) -> None: """ Stub a DynamoDB query call for listing BSOs in a collection. @@ -391,7 +400,7 @@ def stub_query_collection( # ============================================================================ -def assert_successful_response(response: Dict[str, Any], expected_status: int = 200): +def assert_successful_response(response: Dict[str, Any], expected_status: int = 200) -> None: """Assert that a Lambda response is successful""" assert response["statusCode"] == expected_status assert "headers" in response @@ -402,7 +411,7 @@ def assert_error_response( response: Dict[str, Any], expected_status: int, expected_code: Optional[int] = None, -): +) -> None: """ Assert that a Lambda response is an error. @@ -418,7 +427,7 @@ def assert_error_response( assert body == expected_code -def assert_bso_response(response: Dict[str, Any], expected_bso: Dict[str, Any]): +def assert_bso_response(response: Dict[str, Any], expected_bso: Dict[str, Any]) -> None: """Assert that a response contains the expected BSO""" assert_successful_response(response) @@ -439,7 +448,7 @@ def assert_collection_response( response: Dict[str, Any], expected_ids: List[str], full: bool = False, -): +) -> None: """ Assert that a response contains the expected collection data. @@ -471,7 +480,7 @@ def assert_batch_response( response: Dict[str, Any], expected_success: List[str], expected_failed: Optional[Dict[str, str]] = None, -): +) -> None: """Assert that a batch operation response is correct""" assert_successful_response(response) @@ -494,7 +503,7 @@ def assert_batch_response( @pytest.fixture -def multi_user_test_data(): +def multi_user_test_data() -> Dict[str, Any]: """Test data for multiple users to verify isolation""" return { "user1": { @@ -528,7 +537,7 @@ def multi_user_test_data(): # ============================================================================ -def assert_timestamp_format(timestamp_str: str): +def assert_timestamp_format(timestamp_str: str) -> None: """Assert that a timestamp string has the correct format (2 decimal places)""" parts = timestamp_str.split(".") assert len(parts) == 2, "Timestamp must have decimal point" @@ -539,7 +548,7 @@ def assert_timestamp_format(timestamp_str: str): assert timestamp > 0, "Timestamp must be positive" -def assert_timestamp_headers(response: Dict[str, Any], is_write: bool = False): +def assert_timestamp_headers(response: Dict[str, Any], is_write: bool = False) -> None: """ Assert that timestamp headers are present and correct. diff --git a/lambda/tests/integration/test_e2e_flow.py b/lambda/tests/integration/test_e2e_flow.py index 648b1958..1fcce652 100644 --- a/lambda/tests/integration/test_e2e_flow.py +++ b/lambda/tests/integration/test_e2e_flow.py @@ -12,11 +12,14 @@ import json import time +from typing import Generator +from unittest.mock import Mock import pytest -from botocore.stub import ANY +from botocore.stub import ANY, Stubber from src.entrypoint.storage_api import lambda_handler as storage_handler +from src.environment.service_provider import ServiceProvider from src.services.hawk_service import HawkCredentials from src.services.token_generator import TokenGenerator from tests.fixtures.integration import ( @@ -29,8 +32,11 @@ class TestHawkMiddlewareToStorageAPIFlow: """Test end-to-end flow from HAWK middleware to Storage API""" def test_successful_authentication_flow( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test successful HAWK authentication via middleware and Storage API access. @@ -108,7 +114,9 @@ def test_successful_authentication_flow( body = json.loads(storage_response["body"]) assert body == [] - def test_authentication_failure_returns_401(self, mock_service_provider, sample_lambda_context): + def test_authentication_failure_returns_401( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """ Test that invalid HAWK credentials return 401 from middleware. @@ -133,8 +141,8 @@ def test_authentication_failure_returns_401(self, mock_service_provider, sample_ assert storage_response["statusCode"] == 401 def test_missing_authorization_header_returns_401( - self, mock_service_provider, sample_lambda_context - ): + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """ Test that requests without Authorization header are rejected by middleware. """ @@ -157,7 +165,9 @@ def test_missing_authorization_header_returns_401( class TestUidMismatch: """Test that UID mismatch is rejected by middleware""" - def test_uid_mismatch_returns_403(self, mock_service_provider, sample_lambda_context): + def test_uid_mismatch_returns_403( + self, mock_service_provider: ServiceProvider, sample_lambda_context: Mock + ) -> None: """ Test that a request where the URL uid does not match the authenticated user's expected uid returns 403 via the UidMismatchError exception handler. @@ -170,7 +180,7 @@ def test_uid_mismatch_returns_403(self, mock_service_provider, sample_lambda_con expiry=9999999999, hawk_id="test-hawk-id", ) - mock_service_provider.hawk_service.validate = lambda *a, **kw: creds + mock_service_provider.hawk_service.validate = lambda *a, **kw: creds # type: ignore[method-assign] # Build event with WRONG uid in URL path (doesn't match user_id+generation) storage_event = build_storage_event( @@ -193,12 +203,16 @@ class TestUserIsolation: """Test that users can only access their own data""" @pytest.fixture(autouse=True) - def mock_hawk_validate(self, mock_service_provider): + def mock_hawk_validate( + self, mock_service_provider: ServiceProvider + ) -> Generator[None, None, None]: """Mock hawk_service.validate to bypass auth for user isolation tests.""" self._mock_validate = mock_service_provider.hawk_service.validate yield - def _set_hawk_user(self, mock_service_provider, user_id, generation=0): + def _set_hawk_user( + self, mock_service_provider: ServiceProvider, user_id: str, generation: int = 0 + ) -> None: """Configure hawk_service.validate to return credentials for given user.""" creds = HawkCredentials( user_id=user_id, @@ -206,11 +220,14 @@ def _set_hawk_user(self, mock_service_provider, user_id, generation=0): expiry=9999999999, hawk_id="test-hawk-id", ) - mock_service_provider.hawk_service.validate = lambda *a, **kw: creds + mock_service_provider.hawk_service.validate = lambda *a, **kw: creds # type: ignore[method-assign] def test_user_can_read_own_collection( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that a user can successfully list their own collection. """ @@ -248,8 +265,11 @@ def test_user_can_read_own_collection( assert body == ["bso-001"] def test_different_users_query_different_namespaces( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that two different users querying the same collection name actually query different DynamoDB partition keys. @@ -323,8 +343,11 @@ def test_different_users_query_different_namespaces( assert user1_body != user2_body def test_user_cannot_access_missing_bso_in_other_namespace( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that when User2 tries to access a BSO ID that exists in User1's namespace, they get 404. @@ -346,8 +369,11 @@ def test_user_cannot_access_missing_bso_in_other_namespace( assert response["statusCode"] == 404 def test_info_collections_scoped_to_user( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that /info/collections returns only the authenticated user's collections. """ @@ -394,8 +420,11 @@ def test_info_collections_scoped_to_user( assert len(body) == 2 def test_delete_all_scoped_to_user( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that DELETE /storage only deletes the authenticated user's data. """ diff --git a/lambda/tests/integration/test_token_server_mozilla_spec.py b/lambda/tests/integration/test_token_server_mozilla_spec.py index 5a5f2122..c9a0e434 100644 --- a/lambda/tests/integration/test_token_server_mozilla_spec.py +++ b/lambda/tests/integration/test_token_server_mozilla_spec.py @@ -13,11 +13,13 @@ import json import time -from unittest.mock import patch +from unittest.mock import Mock, patch -from botocore.stub import ANY +import pytest +from botocore.stub import ANY, Stubber from src.entrypoint.token_api import lambda_handler as token_handler +from src.environment.service_provider import ServiceProvider class TestGetMethodTokenIssuance: @@ -25,10 +27,10 @@ class TestGetMethodTokenIssuance: def test_get_method_token_issuance_complete_flow( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test complete GET method token issuance flow. @@ -136,10 +138,10 @@ def test_get_method_token_issuance_complete_flow( def test_get_method_response_structure_matches_mozilla_spec( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that response structure exactly matches Mozilla Token Server API v1.0 spec. @@ -214,10 +216,10 @@ class TestClientStateHistory: def test_client_state_change_flow( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test client state change increments generation and updates history. @@ -376,10 +378,10 @@ def test_client_state_change_flow( def test_rejection_of_previously_seen_client_state( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that previously-seen client state is rejected. @@ -462,10 +464,10 @@ def test_rejection_of_previously_seen_client_state( def test_rejection_of_empty_state_when_history_exists( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that empty client state is rejected when history contains non-empty values. @@ -552,10 +554,10 @@ class TestNewErrorStatuses: def test_invalid_timestamp_response( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test invalid-timestamp error status. @@ -597,10 +599,10 @@ def test_invalid_timestamp_response( def test_invalid_generation_response( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test invalid-generation error status. @@ -620,10 +622,10 @@ def test_invalid_generation_response( def test_invalid_client_state_response( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test invalid-client-state error status. @@ -640,11 +642,11 @@ def test_invalid_client_state_response( def test_new_users_disabled_response( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - monkeypatch, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: """ Test new-users-disabled error status. @@ -710,10 +712,10 @@ class TestResponseHeaders: def test_x_timestamp_on_200_response( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test X-Timestamp header on successful 200 response. @@ -791,10 +793,10 @@ def test_x_timestamp_on_200_response( def test_x_timestamp_on_401_response( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test X-Timestamp header on 401 error response. @@ -829,10 +831,10 @@ def test_x_timestamp_on_401_response( def test_www_authenticate_on_401_response( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test WWW-Authenticate header on 401 response. @@ -865,10 +867,10 @@ def test_www_authenticate_on_401_response( def test_all_headers_on_401_response( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that 401 responses include both X-Timestamp and WWW-Authenticate. @@ -902,10 +904,10 @@ class TestNodeReset: def test_uid_changes_when_client_state_changes( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that uid changes when client state changes (node reset). @@ -1068,10 +1070,10 @@ def test_uid_changes_when_client_state_changes( def test_api_endpoint_changes_when_client_state_changes( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that api_endpoint changes when client state changes. diff --git a/lambda/tests/integration/test_token_to_storage_flow.py b/lambda/tests/integration/test_token_to_storage_flow.py index 6f32432d..0ba74fc6 100644 --- a/lambda/tests/integration/test_token_to_storage_flow.py +++ b/lambda/tests/integration/test_token_to_storage_flow.py @@ -12,12 +12,13 @@ import json import time -from unittest.mock import patch +from unittest.mock import Mock, patch -from botocore.stub import ANY +from botocore.stub import ANY, Stubber from src.entrypoint.storage_api import lambda_handler as storage_handler from src.entrypoint.token_api import lambda_handler as token_handler +from src.environment.service_provider import ServiceProvider from src.services.token_generator import TokenGenerator from tests.fixtures.integration import ( build_hawk_auth_header, @@ -30,10 +31,10 @@ class TestTokenServerToStorageServerFlow: def test_complete_token_issuance_and_validation_flow( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test complete flow: Token issuance -> HAWK middleware auth -> Storage access. @@ -194,10 +195,10 @@ def test_complete_token_issuance_and_validation_flow( def test_token_server_stores_credentials_for_middleware_validation( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that Token Server stores HAWK credentials in cache for middleware validation. @@ -317,10 +318,10 @@ def test_token_server_stores_credentials_for_middleware_validation( def test_expired_token_rejected_by_middleware( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that expired HAWK tokens are rejected by StorageHawkMiddleware. @@ -359,10 +360,10 @@ def test_expired_token_rejected_by_middleware( def test_invalid_hawk_signature_rejected_by_middleware( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that invalid HAWK signatures are rejected by StorageHawkMiddleware. @@ -416,10 +417,10 @@ def test_invalid_hawk_signature_rejected_by_middleware( def test_generation_mismatch_rejected_by_middleware( self, - mock_service_provider, - dynamodb_stubber, - sample_lambda_context, - ): + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """ Test that tokens with mismatched generation numbers are rejected. diff --git a/lambda/tests/routes/auth/conftest.py b/lambda/tests/routes/auth/conftest.py new file mode 100644 index 00000000..3fc46be6 --- /dev/null +++ b/lambda/tests/routes/auth/conftest.py @@ -0,0 +1,29 @@ +"""Shared route-test doubles for the FxA auth routes. + +Deliberately bare MagicMocks: a test that needs configured behaviour should configure it +locally, so that adding setup here can never silently change what another test receives. +""" + +from unittest.mock import MagicMock + +import pytest + + +@pytest.fixture +def mock_account_manager() -> MagicMock: + return MagicMock() + + +@pytest.fixture +def mock_token_manager() -> MagicMock: + return MagicMock() + + +@pytest.fixture +def device_manager() -> MagicMock: + return MagicMock() + + +@pytest.fixture +def mock_oauth_code_manager() -> MagicMock: + return MagicMock() diff --git a/lambda/tests/routes/auth/test_account_attached_clients.py b/lambda/tests/routes/auth/test_account_attached_clients.py index 3db9ea28..2fb5bbda 100644 --- a/lambda/tests/routes/auth/test_account_attached_clients.py +++ b/lambda/tests/routes/auth/test_account_attached_clients.py @@ -1,26 +1,23 @@ """Unit tests for AccountAttachedClients route""" -import json from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.account_attached_clients import AccountAttachedClientsRoute +from tests.conftest import json_body @pytest.fixture -def device_manager(): - return MagicMock() - - -@pytest.fixture -def route(device_manager): +def route(device_manager: MagicMock) -> AccountAttachedClientsRoute: return AccountAttachedClientsRoute(device_manager=device_manager, middlewares=[]) class TestAccountAttachedClients: - def test_returns_attached_clients(self, route, device_manager): + def test_returns_attached_clients( + self, route: AccountAttachedClientsRoute, device_manager: MagicMock + ) -> None: device_manager.get_devices.return_value = [ { "id": "dev1", @@ -43,7 +40,7 @@ def test_returns_attached_clients(self, route, device_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert len(body) == 1 client = body[0] assert client["clientId"] is None @@ -59,7 +56,9 @@ def test_returns_attached_clients(self, route, device_manager): assert client["userAgent"] == "" assert client["os"] is None - def test_is_current_session_set(self, route, device_manager): + def test_is_current_session_set( + self, route: AccountAttachedClientsRoute, device_manager: MagicMock + ) -> None: device_manager.get_devices.return_value = [ { "id": "dev1", @@ -89,11 +88,13 @@ def test_is_current_session_set(self, route, device_manager): } ) response = route.handle(event) - body = json.loads(response.body) + body = json_body(response) assert body[0]["isCurrentSession"] is True assert body[1]["isCurrentSession"] is False - def test_returns_empty_list(self, route, device_manager): + def test_returns_empty_list( + self, route: AccountAttachedClientsRoute, device_manager: MagicMock + ) -> None: device_manager.get_devices.return_value = [] event = APIGatewayProxyEvent( { @@ -107,12 +108,12 @@ def test_returns_empty_list(self, route, device_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body == [] class TestAccountAttachedClientsBind: - def test_bind_registers_get_route(self, route): + def test_bind_registers_get_route(self, route: AccountAttachedClientsRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_account_create.py b/lambda/tests/routes/auth/test_account_create.py index a15b52f2..8faeee9a 100644 --- a/lambda/tests/routes/auth/test_account_create.py +++ b/lambda/tests/routes/auth/test_account_create.py @@ -8,25 +8,18 @@ from src.routes.auth.account_create import AccountCreateRoute from src.shared.oidc import OIDCTokenClaims +from tests.conftest import json_body @pytest.fixture -def mock_account_manager(): +def mock_oidc_validator() -> MagicMock: return MagicMock() @pytest.fixture -def mock_token_manager(): - return MagicMock() - - -@pytest.fixture -def mock_oidc_validator(): - return MagicMock() - - -@pytest.fixture -def route(mock_account_manager, mock_token_manager, mock_oidc_validator): +def route( + mock_account_manager: MagicMock, mock_token_manager: MagicMock, mock_oidc_validator: MagicMock +) -> AccountCreateRoute: return AccountCreateRoute( account_manager=mock_account_manager, token_manager=mock_token_manager, @@ -35,7 +28,7 @@ def route(mock_account_manager, mock_token_manager, mock_oidc_validator): @pytest.fixture -def valid_claims(): +def valid_claims() -> OIDCTokenClaims: return OIDCTokenClaims( sub="oidc-sub-123", iss="https://auth.example.com", @@ -48,8 +41,13 @@ def valid_claims(): class TestAccountCreate: def test_success_returns_uid_and_tokens( - self, route, mock_account_manager, mock_token_manager, mock_oidc_validator, valid_claims - ): + self, + route: AccountCreateRoute, + mock_account_manager: MagicMock, + mock_token_manager: MagicMock, + mock_oidc_validator: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims mock_token_manager.create_session_token.return_value = b"\xaa" * 32 mock_token_manager.create_key_fetch_token.return_value = b"\xbb" * 32 @@ -64,13 +62,13 @@ def test_success_returns_uid_and_tokens( ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert "uid" in body assert body["sessionToken"] == "aa" * 32 assert body["keyFetchToken"] == "bb" * 32 assert body["verified"] is True - def test_missing_auth_header_returns_401(self, route): + def test_missing_auth_header_returns_401(self, route: AccountCreateRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -82,7 +80,9 @@ def test_missing_auth_header_returns_401(self, route): response = route.handle(event) assert response.status_code == 401 - def test_invalid_oidc_token_returns_401(self, route, mock_oidc_validator): + def test_invalid_oidc_token_returns_401( + self, route: AccountCreateRoute, mock_oidc_validator: MagicMock + ) -> None: mock_oidc_validator.validate_token.side_effect = Exception("Invalid token") event = APIGatewayProxyEvent( { @@ -96,8 +96,12 @@ def test_invalid_oidc_token_returns_401(self, route, mock_oidc_validator): assert response.status_code == 401 def test_duplicate_email_returns_409( - self, route, mock_account_manager, mock_oidc_validator, valid_claims - ): + self, + route: AccountCreateRoute, + mock_account_manager: MagicMock, + mock_oidc_validator: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims mock_account_manager.create_account.side_effect = ValueError("Email already exists") event = APIGatewayProxyEvent( @@ -111,7 +115,12 @@ def test_duplicate_email_returns_409( response = route.handle(event) assert response.status_code == 409 - def test_missing_email_returns_400(self, route, mock_oidc_validator, valid_claims): + def test_missing_email_returns_400( + self, + route: AccountCreateRoute, + mock_oidc_validator: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims event = APIGatewayProxyEvent( { @@ -124,7 +133,12 @@ def test_missing_email_returns_400(self, route, mock_oidc_validator, valid_claim response = route.handle(event) assert response.status_code == 400 - def test_missing_authpw_returns_400(self, route, mock_oidc_validator, valid_claims): + def test_missing_authpw_returns_400( + self, + route: AccountCreateRoute, + mock_oidc_validator: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims event = APIGatewayProxyEvent( { @@ -137,7 +151,7 @@ def test_missing_authpw_returns_400(self, route, mock_oidc_validator, valid_clai response = route.handle(event) assert response.status_code == 400 - def test_malformed_auth_header_returns_401(self, route): + def test_malformed_auth_header_returns_401(self, route: AccountCreateRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -149,7 +163,12 @@ def test_malformed_auth_header_returns_401(self, route): response = route.handle(event) assert response.status_code == 401 - def test_invalid_json_body_returns_400(self, route, mock_oidc_validator, valid_claims): + def test_invalid_json_body_returns_400( + self, + route: AccountCreateRoute, + mock_oidc_validator: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims event = APIGatewayProxyEvent( { @@ -162,7 +181,12 @@ def test_invalid_json_body_returns_400(self, route, mock_oidc_validator, valid_c response = route.handle(event) assert response.status_code == 400 - def test_missing_body_returns_400(self, route, mock_oidc_validator, valid_claims): + def test_missing_body_returns_400( + self, + route: AccountCreateRoute, + mock_oidc_validator: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims event = APIGatewayProxyEvent( { @@ -176,8 +200,13 @@ def test_missing_body_returns_400(self, route, mock_oidc_validator, valid_claims assert response.status_code == 400 def test_creates_account_with_correct_params( - self, route, mock_account_manager, mock_oidc_validator, mock_token_manager, valid_claims - ): + self, + route: AccountCreateRoute, + mock_account_manager: MagicMock, + mock_oidc_validator: MagicMock, + mock_token_manager: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims mock_token_manager.create_session_token.return_value = b"\xaa" * 32 mock_token_manager.create_key_fetch_token.return_value = b"\xbb" * 32 @@ -200,7 +229,12 @@ def test_creates_account_with_correct_params( assert len(call_kwargs["wrap_kb"]) == 64 assert len(call_kwargs["key_rotation_secret"]) == 64 - def test_invalid_authpw_format_returns_400(self, route, mock_oidc_validator, valid_claims): + def test_invalid_authpw_format_returns_400( + self, + route: AccountCreateRoute, + mock_oidc_validator: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims event = APIGatewayProxyEvent( { @@ -212,10 +246,15 @@ def test_invalid_authpw_format_returns_400(self, route, mock_oidc_validator, val ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert "authPW" in body["message"] - def test_invalid_email_format_returns_400(self, route, mock_oidc_validator, valid_claims): + def test_invalid_email_format_returns_400( + self, + route: AccountCreateRoute, + mock_oidc_validator: MagicMock, + valid_claims: OIDCTokenClaims, + ) -> None: mock_oidc_validator.validate_token.return_value = valid_claims event = APIGatewayProxyEvent( { @@ -227,12 +266,12 @@ def test_invalid_email_format_returns_400(self, route, mock_oidc_validator, vali ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert "email" in body["message"] class TestAccountCreateBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: AccountCreateRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_account_device.py b/lambda/tests/routes/auth/test_account_device.py index f727da5c..42a738d9 100644 --- a/lambda/tests/routes/auth/test_account_device.py +++ b/lambda/tests/routes/auth/test_account_device.py @@ -7,20 +7,18 @@ from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.account_device import AccountDeviceRoute +from tests.conftest import json_body @pytest.fixture -def device_manager(): - return MagicMock() - - -@pytest.fixture -def route(device_manager): +def route(device_manager: MagicMock) -> AccountDeviceRoute: return AccountDeviceRoute(device_manager=device_manager, middlewares=[]) class TestAccountDevice: - def test_create_device_returns_200(self, route, device_manager): + def test_create_device_returns_200( + self, route: AccountDeviceRoute, device_manager: MagicMock + ) -> None: device_manager.upsert_device.return_value = { "id": "dev1", "name": "My Firefox", @@ -40,7 +38,7 @@ def test_create_device_returns_200(self, route, device_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["id"] == "dev1" assert body["name"] == "My Firefox" device_manager.upsert_device.assert_called_once_with( @@ -49,7 +47,9 @@ def test_create_device_returns_200(self, route, device_manager): {"name": "My Firefox", "type": "desktop"}, ) - def test_update_device_returns_200(self, route, device_manager): + def test_update_device_returns_200( + self, route: AccountDeviceRoute, device_manager: MagicMock + ) -> None: device_manager.upsert_device.return_value = { "id": "existing-dev", "name": "Updated Name", @@ -71,7 +71,7 @@ def test_update_device_returns_200(self, route, device_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["id"] == "existing-dev" assert body["name"] == "Updated Name" device_manager.upsert_device.assert_called_once_with( @@ -80,7 +80,9 @@ def test_update_device_returns_200(self, route, device_manager): {"id": "existing-dev", "name": "Updated Name", "type": "mobile"}, ) - def test_missing_body_returns_200(self, route, device_manager): + def test_missing_body_returns_200( + self, route: AccountDeviceRoute, device_manager: MagicMock + ) -> None: device_manager.upsert_device.return_value = { "id": "auto-id", "name": "", @@ -100,13 +102,13 @@ def test_missing_body_returns_200(self, route, device_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["id"] == "auto-id" device_manager.upsert_device.assert_called_once_with("uid1", "token123", {}) class TestAccountDeviceBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: AccountDeviceRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_account_devices.py b/lambda/tests/routes/auth/test_account_devices.py index da56497f..5df927e4 100644 --- a/lambda/tests/routes/auth/test_account_devices.py +++ b/lambda/tests/routes/auth/test_account_devices.py @@ -1,26 +1,23 @@ """Unit tests for AccountDevices route""" -import json from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.account_devices import AccountDevicesRoute +from tests.conftest import json_body @pytest.fixture -def device_manager(): - return MagicMock() - - -@pytest.fixture -def route(device_manager): +def route(device_manager: MagicMock) -> AccountDevicesRoute: return AccountDevicesRoute(device_manager=device_manager, middlewares=[]) class TestAccountDevices: - def test_returns_device_list(self, route, device_manager): + def test_returns_device_list( + self, route: AccountDevicesRoute, device_manager: MagicMock + ) -> None: device_manager.get_devices.return_value = [ { "id": "dev1", @@ -51,13 +48,15 @@ def test_returns_device_list(self, route, device_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert len(body) == 2 assert body[0]["isCurrentDevice"] is True assert body[1]["isCurrentDevice"] is False device_manager.get_devices.assert_called_once_with("uid1", None) - def test_filters_idle_devices(self, route, device_manager): + def test_filters_idle_devices( + self, route: AccountDevicesRoute, device_manager: MagicMock + ) -> None: device_manager.get_devices.return_value = [] event = APIGatewayProxyEvent( { @@ -73,7 +72,9 @@ def test_filters_idle_devices(self, route, device_manager): assert response.status_code == 200 device_manager.get_devices.assert_called_once_with("uid1", 1609459200000) - def test_returns_empty_list(self, route, device_manager): + def test_returns_empty_list( + self, route: AccountDevicesRoute, device_manager: MagicMock + ) -> None: device_manager.get_devices.return_value = [] event = APIGatewayProxyEvent( { @@ -87,12 +88,12 @@ def test_returns_empty_list(self, route, device_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body == [] class TestAccountDevicesBind: - def test_bind_registers_get_route(self, route): + def test_bind_registers_get_route(self, route: AccountDevicesRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_account_devices_notify.py b/lambda/tests/routes/auth/test_account_devices_notify.py index 9182c9fd..536cc69d 100644 --- a/lambda/tests/routes/auth/test_account_devices_notify.py +++ b/lambda/tests/routes/auth/test_account_devices_notify.py @@ -7,15 +7,16 @@ from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.account_devices_notify import AccountDevicesNotifyRoute +from tests.conftest import json_body @pytest.fixture -def route(): +def route() -> AccountDevicesNotifyRoute: return AccountDevicesNotifyRoute(middlewares=[]) class TestAccountDevicesNotify: - def test_returns_empty_object(self, route): + def test_returns_empty_object(self, route: AccountDevicesNotifyRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -27,12 +28,12 @@ def test_returns_empty_object(self, route): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body == {} class TestAccountDevicesNotifyBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: AccountDevicesNotifyRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_account_keys.py b/lambda/tests/routes/auth/test_account_keys.py index ab5ca608..23913d3d 100644 --- a/lambda/tests/routes/auth/test_account_keys.py +++ b/lambda/tests/routes/auth/test_account_keys.py @@ -1,26 +1,16 @@ """Unit tests for AccountKeys route""" -import json from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.account_keys import AccountKeysRoute +from tests.conftest import json_body @pytest.fixture -def mock_account_manager(): - return MagicMock() - - -@pytest.fixture -def mock_token_manager(): - return MagicMock() - - -@pytest.fixture -def route(mock_account_manager, mock_token_manager): +def route(mock_account_manager: MagicMock, mock_token_manager: MagicMock) -> AccountKeysRoute: return AccountKeysRoute( account_manager=mock_account_manager, token_manager=mock_token_manager, @@ -28,7 +18,12 @@ def route(mock_account_manager, mock_token_manager): class TestAccountKeys: - def test_success_returns_bundle(self, route, mock_account_manager, mock_token_manager): + def test_success_returns_bundle( + self, + route: AccountKeysRoute, + mock_account_manager: MagicMock, + mock_token_manager: MagicMock, + ) -> None: token_id_hex = "aa" * 32 raw_token_hex = "bb" * 32 mock_token_manager.verify_keyfetch_hawk.return_value = { @@ -51,12 +46,14 @@ def test_success_returns_bundle(self, route, mock_account_manager, mock_token_ma ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert "bundle" in body # bundle = 64 bytes ciphertext + 32 bytes HMAC = 96 bytes = 192 hex chars assert len(body["bundle"]) == 192 - def test_invalid_token_returns_401(self, route, mock_token_manager): + def test_invalid_token_returns_401( + self, route: AccountKeysRoute, mock_token_manager: MagicMock + ) -> None: mock_token_manager.verify_keyfetch_hawk.return_value = None event = APIGatewayProxyEvent( { @@ -69,7 +66,7 @@ def test_invalid_token_returns_401(self, route, mock_token_manager): response = route.handle(event) assert response.status_code == 401 - def test_missing_auth_header_returns_401(self, route): + def test_missing_auth_header_returns_401(self, route: AccountKeysRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -81,7 +78,12 @@ def test_missing_auth_header_returns_401(self, route): response = route.handle(event) assert response.status_code == 401 - def test_account_not_found_returns_401(self, route, mock_token_manager, mock_account_manager): + def test_account_not_found_returns_401( + self, + route: AccountKeysRoute, + mock_token_manager: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_token_manager.verify_keyfetch_hawk.return_value = { "uid": "uid1", "keyFetchToken": "bb" * 32, @@ -98,7 +100,9 @@ def test_account_not_found_returns_401(self, route, mock_token_manager, mock_acc response = route.handle(event) assert response.status_code == 401 - def test_consumed_token_second_request_returns_401(self, route, mock_token_manager): + def test_consumed_token_second_request_returns_401( + self, route: AccountKeysRoute, mock_token_manager: MagicMock + ) -> None: mock_token_manager.verify_keyfetch_hawk.return_value = None event = APIGatewayProxyEvent( { @@ -113,7 +117,7 @@ def test_consumed_token_second_request_returns_401(self, route, mock_token_manag class TestAccountKeysBind: - def test_bind_registers_get_route(self, route): + def test_bind_registers_get_route(self, route: AccountKeysRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_account_login.py b/lambda/tests/routes/auth/test_account_login.py index e3eb9fb6..a5f54997 100644 --- a/lambda/tests/routes/auth/test_account_login.py +++ b/lambda/tests/routes/auth/test_account_login.py @@ -8,20 +8,11 @@ from src.routes.auth.account_login import AccountLoginRoute from src.services.fxa_crypto import derive_verify_hash +from tests.conftest import json_body @pytest.fixture -def mock_account_manager(): - return MagicMock() - - -@pytest.fixture -def mock_token_manager(): - return MagicMock() - - -@pytest.fixture -def route(mock_account_manager, mock_token_manager): +def route(mock_account_manager: MagicMock, mock_token_manager: MagicMock) -> AccountLoginRoute: return AccountLoginRoute( account_manager=mock_account_manager, token_manager=mock_token_manager, @@ -44,7 +35,12 @@ def _make_account(auth_pw_hex: str) -> dict: class TestAccountLogin: - def test_success_with_keys(self, route, mock_account_manager, mock_token_manager): + def test_success_with_keys( + self, + route: AccountLoginRoute, + mock_account_manager: MagicMock, + mock_token_manager: MagicMock, + ) -> None: auth_pw_hex = "cc" * 32 mock_account_manager.get_account_by_email.return_value = _make_account(auth_pw_hex) mock_token_manager.create_session_token.return_value = b"\xaa" * 32 @@ -61,13 +57,18 @@ def test_success_with_keys(self, route, mock_account_manager, mock_token_manager ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["uid"] == "uid1" assert body["sessionToken"] == "aa" * 32 assert body["keyFetchToken"] == "bb" * 32 assert body["verified"] is True - def test_success_without_keys(self, route, mock_account_manager, mock_token_manager): + def test_success_without_keys( + self, + route: AccountLoginRoute, + mock_account_manager: MagicMock, + mock_token_manager: MagicMock, + ) -> None: auth_pw_hex = "cc" * 32 mock_account_manager.get_account_by_email.return_value = _make_account(auth_pw_hex) mock_token_manager.create_session_token.return_value = b"\xaa" * 32 @@ -83,11 +84,13 @@ def test_success_without_keys(self, route, mock_account_manager, mock_token_mana ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["uid"] == "uid1" assert "keyFetchToken" not in body - def test_unknown_email_returns_400(self, route, mock_account_manager): + def test_unknown_email_returns_400( + self, route: AccountLoginRoute, mock_account_manager: MagicMock + ) -> None: mock_account_manager.get_account_by_email.return_value = None event = APIGatewayProxyEvent( { @@ -99,10 +102,12 @@ def test_unknown_email_returns_400(self, route, mock_account_manager): ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 102 - def test_wrong_password_returns_400(self, route, mock_account_manager): + def test_wrong_password_returns_400( + self, route: AccountLoginRoute, mock_account_manager: MagicMock + ) -> None: mock_account_manager.get_account_by_email.return_value = _make_account("cc" * 32) event = APIGatewayProxyEvent( { @@ -114,10 +119,10 @@ def test_wrong_password_returns_400(self, route, mock_account_manager): ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 103 - def test_invalid_json_body_returns_400(self, route): + def test_invalid_json_body_returns_400(self, route: AccountLoginRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -129,7 +134,7 @@ def test_invalid_json_body_returns_400(self, route): response = route.handle(event) assert response.status_code == 400 - def test_missing_email_returns_400(self, route): + def test_missing_email_returns_400(self, route: AccountLoginRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -141,7 +146,7 @@ def test_missing_email_returns_400(self, route): response = route.handle(event) assert response.status_code == 400 - def test_missing_authpw_returns_400(self, route): + def test_missing_authpw_returns_400(self, route: AccountLoginRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -153,7 +158,7 @@ def test_missing_authpw_returns_400(self, route): response = route.handle(event) assert response.status_code == 400 - def test_missing_body_returns_400(self, route): + def test_missing_body_returns_400(self, route: AccountLoginRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -165,7 +170,7 @@ def test_missing_body_returns_400(self, route): response = route.handle(event) assert response.status_code == 400 - def test_invalid_authpw_format_returns_400(self, route): + def test_invalid_authpw_format_returns_400(self, route: AccountLoginRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -176,10 +181,10 @@ def test_invalid_authpw_format_returns_400(self, route): ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert "authPW" in body["message"] - def test_invalid_email_format_returns_400(self, route): + def test_invalid_email_format_returns_400(self, route: AccountLoginRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -190,12 +195,12 @@ def test_invalid_email_format_returns_400(self, route): ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert "email" in body["message"] class TestAccountLoginBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: AccountLoginRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_account_status.py b/lambda/tests/routes/auth/test_account_status.py index da84e0ba..021f1ea0 100644 --- a/lambda/tests/routes/auth/test_account_status.py +++ b/lambda/tests/routes/auth/test_account_status.py @@ -1,26 +1,23 @@ """Unit tests for AccountStatus route""" -import json from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.account_status import AccountStatusRoute +from tests.conftest import json_body @pytest.fixture -def mock_account_manager(): - return MagicMock() - - -@pytest.fixture -def route(mock_account_manager): +def route(mock_account_manager: MagicMock) -> AccountStatusRoute: return AccountStatusRoute(account_manager=mock_account_manager) class TestAccountStatus: - def test_returns_true_for_existing_account(self, route, mock_account_manager): + def test_returns_true_for_existing_account( + self, route: AccountStatusRoute, mock_account_manager: MagicMock + ) -> None: mock_account_manager.get_account_by_email.return_value = {"uid": "uid1"} event = APIGatewayProxyEvent( { @@ -32,10 +29,12 @@ def test_returns_true_for_existing_account(self, route, mock_account_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["exists"] is True - def test_returns_false_for_unknown_account(self, route, mock_account_manager): + def test_returns_false_for_unknown_account( + self, route: AccountStatusRoute, mock_account_manager: MagicMock + ) -> None: mock_account_manager.get_account_by_email.return_value = None event = APIGatewayProxyEvent( { @@ -47,10 +46,10 @@ def test_returns_false_for_unknown_account(self, route, mock_account_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["exists"] is False - def test_returns_400_for_missing_email(self, route): + def test_returns_400_for_missing_email(self, route: AccountStatusRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -62,7 +61,7 @@ def test_returns_400_for_missing_email(self, route): response = route.handle(event) assert response.status_code == 400 - def test_returns_400_for_null_query_params(self, route): + def test_returns_400_for_null_query_params(self, route: AccountStatusRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -76,7 +75,7 @@ def test_returns_400_for_null_query_params(self, route): class TestAccountStatusBind: - def test_bind_registers_get_route(self, route): + def test_bind_registers_get_route(self, route: AccountStatusRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_jwks.py b/lambda/tests/routes/auth/test_jwks.py index b00db378..aaec9dba 100644 --- a/lambda/tests/routes/auth/test_jwks.py +++ b/lambda/tests/routes/auth/test_jwks.py @@ -1,16 +1,16 @@ """Unit tests for JWKS route""" -import json from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.jwks import JWKSRoute +from tests.conftest import json_body @pytest.fixture -def mock_jwt_service(): +def mock_jwt_service() -> MagicMock: svc = MagicMock() svc.get_public_key_jwk.return_value = { "kty": "RSA", @@ -24,12 +24,12 @@ def mock_jwt_service(): @pytest.fixture -def route(mock_jwt_service): +def route(mock_jwt_service: MagicMock) -> JWKSRoute: return JWKSRoute(jwt_service=mock_jwt_service) class TestJWKS: - def test_returns_jwks(self, route, mock_jwt_service): + def test_returns_jwks(self, route: JWKSRoute, mock_jwt_service: MagicMock) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -39,7 +39,7 @@ def test_returns_jwks(self, route, mock_jwt_service): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert "keys" in body assert len(body["keys"]) == 1 key = body["keys"][0] @@ -50,7 +50,7 @@ def test_returns_jwks(self, route, mock_jwt_service): class TestJWKSBind: - def test_bind_registers_get_route(self, route): + def test_bind_registers_get_route(self, route: JWKSRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_oauth_authorization.py b/lambda/tests/routes/auth/test_oauth_authorization.py index ff67bf58..d6dbea1b 100644 --- a/lambda/tests/routes/auth/test_oauth_authorization.py +++ b/lambda/tests/routes/auth/test_oauth_authorization.py @@ -1,28 +1,25 @@ """Unit tests for OAuthAuthorization route""" import json +from typing import Optional from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.oauth_authorization import OAuthAuthorizationRoute +from tests.conftest import json_body @pytest.fixture -def mock_oauth_code_manager(): - return MagicMock() - - -@pytest.fixture -def route(mock_oauth_code_manager): +def route(mock_oauth_code_manager: MagicMock) -> OAuthAuthorizationRoute: return OAuthAuthorizationRoute( oauth_code_manager=mock_oauth_code_manager, middlewares=[], ) -def _make_event(body=None, hawk_uid="uid1"): +def _make_event(body: Optional[str] = None, hawk_uid: str = "uid1") -> APIGatewayProxyEvent: """Build an event with hawk_uid pre-injected (middleware handled auth).""" return APIGatewayProxyEvent( { @@ -36,7 +33,9 @@ def _make_event(body=None, hawk_uid="uid1"): class TestOAuthAuthorization: - def test_success_returns_code_and_state(self, route, mock_oauth_code_manager): + def test_success_returns_code_and_state( + self, route: OAuthAuthorizationRoute, mock_oauth_code_manager: MagicMock + ) -> None: mock_oauth_code_manager.create_authorization_code.return_value = "auth-code-123" event = _make_event( @@ -52,43 +51,45 @@ def test_success_returns_code_and_state(self, route, mock_oauth_code_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["code"] == "auth-code-123" assert body["state"] == "state123" assert body["redirect"] == "urn:ietf:wg:oauth:2.0:oob" - def test_missing_body_returns_400(self, route): + def test_missing_body_returns_400(self, route: OAuthAuthorizationRoute) -> None: event = _make_event(body=None) response = route.handle(event) assert response.status_code == 400 - def test_missing_client_id_returns_400(self, route): + def test_missing_client_id_returns_400(self, route: OAuthAuthorizationRoute) -> None: event = _make_event( body=json.dumps({"scope": "s", "state": "st"}), ) response = route.handle(event) assert response.status_code == 400 - def test_invalid_json_body_returns_400(self, route): + def test_invalid_json_body_returns_400(self, route: OAuthAuthorizationRoute) -> None: event = _make_event(body="not-json") response = route.handle(event) assert response.status_code == 400 - def test_missing_scope_returns_400(self, route): + def test_missing_scope_returns_400(self, route: OAuthAuthorizationRoute) -> None: event = _make_event( body=json.dumps({"client_id": "c", "state": "st"}), ) response = route.handle(event) assert response.status_code == 400 - def test_missing_state_returns_400(self, route): + def test_missing_state_returns_400(self, route: OAuthAuthorizationRoute) -> None: event = _make_event( body=json.dumps({"client_id": "c", "scope": "s"}), ) response = route.handle(event) assert response.status_code == 400 - def test_creates_code_with_correct_params(self, route, mock_oauth_code_manager): + def test_creates_code_with_correct_params( + self, route: OAuthAuthorizationRoute, mock_oauth_code_manager: MagicMock + ) -> None: mock_oauth_code_manager.create_authorization_code.return_value = "code" event = _make_event( @@ -112,7 +113,9 @@ def test_creates_code_with_correct_params(self, route, mock_oauth_code_manager): keys_jwe="", ) - def test_passes_keys_jwe_to_code_manager(self, route, mock_oauth_code_manager): + def test_passes_keys_jwe_to_code_manager( + self, route: OAuthAuthorizationRoute, mock_oauth_code_manager: MagicMock + ) -> None: mock_oauth_code_manager.create_authorization_code.return_value = "code" event = _make_event( @@ -139,7 +142,9 @@ def test_passes_keys_jwe_to_code_manager(self, route, mock_oauth_code_manager): class TestOAuthAuthorizationRedirectUri: - def test_pairing_redirect_uri_accepted(self, route, mock_oauth_code_manager): + def test_pairing_redirect_uri_accepted( + self, route: OAuthAuthorizationRoute, mock_oauth_code_manager: MagicMock + ) -> None: """Pairing redirect_uri is accepted and returned in response.""" mock_oauth_code_manager.create_authorization_code.return_value = "code-pair" @@ -155,10 +160,12 @@ def test_pairing_redirect_uri_accepted(self, route, mock_oauth_code_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["redirect"] == "urn:ietf:wg:oauth:2.0:oob:pair-auth-webchannel" - def test_invalid_redirect_uri_returns_400(self, route, mock_oauth_code_manager): + def test_invalid_redirect_uri_returns_400( + self, route: OAuthAuthorizationRoute, mock_oauth_code_manager: MagicMock + ) -> None: """Invalid redirect_uri returns 400.""" event = _make_event( body=json.dumps( @@ -172,11 +179,13 @@ def test_invalid_redirect_uri_returns_400(self, route, mock_oauth_code_manager): ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 107 assert "redirect_uri" in body["message"] - def test_default_redirect_uri_when_not_provided(self, route, mock_oauth_code_manager): + def test_default_redirect_uri_when_not_provided( + self, route: OAuthAuthorizationRoute, mock_oauth_code_manager: MagicMock + ) -> None: """Default redirect_uri used when none provided.""" mock_oauth_code_manager.create_authorization_code.return_value = "code-default" @@ -191,12 +200,12 @@ def test_default_redirect_uri_when_not_provided(self, route, mock_oauth_code_man ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["redirect"] == "urn:ietf:wg:oauth:2.0:oob" class TestOAuthAuthorizationBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: OAuthAuthorizationRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_oauth_destroy.py b/lambda/tests/routes/auth/test_oauth_destroy.py index 0185ac5a..e3293124 100644 --- a/lambda/tests/routes/auth/test_oauth_destroy.py +++ b/lambda/tests/routes/auth/test_oauth_destroy.py @@ -8,20 +8,18 @@ from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.oauth_destroy import OAuthDestroyRoute +from tests.conftest import json_body @pytest.fixture -def mock_oauth_code_manager(): - return MagicMock() - - -@pytest.fixture -def route(mock_oauth_code_manager): +def route(mock_oauth_code_manager: MagicMock) -> OAuthDestroyRoute: return OAuthDestroyRoute(oauth_code_manager=mock_oauth_code_manager) class TestOAuthDestroy: - def test_revokes_token_and_returns_200(self, route, mock_oauth_code_manager): + def test_revokes_token_and_returns_200( + self, route: OAuthDestroyRoute, mock_oauth_code_manager: MagicMock + ) -> None: token = "abc123" event = APIGatewayProxyEvent( { @@ -33,11 +31,13 @@ def test_revokes_token_and_returns_200(self, route, mock_oauth_code_manager): ) response = route.handle(event) assert response.status_code == 200 - assert json.loads(response.body) == {} + assert json_body(response) == {} expected_hash = hashlib.sha256(token.encode("ascii")).hexdigest() mock_oauth_code_manager.delete_refresh_token.assert_called_once_with(expected_hash) - def test_returns_200_even_if_token_does_not_exist(self, route, mock_oauth_code_manager): + def test_returns_200_even_if_token_does_not_exist( + self, route: OAuthDestroyRoute, mock_oauth_code_manager: MagicMock + ) -> None: """Per RFC 7009, revoking an invalid token is not an error.""" mock_oauth_code_manager.delete_refresh_token.return_value = None event = APIGatewayProxyEvent( @@ -50,9 +50,9 @@ def test_returns_200_even_if_token_does_not_exist(self, route, mock_oauth_code_m ) response = route.handle(event) assert response.status_code == 200 - assert json.loads(response.body) == {} + assert json_body(response) == {} - def test_missing_body_returns_400(self, route): + def test_missing_body_returns_400(self, route: OAuthDestroyRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -63,10 +63,10 @@ def test_missing_body_returns_400(self, route): ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 107 - def test_invalid_json_body_returns_400(self, route): + def test_invalid_json_body_returns_400(self, route: OAuthDestroyRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -78,7 +78,7 @@ def test_invalid_json_body_returns_400(self, route): response = route.handle(event) assert response.status_code == 400 - def test_missing_token_field_returns_400(self, route): + def test_missing_token_field_returns_400(self, route: OAuthDestroyRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -89,12 +89,12 @@ def test_missing_token_field_returns_400(self, route): ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 107 class TestOAuthDestroyBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: OAuthDestroyRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_oauth_token.py b/lambda/tests/routes/auth/test_oauth_token.py index 5932472f..19411118 100644 --- a/lambda/tests/routes/auth/test_oauth_token.py +++ b/lambda/tests/routes/auth/test_oauth_token.py @@ -7,30 +7,21 @@ from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.oauth_token import OAuthTokenRoute +from tests.conftest import json_body @pytest.fixture -def mock_oauth_code_manager(): +def mock_jwt_service() -> MagicMock: return MagicMock() @pytest.fixture -def mock_jwt_service(): - return MagicMock() - - -@pytest.fixture -def mock_account_manager(): - return MagicMock() - - -@pytest.fixture -def mock_token_manager(): - return MagicMock() - - -@pytest.fixture -def route(mock_oauth_code_manager, mock_jwt_service, mock_account_manager, mock_token_manager): +def route( + mock_oauth_code_manager: MagicMock, + mock_jwt_service: MagicMock, + mock_account_manager: MagicMock, + mock_token_manager: MagicMock, +) -> OAuthTokenRoute: return OAuthTokenRoute( oauth_code_manager=mock_oauth_code_manager, jwt_service=mock_jwt_service, @@ -47,12 +38,12 @@ class TestOAuthTokenAuthorizationCode: ) def test_success_returns_tokens( self, - mock_verify, - route, - mock_oauth_code_manager, - mock_jwt_service, - mock_account_manager, - ): + mock_verify: MagicMock, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_jwt_service: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_authorization_code.return_value = { "uid": "uid1", "clientId": "client1", @@ -85,7 +76,7 @@ def test_success_returns_tokens( ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["access_token"] == "jwt-access-token" assert body["refresh_token"] == "refresh-tok" assert body["token_type"] == "bearer" @@ -99,12 +90,12 @@ def test_success_returns_tokens( ) def test_omits_keys_jwe_when_empty( self, - mock_verify, - route, - mock_oauth_code_manager, - mock_jwt_service, - mock_account_manager, - ): + mock_verify: MagicMock, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_jwt_service: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_authorization_code.return_value = { "uid": "uid1", "clientId": "client1", @@ -137,10 +128,12 @@ def test_omits_keys_jwe_when_empty( ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert "keys_jwe" not in body - def test_invalid_code_returns_400(self, route, mock_oauth_code_manager): + def test_invalid_code_returns_400( + self, route: OAuthTokenRoute, mock_oauth_code_manager: MagicMock + ) -> None: mock_oauth_code_manager.consume_authorization_code.return_value = None event = APIGatewayProxyEvent( { @@ -165,8 +158,12 @@ def test_invalid_code_returns_400(self, route, mock_oauth_code_manager): return_value=False, ) def test_invalid_pkce_returns_400( - self, mock_verify, route, mock_oauth_code_manager, mock_account_manager - ): + self, + mock_verify: MagicMock, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_authorization_code.return_value = { "uid": "uid1", "clientId": "client1", @@ -200,8 +197,12 @@ def test_invalid_pkce_returns_400( class TestOAuthTokenRefreshToken: def test_success_returns_new_access_token( - self, route, mock_oauth_code_manager, mock_jwt_service, mock_account_manager - ): + self, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_jwt_service: MagicMock, + mock_account_manager: MagicMock, + ) -> None: token = "refresh-token-value" mock_oauth_code_manager.consume_refresh_token.return_value = { "uid": "uid1", @@ -231,11 +232,13 @@ def test_success_returns_new_access_token( ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["access_token"] == "new-jwt" assert body["refresh_token"] == "new-refresh" - def test_invalid_refresh_token_returns_400(self, route, mock_oauth_code_manager): + def test_invalid_refresh_token_returns_400( + self, route: OAuthTokenRoute, mock_oauth_code_manager: MagicMock + ) -> None: mock_oauth_code_manager.consume_refresh_token.return_value = None event = APIGatewayProxyEvent( { @@ -256,7 +259,9 @@ def test_invalid_refresh_token_returns_400(self, route, mock_oauth_code_manager) class TestOAuthTokenAuthCodeEdgeCases: - def test_missing_code_returns_400(self, route, mock_oauth_code_manager): + def test_missing_code_returns_400( + self, route: OAuthTokenRoute, mock_oauth_code_manager: MagicMock + ) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -269,8 +274,11 @@ def test_missing_code_returns_400(self, route, mock_oauth_code_manager): assert response.status_code == 400 def test_missing_code_verifier_when_challenge_present_returns_400( - self, route, mock_oauth_code_manager, mock_account_manager - ): + self, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_authorization_code.return_value = { "uid": "uid1", "clientId": "client1", @@ -300,8 +308,11 @@ def test_missing_code_verifier_when_challenge_present_returns_400( assert response.status_code == 400 def test_account_not_found_returns_400( - self, route, mock_oauth_code_manager, mock_account_manager - ): + self, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_authorization_code.return_value = { "uid": "uid1", "clientId": "client1", @@ -331,8 +342,12 @@ def test_account_not_found_returns_400( class TestOAuthTokenTTLCap: def test_ttl_capped_at_max( - self, route, mock_oauth_code_manager, mock_jwt_service, mock_account_manager - ): + self, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_jwt_service: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_authorization_code.return_value = { "uid": "uid1", "clientId": "client1", @@ -365,14 +380,17 @@ def test_ttl_capped_at_max( ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["expires_in"] == 3600 # MAX_TTL class TestOAuthTokenScopeValidation: def test_refresh_scope_exceeds_grant_returns_400( - self, route, mock_oauth_code_manager, mock_account_manager - ): + self, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_refresh_token.return_value = { "uid": "uid1", "clientId": "client1", @@ -395,12 +413,16 @@ def test_refresh_scope_exceeds_grant_returns_400( ) response = route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 165 def test_refresh_scope_subset_succeeds( - self, route, mock_oauth_code_manager, mock_jwt_service, mock_account_manager - ): + self, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_jwt_service: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_refresh_token.return_value = { "uid": "uid1", "clientId": "client1", @@ -432,7 +454,7 @@ def test_refresh_scope_subset_succeeds( class TestOAuthTokenRefreshEdgeCases: - def test_missing_refresh_token_returns_400(self, route): + def test_missing_refresh_token_returns_400(self, route: OAuthTokenRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -445,8 +467,11 @@ def test_missing_refresh_token_returns_400(self, route): assert response.status_code == 400 def test_account_not_found_on_refresh_returns_400( - self, route, mock_oauth_code_manager, mock_account_manager - ): + self, + route: OAuthTokenRoute, + mock_oauth_code_manager: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_oauth_code_manager.consume_refresh_token.return_value = { "uid": "uid1", "clientId": "client1", @@ -472,7 +497,7 @@ def test_account_not_found_on_refresh_returns_400( class TestOAuthTokenErrors: - def test_invalid_json_body_returns_400(self, route): + def test_invalid_json_body_returns_400(self, route: OAuthTokenRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -484,7 +509,7 @@ def test_invalid_json_body_returns_400(self, route): response = route.handle(event) assert response.status_code == 400 - def test_missing_grant_type_returns_400(self, route): + def test_missing_grant_type_returns_400(self, route: OAuthTokenRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -496,7 +521,7 @@ def test_missing_grant_type_returns_400(self, route): response = route.handle(event) assert response.status_code == 400 - def test_invalid_grant_type_returns_400(self, route): + def test_invalid_grant_type_returns_400(self, route: OAuthTokenRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -508,7 +533,7 @@ def test_invalid_grant_type_returns_400(self, route): response = route.handle(event) assert response.status_code == 400 - def test_missing_body_returns_400(self, route): + def test_missing_body_returns_400(self, route: OAuthTokenRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -523,8 +548,12 @@ def test_missing_body_returns_400(self, route): class TestOAuthTokenFxaCredentials: def test_success_returns_access_token( - self, route, mock_token_manager, mock_jwt_service, mock_account_manager - ): + self, + route: OAuthTokenRoute, + mock_token_manager: MagicMock, + mock_jwt_service: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_token_manager.verify_session_hawk.return_value = "uid1" mock_jwt_service.sign_jwt.return_value = "fxa-cred-jwt" mock_account_manager.get_account_by_uid.return_value = { @@ -548,13 +577,15 @@ def test_success_returns_access_token( ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["access_token"] == "fxa-cred-jwt" assert body["token_type"] == "bearer" assert body["scope"] == "profile" assert "refresh_token" not in body - def test_missing_auth_returns_401(self, route, mock_token_manager): + def test_missing_auth_returns_401( + self, route: OAuthTokenRoute, mock_token_manager: MagicMock + ) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -572,7 +603,9 @@ def test_missing_auth_returns_401(self, route, mock_token_manager): response = route.handle(event) assert response.status_code == 401 - def test_invalid_session_returns_401(self, route, mock_token_manager): + def test_invalid_session_returns_401( + self, route: OAuthTokenRoute, mock_token_manager: MagicMock + ) -> None: mock_token_manager.verify_session_hawk.return_value = None event = APIGatewayProxyEvent( { @@ -592,8 +625,11 @@ def test_invalid_session_returns_401(self, route, mock_token_manager): assert response.status_code == 401 def test_returns_400_when_token_manager_not_configured( - self, mock_oauth_code_manager, mock_jwt_service, mock_account_manager - ): + self, + mock_oauth_code_manager: MagicMock, + mock_jwt_service: MagicMock, + mock_account_manager: MagicMock, + ) -> None: route_no_tm = OAuthTokenRoute( oauth_code_manager=mock_oauth_code_manager, jwt_service=mock_jwt_service, @@ -617,7 +653,12 @@ def test_returns_400_when_token_manager_not_configured( response = route_no_tm.handle(event) assert response.status_code == 400 - def test_account_not_found_returns_400(self, route, mock_token_manager, mock_account_manager): + def test_account_not_found_returns_400( + self, + route: OAuthTokenRoute, + mock_token_manager: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_token_manager.verify_session_hawk.return_value = "uid1" mock_account_manager.get_account_by_uid.return_value = None event = APIGatewayProxyEvent( @@ -639,7 +680,7 @@ def test_account_not_found_returns_400(self, route, mock_token_manager, mock_acc class TestOAuthTokenBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: OAuthTokenRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_oidc_discovery.py b/lambda/tests/routes/auth/test_oidc_discovery.py index 0a0444e4..983c2abc 100644 --- a/lambda/tests/routes/auth/test_oidc_discovery.py +++ b/lambda/tests/routes/auth/test_oidc_discovery.py @@ -1,28 +1,28 @@ """Unit tests for OIDCDiscovery route""" -import json from unittest.mock import MagicMock, PropertyMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.oidc_discovery import OIDCDiscoveryRoute +from tests.conftest import json_body @pytest.fixture -def mock_jwt_service(): +def mock_jwt_service() -> MagicMock: svc = MagicMock() type(svc).issuer = PropertyMock(return_value="https://auth.beta.ffsync.layertwo.dev") return svc @pytest.fixture -def route(mock_jwt_service): +def route(mock_jwt_service: MagicMock) -> OIDCDiscoveryRoute: return OIDCDiscoveryRoute(jwt_service=mock_jwt_service) class TestOIDCDiscovery: - def test_returns_discovery_document(self, route): + def test_returns_discovery_document(self, route: OIDCDiscoveryRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -32,7 +32,7 @@ def test_returns_discovery_document(self, route): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["issuer"] == "https://auth.beta.ffsync.layertwo.dev" assert body["authorization_endpoint"].endswith("/v1/oauth/authorization") assert body["token_endpoint"].endswith("/v1/oauth/token") @@ -42,7 +42,7 @@ def test_returns_discovery_document(self, route): class TestOIDCDiscoveryBind: - def test_bind_registers_get_route(self, route): + def test_bind_registers_get_route(self, route: OIDCDiscoveryRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_oidc_exchange.py b/lambda/tests/routes/auth/test_oidc_exchange.py index 1c751fe3..560d418e 100644 --- a/lambda/tests/routes/auth/test_oidc_exchange.py +++ b/lambda/tests/routes/auth/test_oidc_exchange.py @@ -1,6 +1,7 @@ """Unit tests for OIDC exchange routes""" import json +from typing import Any, Dict, Optional from unittest.mock import MagicMock, patch import pytest @@ -8,6 +9,7 @@ from src.routes.auth.oidc_exchange import OIDCCodeExchangeRoute, OIDCProviderConfigRoute from src.shared.oidc import OIDCProviderConfig +from tests.conftest import json_body # ============================================================================ # Fixtures @@ -15,7 +17,7 @@ @pytest.fixture -def mock_oidc_validator(): +def mock_oidc_validator() -> MagicMock: validator = MagicMock() validator.client_id = "test-client-id" validator.discover_provider_config.return_value = OIDCProviderConfig( @@ -29,17 +31,14 @@ def mock_oidc_validator(): @pytest.fixture -def mock_account_manager(): - return MagicMock() - - -@pytest.fixture -def config_route(mock_oidc_validator): +def config_route(mock_oidc_validator: MagicMock) -> OIDCProviderConfigRoute: return OIDCProviderConfigRoute(oidc_validator=mock_oidc_validator) @pytest.fixture -def exchange_route(mock_oidc_validator, mock_account_manager): +def exchange_route( + mock_oidc_validator: MagicMock, mock_account_manager: MagicMock +) -> OIDCCodeExchangeRoute: return OIDCCodeExchangeRoute( oidc_validator=mock_oidc_validator, account_manager=mock_account_manager, @@ -48,7 +47,12 @@ def exchange_route(mock_oidc_validator, mock_account_manager): ) -def _make_event(method="GET", path="/", body=None, headers=None): +def _make_event( + method: str = "GET", + path: str = "/", + body: Any = None, + headers: Optional[Dict[str, str]] = None, +) -> APIGatewayProxyEvent: event_dict = { "httpMethod": method, "path": path, @@ -65,30 +69,34 @@ def _make_event(method="GET", path="/", body=None, headers=None): class TestOIDCProviderConfig: - def test_returns_authorization_endpoint(self, config_route): + def test_returns_authorization_endpoint(self, config_route: OIDCProviderConfigRoute) -> None: event = _make_event("GET", "/v1/oidc/config") response = config_route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["authorization_endpoint"] == "https://idp.example.com/authorize" - def test_returns_only_authorization_endpoint(self, config_route): + def test_returns_only_authorization_endpoint( + self, config_route: OIDCProviderConfigRoute + ) -> None: event = _make_event("GET", "/v1/oidc/config") response = config_route.handle(event) - body = json.loads(response.body) + body = json_body(response) assert list(body.keys()) == ["authorization_endpoint"] - def test_returns_503_when_provider_unavailable(self, config_route, mock_oidc_validator): + def test_returns_503_when_provider_unavailable( + self, config_route: OIDCProviderConfigRoute, mock_oidc_validator: MagicMock + ) -> None: mock_oidc_validator.discover_provider_config.side_effect = Exception("unreachable") event = _make_event("GET", "/v1/oidc/config") response = config_route.handle(event) assert response.status_code == 503 - body = json.loads(response.body) + body = json_body(response) assert "unavailable" in body["message"].lower() class TestOIDCProviderConfigBind: - def test_bind_registers_get_route(self, config_route): + def test_bind_registers_get_route(self, config_route: OIDCProviderConfigRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) config_route.bind(mock_api) @@ -101,7 +109,7 @@ def test_bind_registers_get_route(self, config_route): class TestOIDCCodeExchange: - def _exchange_event(self, body=None): + def _exchange_event(self, body: Optional[Dict[str, Any]] = None) -> APIGatewayProxyEvent: default_body = { "code": "auth-code-123", "code_verifier": "verifier-456", @@ -111,8 +119,12 @@ def _exchange_event(self, body=None): @patch("src.routes.auth.oidc_exchange.requests") def test_success_account_exists( - self, mock_requests, exchange_route, mock_oidc_validator, mock_account_manager - ): + self, + mock_requests: MagicMock, + exchange_route: OIDCCodeExchangeRoute, + mock_oidc_validator: MagicMock, + mock_account_manager: MagicMock, + ) -> None: # Token exchange response mock_token_resp = MagicMock() mock_token_resp.ok = True @@ -132,15 +144,19 @@ def test_success_account_exists( response = exchange_route.handle(self._exchange_event()) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["email"] == "user@example.com" assert body["access_token"] == "at-789" assert body["account_exists"] is True @patch("src.routes.auth.oidc_exchange.requests") def test_success_account_does_not_exist( - self, mock_requests, exchange_route, mock_oidc_validator, mock_account_manager - ): + self, + mock_requests: MagicMock, + exchange_route: OIDCCodeExchangeRoute, + mock_oidc_validator: MagicMock, + mock_account_manager: MagicMock, + ) -> None: mock_token_resp = MagicMock() mock_token_resp.ok = True mock_token_resp.json.return_value = {"access_token": "at-789"} @@ -158,31 +174,35 @@ def test_success_account_does_not_exist( response = exchange_route.handle(self._exchange_event()) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["email"] == "new@example.com" assert body["account_exists"] is False - def test_invalid_json_body(self, exchange_route): + def test_invalid_json_body(self, exchange_route: OIDCCodeExchangeRoute) -> None: event = _make_event("POST", "/v1/oidc/exchange", body="not-json") response = exchange_route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert "Invalid JSON" in body["message"] - def test_missing_required_fields(self, exchange_route): + def test_missing_required_fields(self, exchange_route: OIDCCodeExchangeRoute) -> None: event = _make_event("POST", "/v1/oidc/exchange", body={"code": "abc"}) response = exchange_route.handle(event) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert "Missing required" in body["message"] - def test_provider_discovery_failure(self, exchange_route, mock_oidc_validator): + def test_provider_discovery_failure( + self, exchange_route: OIDCCodeExchangeRoute, mock_oidc_validator: MagicMock + ) -> None: mock_oidc_validator.discover_provider_config.side_effect = Exception("fail") response = exchange_route.handle(self._exchange_event()) assert response.status_code == 503 @patch("src.routes.auth.oidc_exchange.requests") - def test_token_exchange_network_error(self, mock_requests, exchange_route): + def test_token_exchange_network_error( + self, mock_requests: MagicMock, exchange_route: OIDCCodeExchangeRoute + ) -> None: import requests mock_requests.post.side_effect = requests.exceptions.ConnectionError("timeout") @@ -190,11 +210,13 @@ def test_token_exchange_network_error(self, mock_requests, exchange_route): response = exchange_route.handle(self._exchange_event()) assert response.status_code == 502 - body = json.loads(response.body) + body = json_body(response) assert "exchange" in body["message"].lower() @patch("src.routes.auth.oidc_exchange.requests") - def test_token_exchange_provider_rejects(self, mock_requests, exchange_route): + def test_token_exchange_provider_rejects( + self, mock_requests: MagicMock, exchange_route: OIDCCodeExchangeRoute + ) -> None: mock_token_resp = MagicMock() mock_token_resp.ok = False mock_token_resp.status_code = 400 @@ -206,7 +228,9 @@ def test_token_exchange_provider_rejects(self, mock_requests, exchange_route): assert response.status_code == 401 @patch("src.routes.auth.oidc_exchange.requests") - def test_no_access_token_in_response(self, mock_requests, exchange_route): + def test_no_access_token_in_response( + self, mock_requests: MagicMock, exchange_route: OIDCCodeExchangeRoute + ) -> None: mock_token_resp = MagicMock() mock_token_resp.ok = True mock_token_resp.json.return_value = {"id_token": "something"} @@ -215,11 +239,16 @@ def test_no_access_token_in_response(self, mock_requests, exchange_route): response = exchange_route.handle(self._exchange_event()) assert response.status_code == 502 - body = json.loads(response.body) + body = json_body(response) assert "access token" in body["message"].lower() @patch("src.routes.auth.oidc_exchange.requests") - def test_token_validation_failure(self, mock_requests, exchange_route, mock_oidc_validator): + def test_token_validation_failure( + self, + mock_requests: MagicMock, + exchange_route: OIDCCodeExchangeRoute, + mock_oidc_validator: MagicMock, + ) -> None: mock_token_resp = MagicMock() mock_token_resp.ok = True mock_token_resp.json.return_value = {"access_token": "bad-token"} @@ -232,7 +261,12 @@ def test_token_validation_failure(self, mock_requests, exchange_route, mock_oidc assert response.status_code == 401 @patch("src.routes.auth.oidc_exchange.requests") - def test_userinfo_network_error(self, mock_requests, exchange_route, mock_oidc_validator): + def test_userinfo_network_error( + self, + mock_requests: MagicMock, + exchange_route: OIDCCodeExchangeRoute, + mock_oidc_validator: MagicMock, + ) -> None: import requests mock_token_resp = MagicMock() @@ -248,7 +282,12 @@ def test_userinfo_network_error(self, mock_requests, exchange_route, mock_oidc_v assert response.status_code == 502 @patch("src.routes.auth.oidc_exchange.requests") - def test_userinfo_provider_error(self, mock_requests, exchange_route, mock_oidc_validator): + def test_userinfo_provider_error( + self, + mock_requests: MagicMock, + exchange_route: OIDCCodeExchangeRoute, + mock_oidc_validator: MagicMock, + ) -> None: mock_token_resp = MagicMock() mock_token_resp.ok = True mock_token_resp.json.return_value = {"access_token": "at-789"} @@ -268,7 +307,12 @@ def test_userinfo_provider_error(self, mock_requests, exchange_route, mock_oidc_ assert response.status_code == 502 @patch("src.routes.auth.oidc_exchange.requests") - def test_userinfo_missing_email(self, mock_requests, exchange_route, mock_oidc_validator): + def test_userinfo_missing_email( + self, + mock_requests: MagicMock, + exchange_route: OIDCCodeExchangeRoute, + mock_oidc_validator: MagicMock, + ) -> None: mock_token_resp = MagicMock() mock_token_resp.ok = True mock_token_resp.json.return_value = {"access_token": "at-789"} @@ -285,12 +329,12 @@ def test_userinfo_missing_email(self, mock_requests, exchange_route, mock_oidc_v response = exchange_route.handle(self._exchange_event()) assert response.status_code == 400 - body = json.loads(response.body) + body = json_body(response) assert "email" in body["message"].lower() class TestOIDCCodeExchangeBind: - def test_bind_registers_post_route(self, exchange_route): + def test_bind_registers_post_route(self, exchange_route: OIDCCodeExchangeRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) exchange_route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_route_dispatch.py b/lambda/tests/routes/auth/test_route_dispatch.py index 119cfc23..17d01ea0 100644 --- a/lambda/tests/routes/auth/test_route_dispatch.py +++ b/lambda/tests/routes/auth/test_route_dispatch.py @@ -3,6 +3,7 @@ These tests exercise the bind() closure bodies which unit tests call directly via handle(). """ +from typing import Dict, Optional from unittest.mock import MagicMock from src.routes.auth.account_attached_clients import AccountAttachedClientsRoute @@ -23,9 +24,17 @@ from src.routes.auth.session_destroy import SessionDestroyRoute from src.routes.auth.session_status import SessionStatusRoute from src.services.api_router import ApiRouter +from src.shared.base_route import BaseRoute -def _make_event(method, path, headers=None, body=None, qs=None, hawk_uid=None): +def _make_event( + method: str, + path: str, + headers: Optional[Dict[str, str]] = None, + body: Optional[str] = None, + qs: Optional[Dict[str, str]] = None, + hawk_uid: Optional[str] = None, +) -> dict: ctx = {"requestId": "test"} if hawk_uid: ctx["hawk_uid"] = hawk_uid @@ -40,20 +49,20 @@ def _make_event(method, path, headers=None, body=None, qs=None, hawk_uid=None): } -def _make_context(): +def _make_context() -> MagicMock: ctx = MagicMock() ctx.function_name = "test" return ctx -def _router(route): +def _router(route: BaseRoute) -> ApiRouter: return ApiRouter(routes=[route], middlewares=[]) class TestRouteDispatch: """Verify each route's bind closure is exercised via ApiRouter.""" - def test_account_status_dispatches(self): + def test_account_status_dispatches(self) -> None: mgr = MagicMock() mgr.get_account_by_email.return_value = None router = _router(AccountStatusRoute(account_manager=mgr)) @@ -62,7 +71,7 @@ def test_account_status_dispatches(self): ) assert result["statusCode"] == 200 - def test_account_create_dispatches(self): + def test_account_create_dispatches(self) -> None: route = AccountCreateRoute( account_manager=MagicMock(), token_manager=MagicMock(), oidc_validator=MagicMock() ) @@ -71,19 +80,19 @@ def test_account_create_dispatches(self): ) assert result["statusCode"] == 401 - def test_account_login_dispatches(self): + def test_account_login_dispatches(self) -> None: route = AccountLoginRoute(account_manager=MagicMock(), token_manager=MagicMock()) result = _router(route).handler( _make_event("POST", "/v1/account/login", body="{}"), _make_context() ) assert result["statusCode"] == 400 - def test_account_keys_dispatches(self): + def test_account_keys_dispatches(self) -> None: route = AccountKeysRoute(account_manager=MagicMock(), token_manager=MagicMock()) result = _router(route).handler(_make_event("GET", "/v1/account/keys"), _make_context()) assert result["statusCode"] == 401 - def test_scoped_key_data_dispatches(self): + def test_scoped_key_data_dispatches(self) -> None: mgr = MagicMock() mgr.get_account_by_uid.return_value = None route = ScopedKeyDataRoute(account_manager=mgr, middlewares=[]) @@ -99,7 +108,7 @@ def test_scoped_key_data_dispatches(self): # Account not found returns 401 assert result["statusCode"] == 401 - def test_session_status_dispatches(self): + def test_session_status_dispatches(self) -> None: route = SessionStatusRoute(middlewares=[]) result = _router(route).handler( _make_event("GET", "/v1/session/status", hawk_uid="uid1"), @@ -107,7 +116,7 @@ def test_session_status_dispatches(self): ) assert result["statusCode"] == 200 - def test_session_destroy_dispatches(self): + def test_session_destroy_dispatches(self) -> None: route = SessionDestroyRoute(token_manager=MagicMock(), middlewares=[]) result = _router(route).handler( _make_event("POST", "/v1/session/destroy"), @@ -115,7 +124,7 @@ def test_session_destroy_dispatches(self): ) assert result["statusCode"] == 200 - def test_oauth_authorization_dispatches(self): + def test_oauth_authorization_dispatches(self) -> None: route = OAuthAuthorizationRoute(oauth_code_manager=MagicMock(), middlewares=[]) result = _router(route).handler( _make_event( @@ -129,7 +138,7 @@ def test_oauth_authorization_dispatches(self): # Missing client_id returns 400 assert result["statusCode"] == 400 - def test_oauth_token_dispatches(self): + def test_oauth_token_dispatches(self) -> None: route = OAuthTokenRoute( oauth_code_manager=MagicMock(), jwt_service=MagicMock(), @@ -141,14 +150,14 @@ def test_oauth_token_dispatches(self): ) assert result["statusCode"] == 400 - def test_oauth_destroy_dispatches(self): + def test_oauth_destroy_dispatches(self) -> None: route = OAuthDestroyRoute(oauth_code_manager=MagicMock()) result = _router(route).handler( _make_event("POST", "/v1/oauth/destroy", body="{}"), _make_context() ) assert result["statusCode"] in (200, 400) - def test_oidc_discovery_dispatches(self): + def test_oidc_discovery_dispatches(self) -> None: jwt_svc = MagicMock() jwt_svc.issuer = "https://auth.example.com" route = OIDCDiscoveryRoute(jwt_service=jwt_svc) @@ -157,14 +166,14 @@ def test_oidc_discovery_dispatches(self): ) assert result["statusCode"] == 200 - def test_jwks_dispatches(self): + def test_jwks_dispatches(self) -> None: jwt_svc = MagicMock() jwt_svc.get_public_key_jwk.return_value = {"kty": "RSA", "kid": "test"} route = JWKSRoute(jwt_service=jwt_svc) result = _router(route).handler(_make_event("GET", "/v1/jwks"), _make_context()) assert result["statusCode"] == 200 - def test_oidc_provider_config_dispatches(self): + def test_oidc_provider_config_dispatches(self) -> None: validator = MagicMock() validator.discover_provider_config.return_value = MagicMock( authorization_endpoint="https://idp.example.com/authorize" @@ -173,7 +182,7 @@ def test_oidc_provider_config_dispatches(self): result = _router(route).handler(_make_event("GET", "/v1/oidc/config"), _make_context()) assert result["statusCode"] == 200 - def test_oidc_code_exchange_dispatches(self): + def test_oidc_code_exchange_dispatches(self) -> None: route = OIDCCodeExchangeRoute( oidc_validator=MagicMock(), account_manager=MagicMock(), @@ -185,7 +194,7 @@ def test_oidc_code_exchange_dispatches(self): ) assert result["statusCode"] == 400 - def test_account_device_dispatches(self): + def test_account_device_dispatches(self) -> None: mgr = MagicMock() mgr.upsert_device.return_value = {"id": "dev1", "name": "Test"} route = AccountDeviceRoute(device_manager=mgr, middlewares=[]) @@ -195,7 +204,7 @@ def test_account_device_dispatches(self): ) assert result["statusCode"] == 200 - def test_account_devices_dispatches(self): + def test_account_devices_dispatches(self) -> None: mgr = MagicMock() mgr.get_devices.return_value = [] route = AccountDevicesRoute(device_manager=mgr, middlewares=[]) @@ -205,7 +214,7 @@ def test_account_devices_dispatches(self): ) assert result["statusCode"] == 200 - def test_account_attached_clients_dispatches(self): + def test_account_attached_clients_dispatches(self) -> None: mgr = MagicMock() mgr.get_devices.return_value = [] route = AccountAttachedClientsRoute(device_manager=mgr, middlewares=[]) @@ -215,7 +224,7 @@ def test_account_attached_clients_dispatches(self): ) assert result["statusCode"] == 200 - def test_account_devices_notify_dispatches(self): + def test_account_devices_notify_dispatches(self) -> None: route = AccountDevicesNotifyRoute(middlewares=[]) result = _router(route).handler( _make_event("POST", "/v1/account/devices/notify", body="{}", hawk_uid="uid1"), diff --git a/lambda/tests/routes/auth/test_scoped_key_data.py b/lambda/tests/routes/auth/test_scoped_key_data.py index f5cb8342..be00ffdb 100644 --- a/lambda/tests/routes/auth/test_scoped_key_data.py +++ b/lambda/tests/routes/auth/test_scoped_key_data.py @@ -2,28 +2,25 @@ import json from decimal import Decimal +from typing import Optional from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.scoped_key_data import ScopedKeyDataRoute +from tests.conftest import json_body @pytest.fixture -def mock_account_manager(): - return MagicMock() - - -@pytest.fixture -def route(mock_account_manager): +def route(mock_account_manager: MagicMock) -> ScopedKeyDataRoute: return ScopedKeyDataRoute( account_manager=mock_account_manager, middlewares=[], ) -def _make_event(body=None, hawk_uid="uid1"): +def _make_event(body: Optional[str] = None, hawk_uid: str = "uid1") -> APIGatewayProxyEvent: """Build an event with hawk_uid pre-injected (middleware handled auth).""" return APIGatewayProxyEvent( { @@ -37,7 +34,9 @@ def _make_event(body=None, hawk_uid="uid1"): class TestScopedKeyData: - def test_success_returns_key_data(self, route, mock_account_manager): + def test_success_returns_key_data( + self, route: ScopedKeyDataRoute, mock_account_manager: MagicMock + ) -> None: mock_account_manager.get_account_by_uid.return_value = { "uid": "uid1", "createdAt": 1234567890, @@ -53,24 +52,26 @@ def test_success_returns_key_data(self, route, mock_account_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) scope_key = "https://identity.mozilla.com/apps/oldsync" assert scope_key in body assert body[scope_key]["identifier"] == scope_key assert body[scope_key]["keyRotationTimestamp"] == 1234567890 assert body[scope_key]["keyRotationSecret"] == "ab" * 32 - def test_missing_body_returns_400(self, route): + def test_missing_body_returns_400(self, route: ScopedKeyDataRoute) -> None: event = _make_event(body=None) response = route.handle(event) assert response.status_code == 400 - def test_invalid_json_body_returns_400(self, route): + def test_invalid_json_body_returns_400(self, route: ScopedKeyDataRoute) -> None: event = _make_event(body="not-json") response = route.handle(event) assert response.status_code == 400 - def test_account_not_found_returns_401(self, route, mock_account_manager): + def test_account_not_found_returns_401( + self, route: ScopedKeyDataRoute, mock_account_manager: MagicMock + ) -> None: mock_account_manager.get_account_by_uid.return_value = None event = _make_event( body=json.dumps({"client_id": "c", "scope": "s"}), @@ -78,7 +79,9 @@ def test_account_not_found_returns_401(self, route, mock_account_manager): response = route.handle(event) assert response.status_code == 401 - def test_handles_decimal_created_at_from_dynamodb(self, route, mock_account_manager): + def test_handles_decimal_created_at_from_dynamodb( + self, route: ScopedKeyDataRoute, mock_account_manager: MagicMock + ) -> None: mock_account_manager.get_account_by_uid.return_value = { "uid": "uid1", "createdAt": Decimal("1234567890"), @@ -94,11 +97,11 @@ def test_handles_decimal_created_at_from_dynamodb(self, route, mock_account_mana ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) scope_key = "https://identity.mozilla.com/apps/oldsync" assert body[scope_key]["keyRotationTimestamp"] == 1234567890.0 - def test_missing_scope_returns_400(self, route): + def test_missing_scope_returns_400(self, route: ScopedKeyDataRoute) -> None: event = _make_event( body=json.dumps({"client_id": "c"}), ) @@ -107,7 +110,7 @@ def test_missing_scope_returns_400(self, route): class TestScopedKeyDataBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: ScopedKeyDataRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_session_destroy.py b/lambda/tests/routes/auth/test_session_destroy.py index 140d4360..40a20b77 100644 --- a/lambda/tests/routes/auth/test_session_destroy.py +++ b/lambda/tests/routes/auth/test_session_destroy.py @@ -1,26 +1,23 @@ """Unit tests for SessionDestroy route""" -import json from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.session_destroy import SessionDestroyRoute +from tests.conftest import json_body @pytest.fixture -def mock_token_manager(): - return MagicMock() - - -@pytest.fixture -def route(mock_token_manager): +def route(mock_token_manager: MagicMock) -> SessionDestroyRoute: return SessionDestroyRoute(token_manager=mock_token_manager, middlewares=[]) class TestSessionDestroy: - def test_success_deletes_session(self, route, mock_token_manager): + def test_success_deletes_session( + self, route: SessionDestroyRoute, mock_token_manager: MagicMock + ) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -32,13 +29,13 @@ def test_success_deletes_session(self, route, mock_token_manager): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body == {} mock_token_manager.delete_session.assert_called_once_with("tokenid") class TestSessionDestroyBind: - def test_bind_registers_post_route(self, route): + def test_bind_registers_post_route(self, route: SessionDestroyRoute) -> None: mock_api = MagicMock() mock_api.post = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/auth/test_session_status.py b/lambda/tests/routes/auth/test_session_status.py index 0cfdc4a3..33874f3b 100644 --- a/lambda/tests/routes/auth/test_session_status.py +++ b/lambda/tests/routes/auth/test_session_status.py @@ -1,21 +1,21 @@ """Unit tests for SessionStatus route""" -import json from unittest.mock import MagicMock import pytest from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.auth.session_status import SessionStatusRoute +from tests.conftest import json_body @pytest.fixture -def route(): +def route() -> SessionStatusRoute: return SessionStatusRoute(middlewares=[]) class TestSessionStatus: - def test_success_returns_state_and_uid(self, route): + def test_success_returns_state_and_uid(self, route: SessionStatusRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -27,13 +27,13 @@ def test_success_returns_state_and_uid(self, route): ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["state"] == "verified" assert body["uid"] == "uid1" class TestSessionStatusBind: - def test_bind_registers_get_route(self, route): + def test_bind_registers_get_route(self, route: SessionStatusRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/profile/test_get_profile.py b/lambda/tests/routes/profile/test_get_profile.py index ed98ec0e..ebf8897a 100644 --- a/lambda/tests/routes/profile/test_get_profile.py +++ b/lambda/tests/routes/profile/test_get_profile.py @@ -1,6 +1,5 @@ """Unit tests for GetProfile route""" -import json from unittest.mock import MagicMock import pytest @@ -9,20 +8,21 @@ from src.routes.profile.get_profile import GetProfileRoute from src.shared.exceptions import InvalidTokenError from src.shared.oidc import OIDCTokenClaims +from tests.conftest import json_body @pytest.fixture -def mock_jwt_verifier(): +def mock_jwt_verifier() -> MagicMock: return MagicMock() @pytest.fixture -def mock_auth_account_manager(): +def mock_auth_account_manager() -> MagicMock: return MagicMock() @pytest.fixture -def route(mock_jwt_verifier, mock_auth_account_manager): +def route(mock_jwt_verifier: MagicMock, mock_auth_account_manager: MagicMock) -> GetProfileRoute: return GetProfileRoute( jwt_verifier=mock_jwt_verifier, auth_account_manager=mock_auth_account_manager, @@ -32,8 +32,11 @@ def route(mock_jwt_verifier, mock_auth_account_manager): class TestGetProfile: def test_valid_bearer_token_with_fxa_uid_returns_200( - self, route, mock_jwt_verifier, mock_auth_account_manager - ): + self, + route: GetProfileRoute, + mock_jwt_verifier: MagicMock, + mock_auth_account_manager: MagicMock, + ) -> None: mock_jwt_verifier.validate_token.return_value = OIDCTokenClaims( sub="oidc-sub-123", iss="https://auth.example.com", @@ -56,7 +59,7 @@ def test_valid_bearer_token_with_fxa_uid_returns_200( ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["email"] == "user@example.com" assert body["uid"] == "uid1" assert body["locale"] == "en-US" @@ -64,7 +67,12 @@ def test_valid_bearer_token_with_fxa_uid_returns_200( assert body["sub"] == "uid1" mock_auth_account_manager.get_account_by_uid.assert_called_once_with("uid1") - def test_fallback_to_oidc_sub_lookup(self, route, mock_jwt_verifier, mock_auth_account_manager): + def test_fallback_to_oidc_sub_lookup( + self, + route: GetProfileRoute, + mock_jwt_verifier: MagicMock, + mock_auth_account_manager: MagicMock, + ) -> None: """When fxa_uid is absent (older token), fall back to oidcSub lookup.""" mock_jwt_verifier.validate_token.return_value = OIDCTokenClaims( sub="oidc-sub-123", @@ -88,11 +96,11 @@ def test_fallback_to_oidc_sub_lookup(self, route, mock_jwt_verifier, mock_auth_a ) response = route.handle(event) assert response.status_code == 200 - body = json.loads(response.body) + body = json_body(response) assert body["uid"] == "uid1" mock_auth_account_manager.get_account_by_oidc_sub.assert_called_once_with("oidc-sub-123") - def test_missing_auth_returns_401(self, route): + def test_missing_auth_returns_401(self, route: GetProfileRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -103,10 +111,10 @@ def test_missing_auth_returns_401(self, route): ) response = route.handle(event) assert response.status_code == 401 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 110 - def test_non_bearer_auth_returns_401(self, route): + def test_non_bearer_auth_returns_401(self, route: GetProfileRoute) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -117,10 +125,12 @@ def test_non_bearer_auth_returns_401(self, route): ) response = route.handle(event) assert response.status_code == 401 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 110 - def test_invalid_jwt_returns_401(self, route, mock_jwt_verifier): + def test_invalid_jwt_returns_401( + self, route: GetProfileRoute, mock_jwt_verifier: MagicMock + ) -> None: mock_jwt_verifier.validate_token.side_effect = InvalidTokenError("expired") event = APIGatewayProxyEvent( { @@ -132,13 +142,16 @@ def test_invalid_jwt_returns_401(self, route, mock_jwt_verifier): ) response = route.handle(event) assert response.status_code == 401 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 110 assert "Invalid or expired" in body["message"] def test_account_not_found_returns_401( - self, route, mock_jwt_verifier, mock_auth_account_manager - ): + self, + route: GetProfileRoute, + mock_jwt_verifier: MagicMock, + mock_auth_account_manager: MagicMock, + ) -> None: mock_jwt_verifier.validate_token.return_value = OIDCTokenClaims( sub="uid-gone", iss="https://auth.example.com", @@ -158,13 +171,13 @@ def test_account_not_found_returns_401( ) response = route.handle(event) assert response.status_code == 401 - body = json.loads(response.body) + body = json_body(response) assert body["errno"] == 110 assert "Account not found" in body["message"] class TestGetProfileBind: - def test_bind_registers_get_route(self, route): + def test_bind_registers_get_route(self, route: GetProfileRoute) -> None: mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) route.bind(mock_api) diff --git a/lambda/tests/routes/storage/test_delete_root.py b/lambda/tests/routes/storage/test_delete_root.py index f1af2063..79aea864 100644 --- a/lambda/tests/routes/storage/test_delete_root.py +++ b/lambda/tests/routes/storage/test_delete_root.py @@ -1,15 +1,19 @@ """Tests for DeleteAllRootRoute""" import json -from typing import Any -from unittest.mock import MagicMock, patch +from typing import Any, Generator +from unittest.mock import MagicMock, Mock, patch import pytest +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent +from botocore.stub import Stubber from src.entrypoint.storage_api import lambda_handler as storage_handler +from src.environment.service_provider import ServiceProvider from src.routes.storage.delete_root import DeleteAllRootRoute from src.services.hawk_service import HawkCredentials from src.services.token_generator import TokenGenerator +from tests.conftest import json_body TEST_USER_ID = "test-user-123" TEST_GENERATION = 0 @@ -34,7 +38,7 @@ def build_storage_event(method: str, path: str, user_id: str = TEST_USER_ID) -> @pytest.fixture(autouse=True) -def mock_hawk_validate(mock_service_provider): +def mock_hawk_validate(mock_service_provider: ServiceProvider) -> Generator[None, None, None]: """Mock hawk_service.validate to bypass Hawk auth in storage handler tests.""" creds = HawkCredentials( user_id=TEST_USER_ID, @@ -49,7 +53,12 @@ def mock_hawk_validate(mock_service_provider): class TestDeleteAllRootRoute: """Tests for DeleteAllRootRoute""" - def test_handle_success(self, mock_service_provider, dynamodb_stubber, sample_lambda_context): + def test_handle_success( + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test successful deletion of all storage via root endpoint""" event = build_storage_event(method="DELETE", path="/") @@ -110,8 +119,11 @@ def test_handle_success(self, mock_service_provider, dynamodb_stubber, sample_la assert isinstance(body["modified"], (int, float)) def test_handle_with_empty_storage( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test deletion when storage is already empty""" event = build_storage_event(method="DELETE", path="/") @@ -125,8 +137,11 @@ def test_handle_with_empty_storage( assert "modified" in body def test_handle_unauthorized_missing_user_id( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test handling when hawk_uid is missing (no auth header -> middleware rejects)""" event: dict[str, Any] = { "httpMethod": "DELETE", @@ -145,8 +160,11 @@ def test_handle_unauthorized_missing_user_id( assert body["error"] == "Unauthorized" def test_root_and_storage_endpoints_behave_identically( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test that DELETE / and DELETE /storage behave the same way""" # Test DELETE / — empty storage root_event = build_storage_event(method="DELETE", path="/") @@ -171,8 +189,11 @@ def test_root_and_storage_endpoints_behave_identically( assert "modified" in storage_body def test_handle_internal_error( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test handling of internal server errors""" event = build_storage_event(method="DELETE", path="/") @@ -189,13 +210,13 @@ def test_handle_internal_error( class TestDeleteAllRootRouteUnit: """Unit tests for DeleteAllRootRoute.handle() called directly (bypassing middleware)""" - def test_missing_user_id_returns_401(self): + def test_missing_user_id_returns_401(self) -> None: """Route returns 401 when hawk_uid is not in requestContext.""" route = DeleteAllRootRoute(storage_manager=MagicMock()) event: dict = { "requestContext": {}, } - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 401 - body = json.loads(response.body) # type: ignore[arg-type] + body = json_body(response) assert body["error"] == "Unauthorized" diff --git a/lambda/tests/routes/test_bso_routes.py b/lambda/tests/routes/test_bso_routes.py index a5e79ec8..407c14fe 100644 --- a/lambda/tests/routes/test_bso_routes.py +++ b/lambda/tests/routes/test_bso_routes.py @@ -16,6 +16,7 @@ ValidationException, ) from src.shared.models import BasicStorageObject +from tests.conftest import json_body TEST_USER_ID = "test-user-123" @@ -31,7 +32,7 @@ def with_auth(event_dict: dict) -> dict: class TestReadBSORoute: """Tests for ReadBSORoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the GET route and handler works through resolver""" route = ReadBSORoute(mock_storage_manager) app = APIGatewayRestResolver() @@ -53,7 +54,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success(self, mock_storage_manager): + def test_handle_success(self, mock_storage_manager: MagicMock) -> None: """Test successful BSO retrieval""" route = ReadBSORoute(mock_storage_manager) @@ -84,8 +85,7 @@ def test_handle_success(self, mock_storage_manager): ) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["id"] == "item123" assert body["payload"] == "bookmark_data" assert body["modified"] == 1234567890.12 @@ -94,7 +94,7 @@ def test_handle_success(self, mock_storage_manager): assert "ttl" not in body assert response.headers["X-Last-Modified"] == "1234567890.12" - def test_handle_success_without_optional_fields(self, mock_storage_manager): + def test_handle_success_without_optional_fields(self, mock_storage_manager: MagicMock) -> None: """Test BSO retrieval when sortindex and ttl are None""" route = ReadBSORoute(mock_storage_manager) @@ -121,12 +121,11 @@ def test_handle_success_without_optional_fields(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "sortindex" not in body assert "ttl" not in body - def test_handle_validation_exception(self, mock_storage_manager): + def test_handle_validation_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of ValidationException""" route = ReadBSORoute(mock_storage_manager) @@ -148,11 +147,10 @@ def test_handle_validation_exception(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "error" in body - def test_handle_collection_not_found(self, mock_storage_manager): + def test_handle_collection_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of CollectionNotFoundException""" route = ReadBSORoute(mock_storage_manager) @@ -174,11 +172,10 @@ def test_handle_collection_not_found(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 404 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "error" in body - def test_handle_storage_object_not_found(self, mock_storage_manager): + def test_handle_storage_object_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of StorageObjectNotFoundException""" route = ReadBSORoute(mock_storage_manager) @@ -200,11 +197,10 @@ def test_handle_storage_object_not_found(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 404 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "error" in body - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = ReadBSORoute(mock_storage_manager) @@ -224,15 +220,14 @@ def test_handle_generic_exception(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 500 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Internal server error" class TestUpdateBSORoute: """Tests for UpdateBSORoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the PUT route and handler works through resolver""" updated_bso = BasicStorageObject( id="item123", @@ -262,7 +257,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success(self, mock_storage_manager): + def test_handle_success(self, mock_storage_manager: MagicMock) -> None: """Test successful BSO update""" route = UpdateBSORoute(mock_storage_manager) @@ -298,11 +293,10 @@ def test_handle_success(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body == 1234567891.0 - def test_handle_invalid_json(self, mock_storage_manager): + def test_handle_invalid_json(self, mock_storage_manager: MagicMock) -> None: """Test handling of invalid JSON body""" route = UpdateBSORoute(mock_storage_manager) @@ -323,7 +317,7 @@ def test_handle_invalid_json(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_non_object_body(self, mock_storage_manager): + def test_handle_non_object_body(self, mock_storage_manager: MagicMock) -> None: """Test handling of non-object JSON body (e.g. array or string)""" route = UpdateBSORoute(mock_storage_manager) @@ -344,7 +338,7 @@ def test_handle_non_object_body(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_object_id_mismatch(self, mock_storage_manager): + def test_handle_object_id_mismatch(self, mock_storage_manager: MagicMock) -> None: """Test validation when object ID doesn't match path""" route = UpdateBSORoute(mock_storage_manager) @@ -365,7 +359,7 @@ def test_handle_object_id_mismatch(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_with_precondition_header(self, mock_storage_manager): + def test_handle_with_precondition_header(self, mock_storage_manager: MagicMock) -> None: """Test with X-If-Unmodified-Since header""" route = UpdateBSORoute(mock_storage_manager) @@ -395,7 +389,7 @@ def test_handle_with_precondition_header(self, mock_storage_manager): assert response.status_code == 200 - def test_handle_invalid_precondition_header(self, mock_storage_manager): + def test_handle_invalid_precondition_header(self, mock_storage_manager: MagicMock) -> None: """Test with invalid X-If-Unmodified-Since header""" route = UpdateBSORoute(mock_storage_manager) @@ -416,7 +410,7 @@ def test_handle_invalid_precondition_header(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_collection_not_found(self, mock_storage_manager): + def test_handle_collection_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of CollectionNotFoundException""" route = UpdateBSORoute(mock_storage_manager) @@ -441,7 +435,7 @@ def test_handle_collection_not_found(self, mock_storage_manager): assert response.status_code == 404 - def test_handle_object_not_found(self, mock_storage_manager): + def test_handle_object_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of StorageObjectNotFoundException""" route = UpdateBSORoute(mock_storage_manager) @@ -466,7 +460,7 @@ def test_handle_object_not_found(self, mock_storage_manager): assert response.status_code == 404 - def test_handle_precondition_failed(self, mock_storage_manager): + def test_handle_precondition_failed(self, mock_storage_manager: MagicMock) -> None: """Test handling of PreconditionFailedException""" route = UpdateBSORoute(mock_storage_manager) @@ -491,7 +485,7 @@ def test_handle_precondition_failed(self, mock_storage_manager): assert response.status_code == 412 - def test_handle_validation_exception(self, mock_storage_manager): + def test_handle_validation_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of ValidationException""" route = UpdateBSORoute(mock_storage_manager) @@ -514,7 +508,7 @@ def test_handle_validation_exception(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = UpdateBSORoute(mock_storage_manager) @@ -541,7 +535,7 @@ def test_handle_generic_exception(self, mock_storage_manager): class TestDeleteBSORoute: """Tests for DeleteBSORoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the DELETE route and handler works through resolver""" mock_storage_manager.delete_storage_object.return_value = 1234567890.12 route = DeleteBSORoute(mock_storage_manager) @@ -563,7 +557,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success(self, mock_storage_manager): + def test_handle_success(self, mock_storage_manager: MagicMock) -> None: """Test successful BSO deletion""" route = DeleteBSORoute(mock_storage_manager) @@ -586,11 +580,10 @@ def test_handle_success(self, mock_storage_manager): "test-user-123", "bookmarks", "item123" ) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["modified"] == 1234567892.00 - def test_handle_validation_exception(self, mock_storage_manager): + def test_handle_validation_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of ValidationException""" route = DeleteBSORoute(mock_storage_manager) @@ -611,7 +604,7 @@ def test_handle_validation_exception(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_collection_not_found(self, mock_storage_manager): + def test_handle_collection_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of CollectionNotFoundException""" route = DeleteBSORoute(mock_storage_manager) @@ -634,7 +627,7 @@ def test_handle_collection_not_found(self, mock_storage_manager): assert response.status_code == 404 - def test_handle_object_not_found(self, mock_storage_manager): + def test_handle_object_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of StorageObjectNotFoundException""" route = DeleteBSORoute(mock_storage_manager) @@ -657,7 +650,7 @@ def test_handle_object_not_found(self, mock_storage_manager): assert response.status_code == 404 - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = DeleteBSORoute(mock_storage_manager) @@ -682,7 +675,7 @@ def test_handle_generic_exception(self, mock_storage_manager): class TestReadBSORouteUnauthorized: """Tests for ReadBSORoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = ReadBSORoute(mock_storage_manager) @@ -701,15 +694,14 @@ def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestDeleteBSORouteUnauthorized: """Tests for DeleteBSORoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = DeleteBSORoute(mock_storage_manager) @@ -727,15 +719,14 @@ def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestUpdateBSORouteUnauthorized: """Tests for UpdateBSORoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = UpdateBSORoute(mock_storage_manager) @@ -755,15 +746,14 @@ def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestReadBSORouteConditionalGET: """Tests for ReadBSORoute conditional GET support (Requirements 6.1-6.4)""" - def test_handle_if_modified_since_not_modified(self, mock_storage_manager): + def test_handle_if_modified_since_not_modified(self, mock_storage_manager: MagicMock) -> None: """Test 304 Not Modified when resource hasn't changed""" route = ReadBSORoute(mock_storage_manager) @@ -793,7 +783,7 @@ def test_handle_if_modified_since_not_modified(self, mock_storage_manager): assert response.status_code == 304 assert response.headers["X-Last-Modified"] == "1234567890.12" - def test_handle_if_modified_since_modified(self, mock_storage_manager): + def test_handle_if_modified_since_modified(self, mock_storage_manager: MagicMock) -> None: """Test 200 OK when resource has been modified""" route = ReadBSORoute(mock_storage_manager) @@ -821,11 +811,10 @@ def test_handle_if_modified_since_modified(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["id"] == "item123" - def test_handle_if_modified_since_invalid_format(self, mock_storage_manager): + def test_handle_if_modified_since_invalid_format(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for invalid X-If-Modified-Since header""" route = ReadBSORoute(mock_storage_manager) @@ -844,11 +833,10 @@ def test_handle_if_modified_since_invalid_format(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "Invalid X-If-Modified-Since header" in body["error"] - def test_handle_if_modified_since_negative(self, mock_storage_manager): + def test_handle_if_modified_since_negative(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for negative X-If-Modified-Since value""" route = ReadBSORoute(mock_storage_manager) @@ -867,11 +855,10 @@ def test_handle_if_modified_since_negative(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "Invalid X-If-Modified-Since header" in body["error"] - def test_handle_both_conditional_headers(self, mock_storage_manager): + def test_handle_both_conditional_headers(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request when both X-If-Modified-Since and X-If-Unmodified-Since are present""" route = ReadBSORoute(mock_storage_manager) @@ -893,15 +880,14 @@ def test_handle_both_conditional_headers(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "Cannot specify both" in body["error"] class TestReadBSORouteInvalidInputs: """Tests that validate_collection_name and validate_bso_id are called in ReadBSORoute""" - def test_handle_invalid_collection_name(self, mock_storage_manager): + def test_handle_invalid_collection_name(self, mock_storage_manager: MagicMock) -> None: """Test that invalid collection name returns 400 without calling storage""" route = ReadBSORoute(mock_storage_manager) @@ -923,7 +909,7 @@ def test_handle_invalid_collection_name(self, mock_storage_manager): assert response.status_code == 400 mock_storage_manager.get_storage_object.assert_not_called() - def test_handle_invalid_bso_id(self, mock_storage_manager): + def test_handle_invalid_bso_id(self, mock_storage_manager: MagicMock) -> None: """Test that invalid BSO ID returns 400 without calling storage""" route = ReadBSORoute(mock_storage_manager) @@ -949,7 +935,7 @@ def test_handle_invalid_bso_id(self, mock_storage_manager): class TestUpdateBSORouteInvalidCollectionName: """Tests that validate_collection_name is called before body parsing in UpdateBSORoute""" - def test_handle_invalid_collection_name(self, mock_storage_manager): + def test_handle_invalid_collection_name(self, mock_storage_manager: MagicMock) -> None: """Test that invalid collection name returns 400 without calling storage""" route = UpdateBSORoute(mock_storage_manager) @@ -976,7 +962,7 @@ def test_handle_invalid_collection_name(self, mock_storage_manager): class TestDeleteBSORouteInvalidInputs: """Tests that validate_collection_name and validate_bso_id are called in DeleteBSORoute""" - def test_handle_invalid_collection_name(self, mock_storage_manager): + def test_handle_invalid_collection_name(self, mock_storage_manager: MagicMock) -> None: """Test that invalid collection name returns 400 without calling storage""" route = DeleteBSORoute(mock_storage_manager) @@ -997,7 +983,7 @@ def test_handle_invalid_collection_name(self, mock_storage_manager): assert response.status_code == 400 mock_storage_manager.delete_storage_object.assert_not_called() - def test_handle_invalid_bso_id(self, mock_storage_manager): + def test_handle_invalid_bso_id(self, mock_storage_manager: MagicMock) -> None: """Test that invalid BSO ID returns 400 without calling storage""" route = DeleteBSORoute(mock_storage_manager) @@ -1022,7 +1008,7 @@ def test_handle_invalid_bso_id(self, mock_storage_manager): class TestUpdateBSORouteValidation: """Tests for UpdateBSORoute validation (Requirements 10.1-10.5)""" - def test_handle_payload_too_large(self, mock_storage_manager): + def test_handle_payload_too_large(self, mock_storage_manager: MagicMock) -> None: """Test 413 Request Too Large for oversized payload""" route = UpdateBSORoute(mock_storage_manager) @@ -1045,11 +1031,10 @@ def test_handle_payload_too_large(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 413 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "Payload size" in body["error"] or "payload" in body["error"].lower() - def test_handle_bso_id_too_long(self, mock_storage_manager): + def test_handle_bso_id_too_long(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for BSO ID exceeding 64 characters""" route = UpdateBSORoute(mock_storage_manager) @@ -1071,11 +1056,10 @@ def test_handle_bso_id_too_long(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "BSO ID length" in body["error"] - def test_handle_bso_id_non_printable_ascii(self, mock_storage_manager): + def test_handle_bso_id_non_printable_ascii(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for BSO ID with non-printable ASCII""" route = UpdateBSORoute(mock_storage_manager) @@ -1097,11 +1081,10 @@ def test_handle_bso_id_non_printable_ascii(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "non-printable ASCII" in body["error"] - def test_handle_sortindex_invalid(self, mock_storage_manager): + def test_handle_sortindex_invalid(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for invalid sortindex (Pydantic rejects non-int)""" route = UpdateBSORoute(mock_storage_manager) @@ -1121,11 +1104,10 @@ def test_handle_sortindex_invalid(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "sortindex" in body["error"] - def test_handle_sortindex_exceeds_max(self, mock_storage_manager): + def test_handle_sortindex_exceeds_max(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for sortindex exceeding 9 digits""" route = UpdateBSORoute(mock_storage_manager) @@ -1145,11 +1127,10 @@ def test_handle_sortindex_exceeds_max(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "sortindex" in body["error"] - def test_handle_ttl_invalid(self, mock_storage_manager): + def test_handle_ttl_invalid(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for invalid TTL (Pydantic rejects non-int)""" route = UpdateBSORoute(mock_storage_manager) @@ -1169,11 +1150,10 @@ def test_handle_ttl_invalid(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "ttl" in body["error"] - def test_handle_ttl_negative(self, mock_storage_manager): + def test_handle_ttl_negative(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for negative TTL (Pydantic enforces gt=0)""" route = UpdateBSORoute(mock_storage_manager) @@ -1193,11 +1173,10 @@ def test_handle_ttl_negative(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "ttl" in body["error"] - def test_handle_ttl_exceeds_max(self, mock_storage_manager): + def test_handle_ttl_exceeds_max(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request for TTL exceeding 9 digits (Pydantic enforces le=999999999)""" route = UpdateBSORoute(mock_storage_manager) @@ -1217,11 +1196,10 @@ def test_handle_ttl_exceeds_max(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "ttl" in body["error"] - def test_handle_none_body(self, mock_storage_manager): + def test_handle_none_body(self, mock_storage_manager: MagicMock) -> None: """Test 400 Bad Request when body is None""" route = UpdateBSORoute(mock_storage_manager) @@ -1241,11 +1219,10 @@ def test_handle_none_body(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "Invalid request body" in body["error"] - def test_handle_payload_too_large_multibyte(self, mock_storage_manager): + def test_handle_payload_too_large_multibyte(self, mock_storage_manager: MagicMock) -> None: """Test 413 for payload within char limit but exceeding byte limit (multi-byte)""" route = UpdateBSORoute(mock_storage_manager) @@ -1269,11 +1246,10 @@ def test_handle_payload_too_large_multibyte(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 413 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert "Payload size" in body["error"] - def test_handle_no_payload(self, mock_storage_manager): + def test_handle_no_payload(self, mock_storage_manager: MagicMock) -> None: """Test successful update without payload (partial update)""" route = UpdateBSORoute(mock_storage_manager) diff --git a/lambda/tests/routes/test_collection_routes.py b/lambda/tests/routes/test_collection_routes.py index 1f96c0dc..ccc5dae7 100644 --- a/lambda/tests/routes/test_collection_routes.py +++ b/lambda/tests/routes/test_collection_routes.py @@ -19,6 +19,7 @@ ValidationException, ) from src.shared.models import BasicStorageObject, BatchResult, CollectionData +from tests.conftest import json_body TEST_USER_ID = "test-user-123" AUTH_CONTEXT = {"requestContext": {"hawk_uid": TEST_USER_ID}} @@ -33,7 +34,7 @@ def with_auth(event_dict: dict) -> dict: class TestCreateCollectionRoute: """Tests for CreateCollectionRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the POST route and handler works through resolver""" route = CreateCollectionRoute(mock_storage_manager) app = APIGatewayRestResolver() @@ -51,7 +52,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 201 - def test_handle_success_with_objects(self, mock_storage_manager): + def test_handle_success_with_objects(self, mock_storage_manager: MagicMock) -> None: """Test successful collection creation with objects""" route = CreateCollectionRoute(mock_storage_manager) @@ -91,14 +92,13 @@ def test_handle_success_with_objects(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 201 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla-compliant response format assert body["modified"] == 1234567890.12 assert body["success"] == ["obj1", "obj2"] assert body["failed"] == {} - def test_handle_success_with_array_format(self, mock_storage_manager): + def test_handle_success_with_array_format(self, mock_storage_manager: MagicMock) -> None: """Test collection creation with direct array format""" route = CreateCollectionRoute(mock_storage_manager) @@ -134,7 +134,7 @@ def test_handle_success_with_array_format(self, mock_storage_manager): assert response.status_code == 201 - def test_handle_success_without_objects(self, mock_storage_manager): + def test_handle_success_without_objects(self, mock_storage_manager: MagicMock) -> None: """Test collection creation without objects""" route = CreateCollectionRoute(mock_storage_manager) @@ -164,7 +164,7 @@ def test_handle_success_without_objects(self, mock_storage_manager): assert response.status_code == 201 - def test_handle_with_precondition_header(self, mock_storage_manager): + def test_handle_with_precondition_header(self, mock_storage_manager: MagicMock) -> None: """Test handling of X-If-Unmodified-Since header""" route = CreateCollectionRoute(mock_storage_manager) @@ -190,7 +190,7 @@ def test_handle_with_precondition_header(self, mock_storage_manager): assert response.status_code == 412 - def test_handle_precondition_check_passes(self, mock_storage_manager): + def test_handle_precondition_check_passes(self, mock_storage_manager: MagicMock) -> None: """Test precondition check when collection hasn't been modified""" route = CreateCollectionRoute(mock_storage_manager) @@ -232,7 +232,7 @@ def test_handle_precondition_check_passes(self, mock_storage_manager): assert response.status_code == 201 - def test_handle_invalid_json(self, mock_storage_manager): + def test_handle_invalid_json(self, mock_storage_manager: MagicMock) -> None: """Test handling of invalid JSON in body""" route = CreateCollectionRoute(mock_storage_manager) @@ -250,7 +250,7 @@ def test_handle_invalid_json(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_validation_exception(self, mock_storage_manager): + def test_handle_validation_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of ValidationException""" route = CreateCollectionRoute(mock_storage_manager) @@ -272,7 +272,7 @@ def test_handle_validation_exception(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_conflict_exception(self, mock_storage_manager): + def test_handle_conflict_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of ConflictException""" route = CreateCollectionRoute(mock_storage_manager) @@ -292,7 +292,7 @@ def test_handle_conflict_exception(self, mock_storage_manager): assert response.status_code == 409 - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = CreateCollectionRoute(mock_storage_manager) @@ -312,7 +312,7 @@ def test_handle_generic_exception(self, mock_storage_manager): assert response.status_code == 500 - def test_handle_x_weave_records_exceeds_limit(self, mock_storage_manager): + def test_handle_x_weave_records_exceeds_limit(self, mock_storage_manager: MagicMock) -> None: """Test X-Weave-Records header exceeding limit returns 400 with code 17""" route = CreateCollectionRoute(mock_storage_manager) @@ -329,11 +329,10 @@ def test_handle_x_weave_records_exceeds_limit(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body == 17 # CODE_SERVER_LIMIT_EXCEEDED - def test_handle_x_weave_bytes_exceeds_limit(self, mock_storage_manager): + def test_handle_x_weave_bytes_exceeds_limit(self, mock_storage_manager: MagicMock) -> None: """Test X-Weave-Bytes header exceeding limit returns 400 with code 17""" route = CreateCollectionRoute(mock_storage_manager) @@ -350,11 +349,10 @@ def test_handle_x_weave_bytes_exceeds_limit(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 400 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body == 17 # CODE_SERVER_LIMIT_EXCEEDED - def test_handle_x_weave_records_invalid_format(self, mock_storage_manager): + def test_handle_x_weave_records_invalid_format(self, mock_storage_manager: MagicMock) -> None: """Test X-Weave-Records header with invalid format returns 400""" route = CreateCollectionRoute(mock_storage_manager) @@ -372,7 +370,7 @@ def test_handle_x_weave_records_invalid_format(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_x_weave_bytes_invalid_format(self, mock_storage_manager): + def test_handle_x_weave_bytes_invalid_format(self, mock_storage_manager: MagicMock) -> None: """Test X-Weave-Bytes header with invalid format returns 400""" route = CreateCollectionRoute(mock_storage_manager) @@ -390,7 +388,7 @@ def test_handle_x_weave_bytes_invalid_format(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_x_weave_records_mismatch(self, mock_storage_manager): + def test_handle_x_weave_records_mismatch(self, mock_storage_manager: MagicMock) -> None: """Test X-Weave-Records header mismatch with actual records returns 400""" route = CreateCollectionRoute(mock_storage_manager) @@ -408,7 +406,7 @@ def test_handle_x_weave_records_mismatch(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_x_weave_bytes_valid(self, mock_storage_manager): + def test_handle_x_weave_bytes_valid(self, mock_storage_manager: MagicMock) -> None: """Test X-Weave-Bytes header with valid value proceeds normally""" route = CreateCollectionRoute(mock_storage_manager) @@ -442,7 +440,7 @@ def test_handle_x_weave_bytes_valid(self, mock_storage_manager): assert response.status_code == 201 - def test_handle_x_weave_records_valid_match(self, mock_storage_manager): + def test_handle_x_weave_records_valid_match(self, mock_storage_manager: MagicMock) -> None: """Test X-Weave-Records header matching actual records proceeds normally""" route = CreateCollectionRoute(mock_storage_manager) @@ -480,7 +478,7 @@ def test_handle_x_weave_records_valid_match(self, mock_storage_manager): class TestReadCollectionRoute: """Tests for ReadCollectionRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the GET route and handler works through resolver""" # Set up mock to return empty collection objects = { @@ -507,7 +505,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_metadata_only(self, mock_storage_manager): + def test_handle_metadata_only(self, mock_storage_manager: MagicMock) -> None: """Test getting collection - returns empty list when no query params""" route = ReadCollectionRoute(mock_storage_manager) @@ -532,12 +530,11 @@ def test_handle_metadata_only(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla-compliant: returns array of IDs (empty in this case) assert body == [] - def test_handle_with_object_filters(self, mock_storage_manager): + def test_handle_with_object_filters(self, mock_storage_manager: MagicMock) -> None: """Test getting collection objects with filters""" route = ReadCollectionRoute(mock_storage_manager) @@ -574,13 +571,12 @@ def test_handle_with_object_filters(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla-compliant: returns flat array of BSO objects assert len(body) == 1 assert body[0]["id"] == "obj1" - def test_handle_objects_with_pagination(self, mock_storage_manager): + def test_handle_objects_with_pagination(self, mock_storage_manager: MagicMock) -> None: """Test getting objects with pagination""" route = ReadCollectionRoute(mock_storage_manager) @@ -614,14 +610,13 @@ def test_handle_objects_with_pagination(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla-compliant: returns flat array of IDs (default full=0) assert len(body) == 5 # Check X-Weave-Next-Offset header for pagination assert response.headers.get("X-Weave-Next-Offset") == "15" - def test_handle_objects_without_optional_fields(self, mock_storage_manager): + def test_handle_objects_without_optional_fields(self, mock_storage_manager: MagicMock) -> None: """Test formatting objects without sortindex/ttl""" route = ReadCollectionRoute(mock_storage_manager) @@ -652,13 +647,12 @@ def test_handle_objects_without_optional_fields(self, mock_storage_manager): response = route.handle(event) - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla-compliant: returns flat array of BSO objects assert "sortindex" not in body[0] assert "ttl" not in body[0] - def test_handle_validation_exception(self, mock_storage_manager): + def test_handle_validation_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of ValidationException""" route = ReadCollectionRoute(mock_storage_manager) @@ -678,7 +672,7 @@ def test_handle_validation_exception(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_collection_not_found(self, mock_storage_manager): + def test_handle_collection_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of non-existent collection - returns empty list per Mozilla spec""" route = ReadCollectionRoute(mock_storage_manager) @@ -704,11 +698,10 @@ def test_handle_collection_not_found(self, mock_storage_manager): # Should return 200 with empty list, not 404 assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body == [] - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = ReadCollectionRoute(mock_storage_manager) @@ -728,7 +721,7 @@ def test_handle_generic_exception(self, mock_storage_manager): assert response.status_code == 500 - def test_handle_conditional_get_not_modified(self, mock_storage_manager): + def test_handle_conditional_get_not_modified(self, mock_storage_manager: MagicMock) -> None: """Test conditional GET returns 304 when not modified""" route = ReadCollectionRoute(mock_storage_manager) @@ -753,7 +746,7 @@ def test_handle_conditional_get_not_modified(self, mock_storage_manager): assert response.status_code == 304 - def test_handle_conditional_get_modified(self, mock_storage_manager): + def test_handle_conditional_get_modified(self, mock_storage_manager: MagicMock) -> None: """Test conditional GET returns 200 when modified""" route = ReadCollectionRoute(mock_storage_manager) @@ -778,7 +771,9 @@ def test_handle_conditional_get_modified(self, mock_storage_manager): assert response.status_code == 200 - def test_handle_both_conditional_headers_returns_400(self, mock_storage_manager): + def test_handle_both_conditional_headers_returns_400( + self, mock_storage_manager: MagicMock + ) -> None: """Test both X-If-Modified-Since and X-If-Unmodified-Since returns 400""" route = ReadCollectionRoute(mock_storage_manager) @@ -799,7 +794,9 @@ def test_handle_both_conditional_headers_returns_400(self, mock_storage_manager) assert response.status_code == 400 - def test_handle_invalid_if_modified_since_returns_400(self, mock_storage_manager): + def test_handle_invalid_if_modified_since_returns_400( + self, mock_storage_manager: MagicMock + ) -> None: """Test invalid X-If-Modified-Since header returns 400""" route = ReadCollectionRoute(mock_storage_manager) @@ -817,7 +814,9 @@ def test_handle_invalid_if_modified_since_returns_400(self, mock_storage_manager assert response.status_code == 400 - def test_handle_negative_if_modified_since_returns_400(self, mock_storage_manager): + def test_handle_negative_if_modified_since_returns_400( + self, mock_storage_manager: MagicMock + ) -> None: """Test negative X-If-Modified-Since header returns 400""" route = ReadCollectionRoute(mock_storage_manager) @@ -835,7 +834,7 @@ def test_handle_negative_if_modified_since_returns_400(self, mock_storage_manage assert response.status_code == 400 - def test_handle_with_datetime_last_modified(self, mock_storage_manager): + def test_handle_with_datetime_last_modified(self, mock_storage_manager: MagicMock) -> None: """Test handling when last_modified is a datetime object""" route = ReadCollectionRoute(mock_storage_manager) @@ -860,7 +859,7 @@ def test_handle_with_datetime_last_modified(self, mock_storage_manager): assert response.status_code == 200 - def test_handle_with_none_last_modified(self, mock_storage_manager): + def test_handle_with_none_last_modified(self, mock_storage_manager: MagicMock) -> None: """Test handling when last_modified is None""" route = ReadCollectionRoute(mock_storage_manager) @@ -891,7 +890,7 @@ def test_handle_with_none_last_modified(self, mock_storage_manager): class TestUpdateCollectionRoute: """Tests for UpdateCollectionRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the PUT route and handler works through resolver""" collection_data = CollectionData( name="bookmarks", @@ -922,7 +921,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success(self, mock_storage_manager): + def test_handle_success(self, mock_storage_manager: MagicMock) -> None: """Test successful collection update""" route = UpdateCollectionRoute(mock_storage_manager) @@ -956,7 +955,7 @@ def test_handle_success(self, mock_storage_manager): assert response.status_code == 200 - def test_handle_invalid_json(self, mock_storage_manager): + def test_handle_invalid_json(self, mock_storage_manager: MagicMock) -> None: """Test handling of invalid JSON in body""" route = UpdateCollectionRoute(mock_storage_manager) @@ -974,7 +973,7 @@ def test_handle_invalid_json(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_direct_array_body(self, mock_storage_manager): + def test_handle_direct_array_body(self, mock_storage_manager: MagicMock) -> None: """Test batch update with direct JSON array body (SyncStorage API v1.5 format)""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1008,7 +1007,7 @@ def test_handle_direct_array_body(self, mock_storage_manager): assert response.status_code == 200 - def test_handle_missing_objects_key(self, mock_storage_manager): + def test_handle_missing_objects_key(self, mock_storage_manager: MagicMock) -> None: """Test handling of missing 'objects' key in body""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1026,7 +1025,7 @@ def test_handle_missing_objects_key(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_with_precondition_header(self, mock_storage_manager): + def test_handle_with_precondition_header(self, mock_storage_manager: MagicMock) -> None: """Test with X-If-Unmodified-Since header""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1060,7 +1059,9 @@ def test_handle_with_precondition_header(self, mock_storage_manager): assert response.status_code == 200 - def test_handle_passes_precondition_to_storage_manager(self, mock_storage_manager): + def test_handle_passes_precondition_to_storage_manager( + self, mock_storage_manager: MagicMock + ) -> None: """Test that X-If-Unmodified-Since header value is forwarded to update_collection""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1098,7 +1099,7 @@ def test_handle_passes_precondition_to_storage_manager(self, mock_storage_manage ttls=None, ) - def test_handle_invalid_precondition_header(self, mock_storage_manager): + def test_handle_invalid_precondition_header(self, mock_storage_manager: MagicMock) -> None: """Test with invalid X-If-Unmodified-Since header""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1116,7 +1117,7 @@ def test_handle_invalid_precondition_header(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_validation_exception(self, mock_storage_manager): + def test_handle_validation_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of ValidationException""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1136,7 +1137,7 @@ def test_handle_validation_exception(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_collection_not_found(self, mock_storage_manager): + def test_handle_collection_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of CollectionNotFoundException""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1158,7 +1159,7 @@ def test_handle_collection_not_found(self, mock_storage_manager): assert response.status_code == 404 - def test_handle_precondition_failed(self, mock_storage_manager): + def test_handle_precondition_failed(self, mock_storage_manager: MagicMock) -> None: """Test handling of PreconditionFailedException""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1178,7 +1179,7 @@ def test_handle_precondition_failed(self, mock_storage_manager): assert response.status_code == 412 - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1202,7 +1203,7 @@ def test_handle_generic_exception(self, mock_storage_manager): class TestDeleteCollectionRoute: """Tests for DeleteCollectionRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the DELETE route and handler works through resolver""" mock_storage_manager.delete_collection.return_value = 1234567890.12 route = DeleteCollectionRoute(mock_storage_manager) @@ -1222,7 +1223,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success(self, mock_storage_manager): + def test_handle_success(self, mock_storage_manager: MagicMock) -> None: """Test successful collection deletion""" route = DeleteCollectionRoute(mock_storage_manager) @@ -1241,11 +1242,10 @@ def test_handle_success(self, mock_storage_manager): mock_storage_manager.delete_collection.assert_called_once_with(TEST_USER_ID, "bookmarks") assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["modified"] == 1234567892.00 - def test_handle_validation_exception(self, mock_storage_manager): + def test_handle_validation_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of ValidationException""" route = DeleteCollectionRoute(mock_storage_manager) @@ -1264,7 +1264,7 @@ def test_handle_validation_exception(self, mock_storage_manager): assert response.status_code == 400 - def test_handle_collection_not_found(self, mock_storage_manager): + def test_handle_collection_not_found(self, mock_storage_manager: MagicMock) -> None: """Test handling of CollectionNotFoundException""" route = DeleteCollectionRoute(mock_storage_manager) @@ -1285,7 +1285,7 @@ def test_handle_collection_not_found(self, mock_storage_manager): assert response.status_code == 404 - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = DeleteCollectionRoute(mock_storage_manager) @@ -1304,7 +1304,7 @@ def test_handle_generic_exception(self, mock_storage_manager): assert response.status_code == 500 - def test_handle_selective_deletion(self, mock_storage_manager): + def test_handle_selective_deletion(self, mock_storage_manager: MagicMock) -> None: """Test selective deletion with ids parameter""" route = DeleteCollectionRoute(mock_storage_manager) @@ -1325,15 +1325,14 @@ def test_handle_selective_deletion(self, mock_storage_manager): TEST_USER_ID, "bookmarks", ["obj1", "obj2", "obj3"] ) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["modified"] == 1234567892.00 class TestListCollectionsRoute: """Tests for ListCollectionsRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the GET route and handler works through resolver""" mock_storage_manager.list_collections.return_value = [] route = ListCollectionsRoute(mock_storage_manager) @@ -1353,7 +1352,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success(self, mock_storage_manager): + def test_handle_success(self, mock_storage_manager: MagicMock) -> None: """Test successful collection listing""" route = ListCollectionsRoute(mock_storage_manager) @@ -1384,12 +1383,11 @@ def test_handle_success(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert len(body["collections"]) == 2 assert body["collections"][0]["name"] == "bookmarks" - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = ListCollectionsRoute(mock_storage_manager) @@ -1411,7 +1409,7 @@ def test_handle_generic_exception(self, mock_storage_manager): class TestCreateCollectionRouteUnauthorized: """Tests for CreateCollectionRoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = CreateCollectionRoute(mock_storage_manager) @@ -1427,15 +1425,14 @@ def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestDeleteCollectionRouteUnauthorized: """Tests for DeleteCollectionRoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = DeleteCollectionRoute(mock_storage_manager) @@ -1449,15 +1446,14 @@ def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestListCollectionsRouteUnauthorized: """Tests for ListCollectionsRoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = ListCollectionsRoute(mock_storage_manager) @@ -1470,15 +1466,14 @@ def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestReadCollectionRouteUnauthorized: """Tests for ReadCollectionRoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = ReadCollectionRoute(mock_storage_manager) @@ -1492,15 +1487,14 @@ def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestUpdateCollectionRouteUnauthorized: """Tests for UpdateCollectionRoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1516,15 +1510,14 @@ def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): response = route.handle(event) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestCreateCollectionRouteInvalidCollectionName: """Tests that validate_collection_name is called before storage in CreateCollectionRoute""" - def test_handle_invalid_collection_name(self, mock_storage_manager): + def test_handle_invalid_collection_name(self, mock_storage_manager: MagicMock) -> None: """Test that invalid collection name returns 400 without calling storage""" route = CreateCollectionRoute(mock_storage_manager) @@ -1547,7 +1540,7 @@ def test_handle_invalid_collection_name(self, mock_storage_manager): class TestReadCollectionRouteInvalidCollectionName: """Tests that validate_collection_name is called before storage in ReadCollectionRoute""" - def test_handle_invalid_collection_name(self, mock_storage_manager): + def test_handle_invalid_collection_name(self, mock_storage_manager: MagicMock) -> None: """Test that invalid collection name returns 400 without calling storage""" route = ReadCollectionRoute(mock_storage_manager) @@ -1570,7 +1563,7 @@ def test_handle_invalid_collection_name(self, mock_storage_manager): class TestUpdateCollectionRouteInvalidCollectionName: """Tests that validate_collection_name is called before storage in UpdateCollectionRoute""" - def test_handle_invalid_collection_name(self, mock_storage_manager): + def test_handle_invalid_collection_name(self, mock_storage_manager: MagicMock) -> None: """Test that invalid collection name returns 400 without calling storage""" route = UpdateCollectionRoute(mock_storage_manager) @@ -1593,7 +1586,7 @@ def test_handle_invalid_collection_name(self, mock_storage_manager): class TestDeleteCollectionRouteInvalidCollectionName: """Tests that validate_collection_name is called before storage in DeleteCollectionRoute""" - def test_handle_invalid_collection_name(self, mock_storage_manager): + def test_handle_invalid_collection_name(self, mock_storage_manager: MagicMock) -> None: """Test that invalid collection name returns 400 without calling storage""" route = DeleteCollectionRoute(mock_storage_manager) diff --git a/lambda/tests/routes/test_info_routes.py b/lambda/tests/routes/test_info_routes.py index e248b6f7..8c467b60 100644 --- a/lambda/tests/routes/test_info_routes.py +++ b/lambda/tests/routes/test_info_routes.py @@ -1,10 +1,10 @@ """Tests for info route handlers""" -import json from typing import Any from unittest.mock import MagicMock from aws_lambda_powertools.event_handler import APIGatewayRestResolver +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent from src.routes.info.read_collections import ReadCollectionsInfoRoute from src.routes.info.read_configuration import ReadConfigurationRoute @@ -12,6 +12,7 @@ from src.routes.info.read_quota import ReadQuotaInfoRoute from src.routes.info.read_usage import ReadCollectionUsageRoute from src.shared.models import CollectionData +from tests.conftest import json_body TEST_USER_ID = "test-user-123" @@ -27,7 +28,7 @@ def with_auth(event_dict: dict) -> dict: class TestReadCollectionsInfoRoute: """Tests for ReadCollectionsInfoRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the GET route and handler works through resolver""" mock_storage_manager.list_collections.return_value = [] route = ReadCollectionsInfoRoute(mock_storage_manager) @@ -46,7 +47,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success_mozilla_format(self, mock_storage_manager): + def test_handle_success_mozilla_format(self, mock_storage_manager: MagicMock) -> None: """Test successful retrieval of collections info in Mozilla format (name -> timestamp)""" route = ReadCollectionsInfoRoute(mock_storage_manager) @@ -74,13 +75,12 @@ def test_handle_success_mozilla_format(self, mock_storage_manager): ] mock_storage_manager.list_collections.return_value = collections - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) mock_storage_manager.list_collections.assert_called_once() assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: object mapping collection names to timestamps assert body == { "bookmarks": 1234567890.12, @@ -88,7 +88,7 @@ def test_handle_success_mozilla_format(self, mock_storage_manager): "tabs": 1234567870.00, } - def test_handle_empty_collections(self, mock_storage_manager): + def test_handle_empty_collections(self, mock_storage_manager: MagicMock) -> None: """Test handling when no collections exist""" route = ReadCollectionsInfoRoute(mock_storage_manager) @@ -96,15 +96,14 @@ def test_handle_empty_collections(self, mock_storage_manager): mock_storage_manager.list_collections.return_value = [] - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: empty object assert body == {} - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = ReadCollectionsInfoRoute(mock_storage_manager) @@ -112,18 +111,17 @@ def test_handle_generic_exception(self, mock_storage_manager): mock_storage_manager.list_collections.side_effect = Exception("Database error") - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 500 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Internal server error" class TestReadCollectionCountsRoute: """Tests for ReadCollectionCountsRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the GET route and handler works through resolver""" mock_storage_manager.list_collections.return_value = [] route = ReadCollectionCountsRoute(mock_storage_manager) @@ -142,7 +140,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success_mozilla_format(self, mock_storage_manager): + def test_handle_success_mozilla_format(self, mock_storage_manager: MagicMock) -> None: """Test successful retrieval of collection counts in Mozilla format (name -> count)""" route = ReadCollectionCountsRoute(mock_storage_manager) @@ -170,15 +168,14 @@ def test_handle_success_mozilla_format(self, mock_storage_manager): ] mock_storage_manager.list_collections.return_value = collections - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: object mapping collection names to counts directly assert body == {"bookmarks": 15, "history": 100, "tabs": 7} - def test_handle_empty_collections(self, mock_storage_manager): + def test_handle_empty_collections(self, mock_storage_manager: MagicMock) -> None: """Test handling when no collections exist""" route = ReadCollectionCountsRoute(mock_storage_manager) @@ -186,15 +183,14 @@ def test_handle_empty_collections(self, mock_storage_manager): mock_storage_manager.list_collections.return_value = [] - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: empty object assert body == {} - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = ReadCollectionCountsRoute(mock_storage_manager) @@ -202,7 +198,7 @@ def test_handle_generic_exception(self, mock_storage_manager): mock_storage_manager.list_collections.side_effect = Exception("Error") - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 500 @@ -210,7 +206,7 @@ def test_handle_generic_exception(self, mock_storage_manager): class TestReadCollectionUsageRoute: """Tests for ReadCollectionUsageRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the GET route and handler works through resolver""" mock_storage_manager.list_collections.return_value = [] route = ReadCollectionUsageRoute(mock_storage_manager) @@ -229,7 +225,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success_mozilla_format(self, mock_storage_manager): + def test_handle_success_mozilla_format(self, mock_storage_manager: MagicMock) -> None: """Test successful retrieval of collection usage in Mozilla format (name -> usage in KB)""" route = ReadCollectionUsageRoute(mock_storage_manager) @@ -257,15 +253,14 @@ def test_handle_success_mozilla_format(self, mock_storage_manager): ] mock_storage_manager.list_collections.return_value = collections - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: object mapping collection names to usage in KB (not bytes) assert body == {"bookmarks": 1.0, "history": 4.0, "tabs": 0.5} - def test_handle_empty_collections(self, mock_storage_manager): + def test_handle_empty_collections(self, mock_storage_manager: MagicMock) -> None: """Test handling when no collections exist""" route = ReadCollectionUsageRoute(mock_storage_manager) @@ -273,15 +268,14 @@ def test_handle_empty_collections(self, mock_storage_manager): mock_storage_manager.list_collections.return_value = [] - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: empty object assert body == {} - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = ReadCollectionUsageRoute(mock_storage_manager) @@ -289,7 +283,7 @@ def test_handle_generic_exception(self, mock_storage_manager): mock_storage_manager.list_collections.side_effect = Exception("Error") - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 500 @@ -297,7 +291,7 @@ def test_handle_generic_exception(self, mock_storage_manager): class TestReadQuotaInfoRoute: """Tests for ReadQuotaInfoRoute""" - def test_bind_registers_route(self, mock_storage_manager): + def test_bind_registers_route(self, mock_storage_manager: MagicMock) -> None: """Test that bind registers the GET route and handler works through resolver""" mock_storage_manager.list_collections.return_value = [] route = ReadQuotaInfoRoute(mock_storage_manager) @@ -316,7 +310,7 @@ def test_bind_registers_route(self, mock_storage_manager): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_success_mozilla_format(self, mock_storage_manager): + def test_handle_success_mozilla_format(self, mock_storage_manager: MagicMock) -> None: """Test successful retrieval of quota information in Mozilla format [usage_kb, quota_kb]""" route = ReadQuotaInfoRoute(mock_storage_manager) @@ -344,11 +338,10 @@ def test_handle_success_mozilla_format(self, mock_storage_manager): ] mock_storage_manager.list_collections.return_value = collections - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: [usage_kb, quota_kb or null] assert isinstance(body, list) @@ -358,7 +351,7 @@ def test_handle_success_mozilla_format(self, mock_storage_manager): # Default quota is None (not enforced) assert body[1] is None - def test_handle_with_quota_limit(self, mock_storage_manager): + def test_handle_with_quota_limit(self, mock_storage_manager: MagicMock) -> None: """Test quota info with a configured quota limit""" quota_kb = 10240 # 10 MB in KB route = ReadQuotaInfoRoute(mock_storage_manager, quota_kb=quota_kb) @@ -375,17 +368,16 @@ def test_handle_with_quota_limit(self, mock_storage_manager): ] mock_storage_manager.list_collections.return_value = collections - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: [usage_kb, quota_kb] assert body[0] == 2.0 # 2048 bytes = 2 KB assert body[1] == 10240 # Configured quota - def test_handle_no_collections(self, mock_storage_manager): + def test_handle_no_collections(self, mock_storage_manager: MagicMock) -> None: """Test quota info when no collections exist""" route = ReadQuotaInfoRoute(mock_storage_manager) @@ -393,11 +385,10 @@ def test_handle_no_collections(self, mock_storage_manager): mock_storage_manager.list_collections.return_value = [] - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: [usage_kb, quota_kb or null] assert isinstance(body, list) @@ -405,7 +396,7 @@ def test_handle_no_collections(self, mock_storage_manager): assert body[0] == 0.0 # No usage assert body[1] is None # No quota enforced - def test_handle_single_collection(self, mock_storage_manager): + def test_handle_single_collection(self, mock_storage_manager: MagicMock) -> None: """Test quota info with single collection""" route = ReadQuotaInfoRoute(mock_storage_manager) @@ -421,17 +412,16 @@ def test_handle_single_collection(self, mock_storage_manager): ] mock_storage_manager.list_collections.return_value = collections - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Mozilla format: [usage_kb, quota_kb or null] assert body[0] == 5.0 # 5120 bytes = 5 KB assert body[1] is None - def test_handle_generic_exception(self, mock_storage_manager): + def test_handle_generic_exception(self, mock_storage_manager: MagicMock) -> None: """Test handling of generic exceptions""" route = ReadQuotaInfoRoute(mock_storage_manager) @@ -439,79 +429,75 @@ def test_handle_generic_exception(self, mock_storage_manager): mock_storage_manager.list_collections.side_effect = Exception("Error") - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 500 - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = ReadQuotaInfoRoute(mock_storage_manager) event: dict[str, Any] = {"requestContext": {}} - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestReadCollectionsInfoRouteUnauthorized: """Tests for ReadCollectionsInfoRoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = ReadCollectionsInfoRoute(mock_storage_manager) event: dict[str, Any] = {"requestContext": {}} - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestReadCollectionCountsRouteUnauthorized: """Tests for ReadCollectionCountsRoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = ReadCollectionCountsRoute(mock_storage_manager) event: dict[str, Any] = {"requestContext": {}} - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestReadCollectionUsageRouteUnauthorized: """Tests for ReadCollectionUsageRoute unauthorized cases""" - def test_handle_unauthorized_missing_user_id(self, mock_storage_manager): + def test_handle_unauthorized_missing_user_id(self, mock_storage_manager: MagicMock) -> None: """Test handling when user_id is missing from authorizer context""" route = ReadCollectionUsageRoute(mock_storage_manager) event: dict[str, Any] = {"requestContext": {}} - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 401 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "Unauthorized" class TestReadConfigurationRoute: """Tests for ReadConfigurationRoute""" - def test_bind_registers_route(self): + def test_bind_registers_route(self) -> None: """Test that bind registers the GET route and handler works through resolver""" route = ReadConfigurationRoute() app = APIGatewayRestResolver() @@ -527,17 +513,16 @@ def test_bind_registers_route(self): result = app.resolve(event, MagicMock()) assert result["statusCode"] == 200 - def test_handle_default_configuration(self): + def test_handle_default_configuration(self) -> None: """Test successful retrieval of default server configuration""" route = ReadConfigurationRoute() event: dict[str, Any] = {} - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Required fields per Mozilla spec assert body["max_request_bytes"] == 2 * 1024 * 1024 # 2 MB @@ -549,7 +534,7 @@ def test_handle_default_configuration(self): assert "max_total_records" not in body assert "max_total_bytes" not in body - def test_handle_custom_configuration(self): + def test_handle_custom_configuration(self) -> None: """Test configuration with custom limits""" route = ReadConfigurationRoute( max_request_bytes=1024 * 1024, # 1 MB @@ -562,11 +547,10 @@ def test_handle_custom_configuration(self): event: dict[str, Any] = {} - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["max_request_bytes"] == 1024 * 1024 assert body["max_post_records"] == 50 @@ -575,7 +559,7 @@ def test_handle_custom_configuration(self): assert body["max_total_records"] == 1000 assert body["max_total_bytes"] == 10 * 1024 * 1024 - def test_handle_partial_optional_configuration(self): + def test_handle_partial_optional_configuration(self) -> None: """Test configuration with only some optional limits""" route = ReadConfigurationRoute( max_total_records=500, @@ -584,11 +568,10 @@ def test_handle_partial_optional_configuration(self): event: dict[str, Any] = {} - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 200 - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) # Required fields present assert "max_request_bytes" in body diff --git a/lambda/tests/routes/test_storage_routes.py b/lambda/tests/routes/test_storage_routes.py index ff1309b1..9dfb494d 100644 --- a/lambda/tests/routes/test_storage_routes.py +++ b/lambda/tests/routes/test_storage_routes.py @@ -1,15 +1,19 @@ """Tests for storage route handlers""" import json -from typing import Any -from unittest.mock import MagicMock, patch +from typing import Any, Generator +from unittest.mock import MagicMock, Mock, patch import pytest +from aws_lambda_powertools.utilities.data_classes import APIGatewayProxyEvent +from botocore.stub import Stubber from src.entrypoint.storage_api import lambda_handler as storage_handler +from src.environment.service_provider import ServiceProvider from src.routes.storage.delete_all import DeleteAllStorageRoute from src.services.hawk_service import HawkCredentials from src.services.token_generator import TokenGenerator +from tests.conftest import json_body TEST_USER_ID = "test-user-123" TEST_GENERATION = 0 @@ -34,7 +38,7 @@ def build_storage_event(method: str, path: str, user_id: str = TEST_USER_ID) -> @pytest.fixture(autouse=True) -def mock_hawk_validate(mock_service_provider): +def mock_hawk_validate(mock_service_provider: ServiceProvider) -> Generator[None, None, None]: """Mock hawk_service.validate to bypass Hawk auth in storage handler tests.""" creds = HawkCredentials( user_id=TEST_USER_ID, @@ -49,7 +53,12 @@ def mock_hawk_validate(mock_service_provider): class TestDeleteAllStorageRoute: """Tests for DeleteAllStorageRoute""" - def test_handle_success(self, mock_service_provider, dynamodb_stubber, sample_lambda_context): + def test_handle_success( + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test successful deletion of all storage""" event = build_storage_event(method="DELETE", path="/storage") @@ -114,8 +123,11 @@ def test_handle_success(self, mock_service_provider, dynamodb_stubber, sample_la assert isinstance(body["modified"], (int, float)) def test_handle_with_empty_storage( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test deletion when storage is already empty""" event = build_storage_event(method="DELETE", path="/storage") @@ -129,8 +141,11 @@ def test_handle_with_empty_storage( assert "modified" in body def test_handle_with_pagination( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test deletion with multiple collections (GSI pagination)""" event = build_storage_event(method="DELETE", path="/storage") @@ -238,8 +253,11 @@ def test_handle_with_pagination( assert "modified" in body def test_handle_unauthorized_missing_user_id( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test handling when hawk_uid is missing (no auth header -> middleware rejects)""" event: dict[str, Any] = { "httpMethod": "DELETE", @@ -258,8 +276,11 @@ def test_handle_unauthorized_missing_user_id( assert body["error"] == "Unauthorized" def test_handle_internal_error( - self, mock_service_provider, dynamodb_stubber, sample_lambda_context - ): + self, + mock_service_provider: ServiceProvider, + dynamodb_stubber: Stubber, + sample_lambda_context: Mock, + ) -> None: """Test handling of internal server errors""" event = build_storage_event(method="DELETE", path="/storage") @@ -276,13 +297,13 @@ def test_handle_internal_error( class TestDeleteAllStorageRouteUnit: """Unit tests for DeleteAllStorageRoute.handle() called directly (bypassing middleware)""" - def test_missing_user_id_returns_401(self): + def test_missing_user_id_returns_401(self) -> None: """Route returns 401 when hawk_uid is not in requestContext.""" route = DeleteAllStorageRoute(storage_manager=MagicMock()) event: dict = { "requestContext": {}, } - response = route.handle(event) + response = route.handle(APIGatewayProxyEvent(event)) assert response.status_code == 401 - body = json.loads(response.body) # type: ignore[arg-type] + body = json_body(response) assert body["error"] == "Unauthorized" diff --git a/lambda/tests/routes/token/test_request.py b/lambda/tests/routes/token/test_request.py index e19f917b..eeecd095 100644 --- a/lambda/tests/routes/token/test_request.py +++ b/lambda/tests/routes/token/test_request.py @@ -1,6 +1,5 @@ """Unit tests for RequestTokenRoute""" -import json from http import HTTPStatus from unittest.mock import MagicMock @@ -21,25 +20,28 @@ from src.shared.oidc import OIDCTokenClaims from src.shared.token import TokenResponse from src.shared.user import UserRecord +from tests.conftest import header, json_body @pytest.fixture -def mock_oidc_validator(): +def mock_oidc_validator() -> MagicMock: return MagicMock() @pytest.fixture -def mock_user_manager(): +def mock_user_manager() -> MagicMock: return MagicMock() @pytest.fixture -def mock_token_generator(): +def mock_token_generator() -> MagicMock: return MagicMock() @pytest.fixture -def request_token_route(mock_oidc_validator, mock_user_manager, mock_token_generator): +def request_token_route( + mock_oidc_validator: MagicMock, mock_user_manager: MagicMock, mock_token_generator: MagicMock +) -> GetTokenRoute: return GetTokenRoute( oidc_validator=mock_oidc_validator, user_manager=mock_user_manager, @@ -49,7 +51,7 @@ def request_token_route(mock_oidc_validator, mock_user_manager, mock_token_gener @pytest.fixture -def valid_event(): +def valid_event() -> APIGatewayProxyEvent: """Valid GET request to token endpoint""" return APIGatewayProxyEvent( { @@ -64,7 +66,7 @@ def valid_event(): @pytest.fixture -def mock_oidc_claims(): +def mock_oidc_claims() -> OIDCTokenClaims: return OIDCTokenClaims( sub="user123", iss="https://auth.example.com", @@ -76,7 +78,7 @@ def mock_oidc_claims(): @pytest.fixture -def mock_user_record(): +def mock_user_record() -> UserRecord: return UserRecord( user_id="user123", generation=0, @@ -87,7 +89,7 @@ def mock_user_record(): @pytest.fixture -def mock_token_response(): +def mock_token_response() -> TokenResponse: return TokenResponse( id="dXNlcjEyMzowOjEyMzQ1Njc4OTA", key="a" * 64, @@ -102,8 +104,11 @@ class TestRequestTokenRouteInit: """Test RequestTokenRoute initialization""" def test_init_stores_dependencies( - self, mock_oidc_validator, mock_user_manager, mock_token_generator - ): + self, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test that dependencies are stored correctly""" route = GetTokenRoute( oidc_validator=mock_oidc_validator, @@ -119,7 +124,7 @@ def test_init_stores_dependencies( class TestRequestTokenRouteBind: """Test bind method""" - def test_bind_registers_get_route(self, request_token_route): + def test_bind_registers_get_route(self, request_token_route: GetTokenRoute) -> None: """Test that bind registers GET route""" mock_api = MagicMock() mock_api.get = MagicMock(return_value=lambda f: f) @@ -135,22 +140,25 @@ class TestRequestTokenRouteHandle: def test_handle_success( self, - request_token_route, - valid_event, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test successful token issuance""" - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.OK assert response.content_type == "application/json" - body = json.loads(response.body) + body = json_body(response) assert body["id"] == mock_token_response.id assert body["key"] == mock_token_response.key assert body["api_endpoint"] == mock_token_response.api_endpoint @@ -158,7 +166,7 @@ def test_handle_success( assert body["duration"] == 300 assert body["hashalg"] == "sha256" - def test_handle_missing_auth_header(self, request_token_route): + def test_handle_missing_auth_header(self, request_token_route: GetTokenRoute) -> None: """Test missing Authorization header returns 401""" event = APIGatewayProxyEvent( { @@ -170,11 +178,11 @@ def test_handle_missing_auth_header(self, request_token_route): response = request_token_route.handle(event) assert response.status_code == HTTPStatus.UNAUTHORIZED - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-credentials" assert "Missing Authorization header" in body["errors"][0]["description"] - def test_handle_malformed_auth_header(self, request_token_route): + def test_handle_malformed_auth_header(self, request_token_route: GetTokenRoute) -> None: """Test malformed Authorization header returns 400""" event = APIGatewayProxyEvent( { @@ -186,108 +194,128 @@ def test_handle_malformed_auth_header(self, request_token_route): response = request_token_route.handle(event) assert response.status_code == HTTPStatus.BAD_REQUEST - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-request" assert "Malformed Authorization header" in body["errors"][0]["description"] - def test_handle_invalid_credentials_error(self, request_token_route, valid_event): + def test_handle_invalid_credentials_error( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test InvalidCredentialsError returns 401""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidCredentialsError( - "Token expired" - ) + mock_oidc_validator.validate_token.side_effect = InvalidCredentialsError("Token expired") response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.UNAUTHORIZED - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-credentials" - def test_handle_invalid_token_error(self, request_token_route, valid_event): + def test_handle_invalid_token_error( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test InvalidTokenError returns 401""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidTokenError( - "Invalid signature" - ) + mock_oidc_validator.validate_token.side_effect = InvalidTokenError("Invalid signature") response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.UNAUTHORIZED - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-credentials" - def test_handle_service_unavailable_error(self, request_token_route, valid_event): + def test_handle_service_unavailable_error( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test ServiceUnavailableError returns 503""" - request_token_route.oidc_validator.validate_token.side_effect = ServiceUnavailableError( + mock_oidc_validator.validate_token.side_effect = ServiceUnavailableError( "OIDC provider unreachable" ) response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.SERVICE_UNAVAILABLE - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "service-unavailable" - def test_handle_validation_exception(self, request_token_route, valid_event): + def test_handle_validation_exception( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test ValidationException returns 400""" - request_token_route.oidc_validator.validate_token.side_effect = ValidationException( + mock_oidc_validator.validate_token.side_effect = ValidationException( "Invalid request format" ) response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.BAD_REQUEST - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-request" - def test_handle_unexpected_error(self, request_token_route, valid_event): + def test_handle_unexpected_error( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test unexpected error returns 500""" - request_token_route.oidc_validator.validate_token.side_effect = RuntimeError( - "Unexpected error" - ) + mock_oidc_validator.validate_token.side_effect = RuntimeError("Unexpected error") response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.INTERNAL_SERVER_ERROR - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "internal-error" def test_handle_calls_services_in_order( self, - request_token_route, - valid_event, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test that services are called in correct order""" - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response # Configure generate_uid to return a known value expected_uid = 7351813628096158130 - request_token_route.token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_uid.return_value = expected_uid request_token_route.handle(valid_event) # Verify OIDC validator was called with the token - request_token_route.oidc_validator.validate_token.assert_called_once_with( - "valid-oidc-token" - ) + mock_oidc_validator.validate_token.assert_called_once_with("valid-oidc-token") # Verify generate_uid was called with user_id from OIDC claims - request_token_route.token_generator.generate_uid.assert_called_once_with("user123", 0) + mock_token_generator.generate_uid.assert_called_once_with("user123", 0) # Verify user manager was called with uid and client_state (empty string default) - request_token_route.user_manager.get_or_create_user.assert_called_once_with("user123", "") + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "") # Verify token generator was called with user_id, uid, and generation - request_token_route.token_generator.generate_token.assert_called_once_with( + mock_token_generator.generate_token.assert_called_once_with( user_id="user123", uid=expected_uid, generation=0, ) - def test_handle_null_headers(self, request_token_route): + def test_handle_null_headers(self, request_token_route: GetTokenRoute) -> None: """Test handling of null headers""" event = APIGatewayProxyEvent( { @@ -302,11 +330,14 @@ def test_handle_null_headers(self, request_token_route): def test_handle_request_context_identity_missing( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test handling when requestContext.identity raises KeyError""" event = APIGatewayProxyEvent( { @@ -316,9 +347,9 @@ def test_handle_request_context_identity_missing( "requestContext": {}, # Empty requestContext triggers KeyError on identity access } ) - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) @@ -326,11 +357,14 @@ def test_handle_request_context_identity_missing( def test_handle_case_insensitive_auth_header( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test that Authorization header lookup is case-insensitive""" event = APIGatewayProxyEvent( { @@ -339,9 +373,9 @@ def test_handle_case_insensitive_auth_header( "headers": {"Authorization": "Bearer valid-token"}, } ) - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) @@ -351,12 +385,14 @@ def test_handle_case_insensitive_auth_header( class TestExtractBearerToken: """Test _extract_bearer_token method""" - def test_extract_bearer_token_valid(self, request_token_route): + def test_extract_bearer_token_valid(self, request_token_route: GetTokenRoute) -> None: """Test extracting valid Bearer token""" token = request_token_route._extract_bearer_token("Bearer my-token-123") assert token == "my-token-123" - def test_extract_bearer_token_case_insensitive(self, request_token_route): + def test_extract_bearer_token_case_insensitive( + self, request_token_route: GetTokenRoute + ) -> None: """Test Bearer keyword is case-insensitive""" token = request_token_route._extract_bearer_token("bearer my-token") assert token == "my-token" @@ -364,13 +400,13 @@ def test_extract_bearer_token_case_insensitive(self, request_token_route): token = request_token_route._extract_bearer_token("BEARER my-token") assert token == "my-token" - def test_extract_bearer_token_with_spaces(self, request_token_route): + def test_extract_bearer_token_with_spaces(self, request_token_route: GetTokenRoute) -> None: """Test token extraction with multiple spaces after Bearer""" # The regex \s+ consumes all whitespace between Bearer and token token = request_token_route._extract_bearer_token("Bearer token-with-spaces") assert token == "token-with-spaces" - def test_extract_bearer_token_invalid_format(self, request_token_route): + def test_extract_bearer_token_invalid_format(self, request_token_route: GetTokenRoute) -> None: """Test None returned for invalid format""" token = request_token_route._extract_bearer_token("Basic dXNlcjpwYXNz") assert token is None @@ -379,7 +415,7 @@ def test_extract_bearer_token_invalid_format(self, request_token_route): class TestErrorResponse: """Test _error_response method""" - def test_error_response_structure(self, request_token_route): + def test_error_response_structure(self, request_token_route: GetTokenRoute) -> None: """Test error response has correct structure""" response = request_token_route._error_response( status_code=HTTPStatus.UNAUTHORIZED, @@ -392,7 +428,7 @@ def test_error_response_structure(self, request_token_route): assert response.status_code == HTTPStatus.UNAUTHORIZED assert response.content_type == "application/json" - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-credentials" assert len(body["errors"]) == 1 assert body["errors"][0]["location"] == "header" @@ -403,7 +439,9 @@ def test_error_response_structure(self, request_token_route): class TestContentTypeValidation: """Test Content-Type validation""" - def test_handle_invalid_content_type_returns_415(self, request_token_route): + def test_handle_invalid_content_type_returns_415( + self, request_token_route: GetTokenRoute + ) -> None: """Test invalid Content-Type returns 415""" event = APIGatewayProxyEvent( { @@ -419,17 +457,20 @@ def test_handle_invalid_content_type_returns_415(self, request_token_route): response = request_token_route.handle(event) assert response.status_code == HTTPStatus.UNSUPPORTED_MEDIA_TYPE - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "unsupported-media-type" assert "Content-Type" in body["errors"][0]["name"] def test_handle_valid_content_type_json( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test application/json Content-Type is accepted""" event = APIGatewayProxyEvent( { @@ -442,9 +483,9 @@ def test_handle_valid_content_type_json( "body": '{"some": "data"}', } ) - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) @@ -452,11 +493,14 @@ def test_handle_valid_content_type_json( def test_handle_valid_content_type_form( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test application/x-www-form-urlencoded Content-Type is accepted""" event = APIGatewayProxyEvent( { @@ -469,9 +513,9 @@ def test_handle_valid_content_type_form( "body": "key=value", } ) - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) @@ -479,11 +523,14 @@ def test_handle_valid_content_type_form( def test_handle_content_type_with_charset( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test Content-Type with charset parameter is accepted""" event = APIGatewayProxyEvent( { @@ -496,9 +543,9 @@ def test_handle_content_type_with_charset( "body": '{"some": "data"}', } ) - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) @@ -506,11 +553,14 @@ def test_handle_content_type_with_charset( def test_handle_no_body_skips_content_type_validation( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test Content-Type validation is skipped when no body""" event = APIGatewayProxyEvent( { @@ -523,9 +573,9 @@ def test_handle_no_body_skips_content_type_validation( "body": None, } ) - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) @@ -533,11 +583,14 @@ def test_handle_no_body_skips_content_type_validation( def test_handle_empty_body_skips_content_type_validation( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test Content-Type validation is skipped when body is empty string""" event = APIGatewayProxyEvent( { @@ -550,9 +603,9 @@ def test_handle_empty_body_skips_content_type_validation( "body": "", } ) - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) @@ -562,14 +615,14 @@ def test_handle_empty_body_skips_content_type_validation( class TestBearerTokenPattern: """Test BEARER_TOKEN_PATTERN regex""" - def test_pattern_matches_valid_bearer(self): + def test_pattern_matches_valid_bearer(self) -> None: """Test pattern matches valid Bearer tokens""" assert BEARER_TOKEN_PATTERN.match("Bearer token123") assert BEARER_TOKEN_PATTERN.match("bearer token123") assert BEARER_TOKEN_PATTERN.match("BEARER token123") assert BEARER_TOKEN_PATTERN.match("Bearer eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.xxx") - def test_pattern_rejects_invalid_formats(self): + def test_pattern_rejects_invalid_formats(self) -> None: """Test pattern rejects invalid formats""" assert not BEARER_TOKEN_PATTERN.match("Basic dXNlcjpwYXNz") assert not BEARER_TOKEN_PATTERN.match("token123") @@ -582,11 +635,14 @@ class TestXClientStateHeader: def test_valid_client_state_passed_to_user_manager( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test valid X-Client-State is passed to user manager""" event = APIGatewayProxyEvent( { @@ -599,25 +655,26 @@ def test_valid_client_state_passed_to_user_manager( } ) expected_uid = 7351813628096158130 - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_uid.return_value = expected_uid - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) assert response.status_code == HTTPStatus.OK - request_token_route.user_manager.get_or_create_user.assert_called_once_with( - "user123", "abcdef123456" - ) + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "abcdef123456") def test_client_state_case_insensitive_header( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test X-Client-State header lookup is case-insensitive""" event = APIGatewayProxyEvent( { @@ -630,25 +687,26 @@ def test_client_state_case_insensitive_header( } ) expected_uid = 7351813628096158130 - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_uid.return_value = expected_uid - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) assert response.status_code == HTTPStatus.OK - request_token_route.user_manager.get_or_create_user.assert_called_once_with( - "user123", "ABCDEF" - ) + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "ABCDEF") def test_missing_client_state_defaults_to_empty_string( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test missing X-Client-State defaults to empty string""" event = APIGatewayProxyEvent( { @@ -660,17 +718,19 @@ def test_missing_client_state_defaults_to_empty_string( } ) expected_uid = 7351813628096158130 - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_uid.return_value = expected_uid - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) assert response.status_code == HTTPStatus.OK - request_token_route.user_manager.get_or_create_user.assert_called_once_with("user123", "") + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "") - def test_invalid_client_state_special_chars_returns_400(self, request_token_route): + def test_invalid_client_state_special_chars_returns_400( + self, request_token_route: GetTokenRoute + ) -> None: """Test X-Client-State with invalid special characters returns 400""" event = APIGatewayProxyEvent( { @@ -685,12 +745,14 @@ def test_invalid_client_state_special_chars_returns_400(self, request_token_rout response = request_token_route.handle(event) assert response.status_code == HTTPStatus.BAD_REQUEST - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-request" assert body["errors"][0]["name"] == "X-Client-State" assert "urlsafe-base64" in body["errors"][0]["description"].lower() - def test_invalid_client_state_too_long_returns_400(self, request_token_route): + def test_invalid_client_state_too_long_returns_400( + self, request_token_route: GetTokenRoute + ) -> None: """Test X-Client-State longer than 32 chars returns 400""" event = APIGatewayProxyEvent( { @@ -705,17 +767,20 @@ def test_invalid_client_state_too_long_returns_400(self, request_token_route): response = request_token_route.handle(event) assert response.status_code == HTTPStatus.BAD_REQUEST - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-request" assert body["errors"][0]["name"] == "X-Client-State" def test_valid_client_state_max_length( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test X-Client-State at max length (32 chars) is accepted""" event = APIGatewayProxyEvent( { @@ -727,9 +792,9 @@ def test_valid_client_state_max_length( }, } ) - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) @@ -737,11 +802,14 @@ def test_valid_client_state_max_length( def test_valid_client_state_with_underscore( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test X-Client-State with underscore is accepted (urlsafe-base64)""" event = APIGatewayProxyEvent( { @@ -754,25 +822,26 @@ def test_valid_client_state_with_underscore( } ) expected_uid = 7351813628096158130 - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_uid.return_value = expected_uid - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) assert response.status_code == HTTPStatus.OK - request_token_route.user_manager.get_or_create_user.assert_called_once_with( - "user123", "abc_def_123" - ) + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "abc_def_123") def test_valid_client_state_with_hyphen( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test X-Client-State with hyphen is accepted (urlsafe-base64)""" event = APIGatewayProxyEvent( { @@ -785,25 +854,26 @@ def test_valid_client_state_with_hyphen( } ) expected_uid = 7351813628096158130 - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_uid.return_value = expected_uid - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) assert response.status_code == HTTPStatus.OK - request_token_route.user_manager.get_or_create_user.assert_called_once_with( - "user123", "abc-def-123" - ) + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "abc-def-123") def test_valid_client_state_with_period( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test X-Client-State with period is accepted (urlsafe-base64 + period)""" event = APIGatewayProxyEvent( { @@ -816,25 +886,26 @@ def test_valid_client_state_with_period( } ) expected_uid = 7351813628096158130 - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_uid.return_value = expected_uid - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) assert response.status_code == HTTPStatus.OK - request_token_route.user_manager.get_or_create_user.assert_called_once_with( - "user123", "abc.def.123" - ) + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "abc.def.123") def test_valid_client_state_mixed_urlsafe_chars( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test X-Client-State with mixed urlsafe-base64 + period chars is accepted""" event = APIGatewayProxyEvent( { @@ -847,25 +918,26 @@ def test_valid_client_state_mixed_urlsafe_chars( } ) expected_uid = 7351813628096158130 - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_uid.return_value = expected_uid - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) assert response.status_code == HTTPStatus.OK - request_token_route.user_manager.get_or_create_user.assert_called_once_with( - "user123", "aB3_xY-z.9Q" - ) + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "aB3_xY-z.9Q") def test_empty_client_state_is_valid( self, - request_token_route, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test empty X-Client-State header value is valid""" event = APIGatewayProxyEvent( { @@ -878,15 +950,15 @@ def test_empty_client_state_is_valid( } ) expected_uid = 7351813628096158130 - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_uid.return_value = expected_uid - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_uid.return_value = expected_uid + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(event) assert response.status_code == HTTPStatus.OK - request_token_route.user_manager.get_or_create_user.assert_called_once_with("user123", "") + mock_user_manager.get_or_create_user.assert_called_once_with("user123", "") class TestXTimestampHeader: @@ -894,26 +966,31 @@ class TestXTimestampHeader: def test_success_response_includes_timestamp_header( self, - request_token_route, - valid_event, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test successful response includes X-Timestamp header""" - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.OK assert "X-Timestamp" in response.headers # Verify it's a valid integer timestamp - timestamp = int(response.headers["X-Timestamp"]) + timestamp = int(header(response, "X-Timestamp")) assert timestamp > 0 - def test_error_response_includes_timestamp_header(self, request_token_route): + def test_error_response_includes_timestamp_header( + self, request_token_route: GetTokenRoute + ) -> None: """Test error response includes X-Timestamp header""" event = APIGatewayProxyEvent( { @@ -927,30 +1004,35 @@ def test_error_response_includes_timestamp_header(self, request_token_route): assert response.status_code == HTTPStatus.UNAUTHORIZED assert "X-Timestamp" in response.headers # Verify it's a valid integer timestamp - timestamp = int(response.headers["X-Timestamp"]) + timestamp = int(header(response, "X-Timestamp")) assert timestamp > 0 def test_timestamp_is_integer_format( self, - request_token_route, - valid_event, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test X-Timestamp value is an integer (no decimal)""" - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(valid_event) - timestamp_str = response.headers["X-Timestamp"] + timestamp_str = header(response, "X-Timestamp") # Should be a string representation of an integer (no decimal point) assert "." not in timestamp_str assert timestamp_str.isdigit() - def test_validation_error_includes_timestamp_header(self, request_token_route): + def test_validation_error_includes_timestamp_header( + self, request_token_route: GetTokenRoute + ) -> None: """Test validation error (400) includes X-Timestamp header""" event = APIGatewayProxyEvent( { @@ -966,9 +1048,14 @@ def test_validation_error_includes_timestamp_header(self, request_token_route): assert response.status_code == HTTPStatus.BAD_REQUEST assert "X-Timestamp" in response.headers - def test_service_unavailable_includes_timestamp_header(self, request_token_route, valid_event): + def test_service_unavailable_includes_timestamp_header( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test service unavailable (503) includes X-Timestamp header""" - request_token_route.oidc_validator.validate_token.side_effect = ServiceUnavailableError( + mock_oidc_validator.validate_token.side_effect = ServiceUnavailableError( "OIDC provider unreachable" ) @@ -981,67 +1068,92 @@ def test_service_unavailable_includes_timestamp_header(self, request_token_route class TestNewErrorStatuses: """Test new error status types per Mozilla spec""" - def test_handle_invalid_timestamp_error(self, request_token_route, valid_event): + def test_handle_invalid_timestamp_error( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test InvalidTimestampError returns 401 with invalid-timestamp status""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidTimestampError( + mock_oidc_validator.validate_token.side_effect = InvalidTimestampError( "Token timestamp differs significantly from server time" ) response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.UNAUTHORIZED - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-timestamp" assert body["errors"][0]["location"] == "header" assert body["errors"][0]["name"] == "Authorization" assert "timestamp" in body["errors"][0]["description"].lower() - def test_handle_invalid_generation_error(self, request_token_route, valid_event): + def test_handle_invalid_generation_error( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test InvalidGenerationError returns 401 with invalid-generation status""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidGenerationError( + mock_oidc_validator.validate_token.side_effect = InvalidGenerationError( "Token generation number is outdated" ) response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.UNAUTHORIZED - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-generation" assert body["errors"][0]["location"] == "header" assert body["errors"][0]["name"] == "Authorization" assert "generation" in body["errors"][0]["description"].lower() - def test_handle_invalid_client_state_error(self, request_token_route, valid_event): + def test_handle_invalid_client_state_error( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test InvalidClientStateError returns 401 with invalid-client-state status""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidClientStateError( + mock_oidc_validator.validate_token.side_effect = InvalidClientStateError( "Client state has been seen before" ) response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.UNAUTHORIZED - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "invalid-client-state" assert body["errors"][0]["location"] == "header" assert body["errors"][0]["name"] == "X-Client-State" - def test_handle_new_users_disabled_error(self, request_token_route, valid_event): + def test_handle_new_users_disabled_error( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test NewUsersDisabledError returns 401 with new-users-disabled status""" - request_token_route.oidc_validator.validate_token.side_effect = NewUsersDisabledError( + mock_oidc_validator.validate_token.side_effect = NewUsersDisabledError( "New user registration is disabled" ) response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.UNAUTHORIZED - body = json.loads(response.body) + body = json_body(response) assert body["status"] == "new-users-disabled" assert body["errors"][0]["location"] == "server" assert body["errors"][0]["name"] == "registration" - def test_invalid_timestamp_includes_x_timestamp_header(self, request_token_route, valid_event): + def test_invalid_timestamp_includes_x_timestamp_header( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test InvalidTimestampError response includes X-Timestamp header""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidTimestampError( + mock_oidc_validator.validate_token.side_effect = InvalidTimestampError( "Token timestamp differs significantly from server time" ) @@ -1049,12 +1161,17 @@ def test_invalid_timestamp_includes_x_timestamp_header(self, request_token_route assert response.status_code == HTTPStatus.UNAUTHORIZED assert "X-Timestamp" in response.headers - timestamp = int(response.headers["X-Timestamp"]) + timestamp = int(header(response, "X-Timestamp")) assert timestamp > 0 - def test_invalid_generation_includes_x_timestamp_header(self, request_token_route, valid_event): + def test_invalid_generation_includes_x_timestamp_header( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test InvalidGenerationError response includes X-Timestamp header""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidGenerationError( + mock_oidc_validator.validate_token.side_effect = InvalidGenerationError( "Token generation number is outdated" ) @@ -1064,10 +1181,13 @@ def test_invalid_generation_includes_x_timestamp_header(self, request_token_rout assert "X-Timestamp" in response.headers def test_invalid_client_state_includes_x_timestamp_header( - self, request_token_route, valid_event - ): + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test InvalidClientStateError response includes X-Timestamp header""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidClientStateError( + mock_oidc_validator.validate_token.side_effect = InvalidClientStateError( "Client state has been seen before" ) @@ -1076,9 +1196,14 @@ def test_invalid_client_state_includes_x_timestamp_header( assert response.status_code == HTTPStatus.UNAUTHORIZED assert "X-Timestamp" in response.headers - def test_new_users_disabled_includes_x_timestamp_header(self, request_token_route, valid_event): + def test_new_users_disabled_includes_x_timestamp_header( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test NewUsersDisabledError response includes X-Timestamp header""" - request_token_route.oidc_validator.validate_token.side_effect = NewUsersDisabledError( + mock_oidc_validator.validate_token.side_effect = NewUsersDisabledError( "New user registration is disabled" ) @@ -1092,10 +1217,13 @@ class TestRetryAfterHeader: """Test Retry-After header on 503 responses""" def test_service_unavailable_includes_retry_after_header( - self, request_token_route, valid_event - ): + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test 503 response includes Retry-After header""" - request_token_route.oidc_validator.validate_token.side_effect = ServiceUnavailableError( + mock_oidc_validator.validate_token.side_effect = ServiceUnavailableError( "OIDC provider unreachable" ) @@ -1104,13 +1232,18 @@ def test_service_unavailable_includes_retry_after_header( assert response.status_code == HTTPStatus.SERVICE_UNAVAILABLE assert "Retry-After" in response.headers # Verify it's a valid integer - retry_after = int(response.headers["Retry-After"]) + retry_after = int(header(response, "Retry-After")) assert retry_after > 0 - def test_retry_after_header_value_is_correct(self, request_token_route, valid_event): + def test_retry_after_header_value_is_correct( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test Retry-After header value matches configured value""" # Default is 30 seconds - request_token_route.oidc_validator.validate_token.side_effect = ServiceUnavailableError( + mock_oidc_validator.validate_token.side_effect = ServiceUnavailableError( "OIDC provider unreachable" ) @@ -1120,8 +1253,12 @@ def test_retry_after_header_value_is_correct(self, request_token_route, valid_ev assert response.headers["Retry-After"] == "30" def test_retry_after_header_custom_value( - self, mock_oidc_validator, mock_user_manager, mock_token_generator, valid_event - ): + self, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + valid_event: APIGatewayProxyEvent, + ) -> None: """Test Retry-After header uses custom configured value""" # Create route with custom retry_after_seconds route = GetTokenRoute( @@ -1140,7 +1277,9 @@ def test_retry_after_header_custom_value( assert response.status_code == HTTPStatus.SERVICE_UNAVAILABLE assert response.headers["Retry-After"] == "60" - def test_non_503_responses_do_not_include_retry_after(self, request_token_route): + def test_non_503_responses_do_not_include_retry_after( + self, request_token_route: GetTokenRoute + ) -> None: """Test non-503 responses do not include Retry-After header""" # Test 401 error event = APIGatewayProxyEvent( @@ -1155,7 +1294,9 @@ def test_non_503_responses_do_not_include_retry_after(self, request_token_route) assert response.status_code == HTTPStatus.UNAUTHORIZED assert "Retry-After" not in response.headers - def test_400_error_does_not_include_retry_after(self, request_token_route): + def test_400_error_does_not_include_retry_after( + self, request_token_route: GetTokenRoute + ) -> None: """Test 400 error does not include Retry-After header""" event = APIGatewayProxyEvent( { @@ -1171,11 +1312,14 @@ def test_400_error_does_not_include_retry_after(self, request_token_route): assert response.status_code == HTTPStatus.BAD_REQUEST assert "Retry-After" not in response.headers - def test_500_error_does_not_include_retry_after(self, request_token_route, valid_event): + def test_500_error_does_not_include_retry_after( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test 500 error does not include Retry-After header""" - request_token_route.oidc_validator.validate_token.side_effect = RuntimeError( - "Unexpected error" - ) + mock_oidc_validator.validate_token.side_effect = RuntimeError("Unexpected error") response = request_token_route.handle(valid_event) @@ -1186,7 +1330,9 @@ def test_500_error_does_not_include_retry_after(self, request_token_route, valid class TestWWWAuthenticateHeader: """Test WWW-Authenticate header on 401 responses""" - def test_missing_auth_header_includes_www_authenticate(self, request_token_route): + def test_missing_auth_header_includes_www_authenticate( + self, request_token_route: GetTokenRoute + ) -> None: """Test 401 response for missing auth header includes WWW-Authenticate""" event = APIGatewayProxyEvent( { @@ -1199,81 +1345,37 @@ def test_missing_auth_header_includes_www_authenticate(self, request_token_route assert response.status_code == HTTPStatus.UNAUTHORIZED assert "WWW-Authenticate" in response.headers - assert response.headers["WWW-Authenticate"].startswith("Bearer") - - def test_invalid_credentials_includes_www_authenticate(self, request_token_route, valid_event): - """Test 401 response for invalid credentials includes WWW-Authenticate""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidCredentialsError( - "Token expired" - ) - - response = request_token_route.handle(valid_event) - - assert response.status_code == HTTPStatus.UNAUTHORIZED - assert "WWW-Authenticate" in response.headers - assert response.headers["WWW-Authenticate"].startswith("Bearer") - - def test_invalid_token_includes_www_authenticate(self, request_token_route, valid_event): - """Test 401 response for invalid token includes WWW-Authenticate""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidTokenError( - "Invalid signature" - ) - - response = request_token_route.handle(valid_event) - - assert response.status_code == HTTPStatus.UNAUTHORIZED - assert "WWW-Authenticate" in response.headers - assert response.headers["WWW-Authenticate"].startswith("Bearer") - - def test_invalid_timestamp_includes_www_authenticate(self, request_token_route, valid_event): - """Test 401 response for invalid timestamp includes WWW-Authenticate""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidTimestampError( - "Token timestamp differs significantly from server time" - ) - - response = request_token_route.handle(valid_event) - - assert response.status_code == HTTPStatus.UNAUTHORIZED - assert "WWW-Authenticate" in response.headers - assert response.headers["WWW-Authenticate"].startswith("Bearer") - - def test_invalid_generation_includes_www_authenticate(self, request_token_route, valid_event): - """Test 401 response for invalid generation includes WWW-Authenticate""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidGenerationError( - "Token generation number is outdated" - ) - - response = request_token_route.handle(valid_event) - - assert response.status_code == HTTPStatus.UNAUTHORIZED - assert "WWW-Authenticate" in response.headers - assert response.headers["WWW-Authenticate"].startswith("Bearer") - - def test_invalid_client_state_includes_www_authenticate(self, request_token_route, valid_event): - """Test 401 response for invalid client state includes WWW-Authenticate""" - request_token_route.oidc_validator.validate_token.side_effect = InvalidClientStateError( - "Client state has been seen before" - ) - - response = request_token_route.handle(valid_event) - - assert response.status_code == HTTPStatus.UNAUTHORIZED - assert "WWW-Authenticate" in response.headers - assert response.headers["WWW-Authenticate"].startswith("Bearer") - - def test_new_users_disabled_includes_www_authenticate(self, request_token_route, valid_event): - """Test 401 response for new users disabled includes WWW-Authenticate""" - request_token_route.oidc_validator.validate_token.side_effect = NewUsersDisabledError( - "New user registration is disabled" - ) + assert header(response, "WWW-Authenticate").startswith("Bearer") + + @pytest.mark.parametrize( + "error", + [ + InvalidCredentialsError("Token expired"), + InvalidTokenError("Invalid signature"), + InvalidTimestampError("Token timestamp differs significantly from server time"), + InvalidGenerationError("Token generation number is outdated"), + InvalidClientStateError("Client state has been seen before"), + NewUsersDisabledError("New user registration is disabled"), + ], + ids=lambda e: type(e).__name__, + ) + def test_validation_errors_include_www_authenticate( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + error: Exception, + ) -> None: + """Every 401 raised by token validation carries WWW-Authenticate: Bearer""" + mock_oidc_validator.validate_token.side_effect = error response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.UNAUTHORIZED assert "WWW-Authenticate" in response.headers - assert response.headers["WWW-Authenticate"].startswith("Bearer") + assert header(response, "WWW-Authenticate").startswith("Bearer") - def test_www_authenticate_header_format(self, request_token_route): + def test_www_authenticate_header_format(self, request_token_route: GetTokenRoute) -> None: """Test WWW-Authenticate header has correct Bearer format""" event = APIGatewayProxyEvent( { @@ -1285,7 +1387,7 @@ def test_www_authenticate_header_format(self, request_token_route): response = request_token_route.handle(event) assert response.status_code == HTTPStatus.UNAUTHORIZED - www_auth = response.headers["WWW-Authenticate"] + www_auth = header(response, "WWW-Authenticate") # Should be in format: Bearer realm="...", error="..." assert www_auth.startswith("Bearer") assert "realm=" in www_auth @@ -1293,24 +1395,29 @@ def test_www_authenticate_header_format(self, request_token_route): def test_non_401_responses_do_not_include_www_authenticate( self, - request_token_route, - valid_event, - mock_oidc_claims, - mock_user_record, - mock_token_response, - ): + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_claims: OIDCTokenClaims, + mock_user_record: UserRecord, + mock_token_response: TokenResponse, + mock_oidc_validator: MagicMock, + mock_user_manager: MagicMock, + mock_token_generator: MagicMock, + ) -> None: """Test non-401 responses do not include WWW-Authenticate header""" # Test 200 success response - request_token_route.oidc_validator.validate_token.return_value = mock_oidc_claims - request_token_route.user_manager.get_or_create_user.return_value = mock_user_record - request_token_route.token_generator.generate_token.return_value = mock_token_response + mock_oidc_validator.validate_token.return_value = mock_oidc_claims + mock_user_manager.get_or_create_user.return_value = mock_user_record + mock_token_generator.generate_token.return_value = mock_token_response response = request_token_route.handle(valid_event) assert response.status_code == HTTPStatus.OK assert "WWW-Authenticate" not in response.headers - def test_400_error_does_not_include_www_authenticate(self, request_token_route): + def test_400_error_does_not_include_www_authenticate( + self, request_token_route: GetTokenRoute + ) -> None: """Test 400 error does not include WWW-Authenticate header""" event = APIGatewayProxyEvent( { @@ -1326,9 +1433,14 @@ def test_400_error_does_not_include_www_authenticate(self, request_token_route): assert response.status_code == HTTPStatus.BAD_REQUEST assert "WWW-Authenticate" not in response.headers - def test_503_error_does_not_include_www_authenticate(self, request_token_route, valid_event): + def test_503_error_does_not_include_www_authenticate( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test 503 error does not include WWW-Authenticate header""" - request_token_route.oidc_validator.validate_token.side_effect = ServiceUnavailableError( + mock_oidc_validator.validate_token.side_effect = ServiceUnavailableError( "OIDC provider unreachable" ) @@ -1337,11 +1449,14 @@ def test_503_error_does_not_include_www_authenticate(self, request_token_route, assert response.status_code == HTTPStatus.SERVICE_UNAVAILABLE assert "WWW-Authenticate" not in response.headers - def test_500_error_does_not_include_www_authenticate(self, request_token_route, valid_event): + def test_500_error_does_not_include_www_authenticate( + self, + request_token_route: GetTokenRoute, + valid_event: APIGatewayProxyEvent, + mock_oidc_validator: MagicMock, + ) -> None: """Test 500 error does not include WWW-Authenticate header""" - request_token_route.oidc_validator.validate_token.side_effect = RuntimeError( - "Unexpected error" - ) + mock_oidc_validator.validate_token.side_effect = RuntimeError("Unexpected error") response = request_token_route.handle(valid_event) diff --git a/lambda/tests/services/test_auth_account_manager.py b/lambda/tests/services/test_auth_account_manager.py index 05d34a1a..dfad1727 100644 --- a/lambda/tests/services/test_auth_account_manager.py +++ b/lambda/tests/services/test_auth_account_manager.py @@ -1,55 +1,60 @@ """Unit tests for AuthAccountManager with DynamoDB stubber""" -from unittest.mock import patch +from typing import TYPE_CHECKING, Generator +from unittest.mock import MagicMock, patch import pytest from botocore.exceptions import ClientError +from botocore.stub import Stubber from src.services.auth_account_manager import AuthAccountManager +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + class TestAuthAccountManager: """Test AuthAccountManager DynamoDB operations""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> AuthAccountManager: """Create AuthAccountManager instance with stubbed table""" return AuthAccountManager(table=dynamodb_table) @pytest.fixture - def sample_uid(self): + def sample_uid(self) -> str: return "abcdef1234567890abcdef1234567890" @pytest.fixture - def sample_email(self): + def sample_email(self) -> str: return "Test.User@Example.com" @pytest.fixture - def sample_normalized_email(self): + def sample_normalized_email(self) -> str: return "test.user@example.com" @pytest.fixture - def sample_verify_hash(self): + def sample_verify_hash(self) -> str: return "a" * 64 @pytest.fixture - def sample_k_a(self): + def sample_k_a(self) -> str: return "b" * 64 @pytest.fixture - def sample_wrap_kb(self): + def sample_wrap_kb(self) -> str: return "c" * 64 @pytest.fixture - def sample_oidc_sub(self): + def sample_oidc_sub(self) -> str: return "oidc-sub-12345" @pytest.fixture - def sample_key_rotation_secret(self): + def sample_key_rotation_secret(self) -> str: return "d" * 64 @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: """Mock time.time() for auth_account_manager""" with patch("src.services.auth_account_manager.time") as mock: mock.time.return_value = 1234567890.0 @@ -59,19 +64,19 @@ def mock_time(self): def test_create_account_stores_email_then_account_records( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_email, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - sample_key_rotation_secret, - mock_time, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_email: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + sample_key_rotation_secret: str, + mock_time: MagicMock, + ) -> None: """Test create_account stores EMAIL# first then ACCOUNT# record""" # Stub put_item for EMAIL# record first (with condition) dynamodb_stubber.add_response( @@ -133,16 +138,16 @@ def test_create_account_stores_email_then_account_records( def test_create_account_normalizes_email( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - mock_time, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + mock_time: MagicMock, + ) -> None: """Test that email is lowercased and stripped during creation""" email = " User@EXAMPLE.COM " normalized = "user@example.com" @@ -206,18 +211,18 @@ def test_create_account_normalizes_email( def test_create_account_rejects_duplicate_email( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_email, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - mock_time, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_email: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + mock_time: MagicMock, + ) -> None: """Test that duplicate email raises ValueError""" # Stub EMAIL# put to fail with ConditionalCheckFailedException dynamodb_stubber.add_client_error( @@ -238,18 +243,18 @@ def test_create_account_rejects_duplicate_email( def test_create_account_cleans_up_email_on_oidcsub_write_failure( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_email, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - mock_time, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_email: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + mock_time: MagicMock, + ) -> None: """If OIDCSUB# put fails, EMAIL# record is cleaned up""" # Stub successful EMAIL# put dynamodb_stubber.add_response( @@ -294,18 +299,18 @@ def test_create_account_cleans_up_email_on_oidcsub_write_failure( def test_create_account_oidcsub_failure_with_email_cleanup_failure( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_email, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - mock_time, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_email: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + mock_time: MagicMock, + ) -> None: """If OIDCSUB# put fails and EMAIL# cleanup also fails, original error is raised""" # Stub successful EMAIL# put dynamodb_stubber.add_response( @@ -349,12 +354,12 @@ def test_create_account_oidcsub_failure_with_email_cleanup_failure( def test_ensure_oidcsub_record_creates_when_missing( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_oidc_sub, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_oidc_sub: str, + ) -> None: """Creates OIDCSUB# record when it doesn't exist""" dynamodb_stubber.add_response( "put_item", @@ -373,12 +378,12 @@ def test_ensure_oidcsub_record_creates_when_missing( def test_ensure_oidcsub_record_noop_when_exists( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_oidc_sub, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_oidc_sub: str, + ) -> None: """No-op when OIDCSUB# record already exists""" dynamodb_stubber.add_client_error( "put_item", @@ -390,12 +395,12 @@ def test_ensure_oidcsub_record_noop_when_exists( def test_ensure_oidcsub_record_raises_unexpected_error( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_oidc_sub, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_oidc_sub: str, + ) -> None: """Unexpected errors are re-raised""" dynamodb_stubber.add_client_error( "put_item", @@ -408,9 +413,9 @@ def test_ensure_oidcsub_record_raises_unexpected_error( def test_ensure_oidcsub_record_noop_when_empty_sub( self, - manager, - sample_uid, - ): + manager: AuthAccountManager, + sample_uid: str, + ) -> None: """No-op when oidc_sub is empty""" manager.ensure_oidcsub_record(sample_uid, "") @@ -418,16 +423,16 @@ def test_ensure_oidcsub_record_noop_when_empty_sub( def test_get_account_by_oidc_sub_returns_account( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + ) -> None: """Test get_account_by_oidc_sub returns account for existing OIDC subject""" # Stub get_item for OIDCSUB# record dynamodb_stubber.add_response( @@ -473,10 +478,10 @@ def test_get_account_by_oidc_sub_returns_account( def test_get_account_by_oidc_sub_returns_none_for_unknown( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """Test get_account_by_oidc_sub returns None for unknown OIDC subject""" oidc_sub = "unknown-oidc-sub" @@ -496,17 +501,17 @@ def test_get_account_by_oidc_sub_returns_none_for_unknown( def test_get_account_by_email_returns_account( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_email, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_email: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + ) -> None: """Test get_account_by_email returns account for existing email""" # Stub get_item for EMAIL# record dynamodb_stubber.add_response( @@ -559,10 +564,10 @@ def test_get_account_by_email_returns_account( def test_get_account_by_email_returns_none_for_unknown( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """Test get_account_by_email returns None for unknown email""" email = "unknown@example.com" @@ -584,16 +589,16 @@ def test_get_account_by_email_returns_none_for_unknown( def test_get_account_by_uid_returns_account( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + ) -> None: """Test get_account_by_uid returns account for existing uid""" # Stub get_item for ACCOUNT# record dynamodb_stubber.add_response( @@ -631,10 +636,10 @@ def test_get_account_by_uid_returns_account( def test_get_account_by_uid_returns_none_for_unknown( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """Test get_account_by_uid returns None for unknown uid""" uid = "nonexistent-uid-00000000000000000" @@ -654,18 +659,18 @@ def test_get_account_by_uid_returns_none_for_unknown( def test_create_account_cleans_up_on_account_write_failure( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_email, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - mock_time, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_email: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + mock_time: MagicMock, + ) -> None: """If ACCOUNT# put fails, EMAIL# and OIDCSUB# records are cleaned up""" # Stub successful EMAIL# put dynamodb_stubber.add_response( @@ -733,18 +738,18 @@ def test_create_account_cleans_up_on_account_write_failure( def test_create_account_cleans_up_even_if_cleanup_fails( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_email, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - mock_time, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_email: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + mock_time: MagicMock, + ) -> None: """If ACCOUNT# put fails and cleanup also fails, the original error is raised""" # Stub successful EMAIL# put dynamodb_stubber.add_response( @@ -808,18 +813,18 @@ def test_create_account_cleans_up_even_if_cleanup_fails( def test_create_account_reraises_unexpected_client_error( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_email, - sample_normalized_email, - sample_verify_hash, - sample_k_a, - sample_wrap_kb, - sample_oidc_sub, - mock_time, - ): + manager: AuthAccountManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_email: str, + sample_normalized_email: str, + sample_verify_hash: str, + sample_k_a: str, + sample_wrap_kb: str, + sample_oidc_sub: str, + mock_time: MagicMock, + ) -> None: """Test that unexpected ClientErrors on EMAIL# put are re-raised""" # Stub EMAIL# put to fail with unexpected error dynamodb_stubber.add_client_error( diff --git a/lambda/tests/services/test_channel_service.py b/lambda/tests/services/test_channel_service.py index a3c12ddf..e6bd9ca7 100644 --- a/lambda/tests/services/test_channel_service.py +++ b/lambda/tests/services/test_channel_service.py @@ -1,10 +1,13 @@ """Unit tests for ChannelService with DynamoDB stubber""" import json -from unittest.mock import patch +from typing import TYPE_CHECKING, Any, Dict, Optional, cast +from unittest.mock import MagicMock, patch +import boto3 import pytest from botocore.exceptions import ClientError +from botocore.stub import Stubber from src.services.channel_service import ( CHANNEL_TTL_SECONDS, @@ -13,12 +16,22 @@ ChannelService, ) +if TYPE_CHECKING: + from types_boto3_apigatewaymanagementapi.client import ApiGatewayManagementApiClient + from types_boto3_dynamodb.client import DynamoDBClient + from types_boto3_dynamodb.service_resource import Table + CHANNEL_TABLE_NAME = "test-channel-table" FIXED_UUID = "aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee" FIXED_TIME = 1700000000 -def _ws_event(route_key="$default", connection_id="conn-1", body=None, query_params=None): +def _ws_event( + route_key: str = "$default", + connection_id: str = "conn-1", + body: Optional[str] = None, + query_params: Optional[Dict[str, str]] = None, +) -> Dict[str, Any]: """Build a WebSocket API Gateway event dict.""" event = { "requestContext": { @@ -38,14 +51,22 @@ class TestChannelService: """Test ChannelService DynamoDB operations""" @pytest.fixture - def channel_table(self, boto_session, dynamodb_stubber): + def channel_table( + self, boto_session: boto3.session.Session, dynamodb_stubber: Stubber + ) -> "Table": resource = boto_session.resource("dynamodb") table = resource.Table(CHANNEL_TABLE_NAME) - table.meta.client = dynamodb_stubber.client + table.meta.client = cast("DynamoDBClient", dynamodb_stubber.client) return table @pytest.fixture - def service(self, channel_table, boto_session, apigw_client, apigw_stubber): + def service( + self, + channel_table: "Table", + boto_session: boto3.session.Session, + apigw_client: "ApiGatewayManagementApiClient", + apigw_stubber: Stubber, + ) -> ChannelService: svc = ChannelService(table=channel_table, session=boto_session) # Pre-populate the APIGW client cache with the shared stubbed client. # The key must match what _get_apigw_client computes from _ws_event(): @@ -57,7 +78,7 @@ def service(self, channel_table, boto_session, apigw_client, apigw_stubber): # -- Constants ------------------------------------------------------------ - def test_constants(self): + def test_constants(self) -> None: assert MAX_CONNECTIONS_PER_CHANNEL == 3 assert MAX_MESSAGES_PER_CHANNEL == 10 assert CHANNEL_TTL_SECONDS == 300 @@ -68,12 +89,12 @@ def test_constants(self): @patch("src.services.channel_service.time.time", return_value=FIXED_TIME) def test_create_channel( self, - mock_time, - mock_uuid, - service, - dynamodb_stubber, - apigw_stubber, - ): + mock_time: MagicMock, + mock_uuid: MagicMock, + service: ChannelService, + dynamodb_stubber: Stubber, + apigw_stubber: Stubber, + ) -> None: """Create channel stores metadata + reverse lookup + sends channelId.""" expiry = FIXED_TIME + CHANNEL_TTL_SECONDS @@ -127,10 +148,10 @@ def test_create_channel( @patch("src.services.channel_service.time.time", return_value=FIXED_TIME) def test_join_channel( self, - mock_time, - service, - dynamodb_stubber, - ): + mock_time: MagicMock, + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Join existing channel via atomic update_item + reverse lookup.""" expiry = FIXED_TIME + CHANNEL_TTL_SECONDS channel_id = "existing-channel" @@ -165,9 +186,9 @@ def test_join_channel( def test_join_nonexistent_channel_returns_404( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """ConditionalCheckFailed + empty get_item => 404.""" channel_id = "no-such-channel" @@ -201,9 +222,9 @@ def test_join_nonexistent_channel_returns_404( def test_join_full_channel_returns_403( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """ConditionalCheckFailed + channel exists => 403.""" channel_id = "full-channel" @@ -244,9 +265,9 @@ def test_join_full_channel_returns_403( def test_disconnect_cleans_up( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Disconnect removes reverse lookup then patches connections list.""" channel_id = "chan-1" @@ -305,9 +326,9 @@ def test_disconnect_cleans_up( def test_disconnect_unknown_connection( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Disconnect with no reverse lookup => no-op.""" # get_item for CONN# => empty dynamodb_stubber.add_response( @@ -328,10 +349,10 @@ def test_disconnect_unknown_connection( def test_relay_message( self, - service, - dynamodb_stubber, - apigw_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + apigw_stubber: Stubber, + ) -> None: """Message relayed to other connections in the channel.""" channel_id = "chan-1" @@ -395,9 +416,9 @@ def test_relay_message( def test_unknown_connection_message_returns_404( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Message from unknown connection => 404.""" # get_item for CONN# => empty dynamodb_stubber.add_response( @@ -418,9 +439,9 @@ def test_unknown_connection_message_returns_404( def test_channel_not_found_on_message_returns_404( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Message with valid connection but missing channel => 404.""" channel_id = "gone-channel" @@ -466,9 +487,9 @@ def test_channel_not_found_on_message_returns_404( def test_message_limit_returns_429( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Message count at limit => 429.""" channel_id = "busy-channel" @@ -521,10 +542,10 @@ def test_message_limit_returns_429( def test_gone_exception_triggers_cleanup( self, - service, - dynamodb_stubber, - apigw_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + apigw_stubber: Stubber, + ) -> None: """When post_to_connection raises GoneException, stale conn is cleaned up.""" channel_id = "chan-1" @@ -632,10 +653,10 @@ def test_gone_exception_triggers_cleanup( def test_empty_body_relay( self, - service, - dynamodb_stubber, - apigw_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + apigw_stubber: Stubber, + ) -> None: """Relay works when body is missing from event.""" channel_id = "chan-1" @@ -697,9 +718,9 @@ def test_empty_body_relay( def test_lazy_apigw_client_init( self, - channel_table, - boto_session, - ): + channel_table: "Table", + boto_session: boto3.session.Session, + ) -> None: """Client is created lazily on first use and cached by endpoint.""" svc = ChannelService(table=channel_table, session=boto_session) assert svc._apigw_clients == {} @@ -718,9 +739,9 @@ def test_lazy_apigw_client_init( def test_disconnect_channel_gone_after_conn_delete( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Disconnect when channel disappears between CONN delete and channel lookup.""" channel_id = "vanished-chan" @@ -769,9 +790,9 @@ def test_disconnect_channel_gone_after_conn_delete( def test_disconnect_connection_not_in_list( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Disconnect when connection is not in the channel's connections list.""" channel_id = "chan-1" @@ -827,9 +848,9 @@ def test_disconnect_connection_not_in_list( def test_join_unexpected_client_error_reraised( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Non-ConditionalCheckFailed ClientError is re-raised on join.""" dynamodb_stubber.add_client_error( "update_item", @@ -849,9 +870,9 @@ def test_join_unexpected_client_error_reraised( def test_message_unexpected_client_error_reraised( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Non-ConditionalCheckFailed ClientError is re-raised on message.""" channel_id = "chan-1" @@ -886,9 +907,9 @@ def test_message_unexpected_client_error_reraised( def test_message_channel_gone_after_count_update( self, - service, - dynamodb_stubber, - ): + service: ChannelService, + dynamodb_stubber: Stubber, + ) -> None: """Channel disappears between message count update and connections fetch.""" channel_id = "ephemeral-chan" @@ -928,7 +949,7 @@ def test_message_channel_gone_after_count_update( # -- Unknown route -------------------------------------------------------- - def test_unknown_route_returns_400(self, service): + def test_unknown_route_returns_400(self, service: ChannelService) -> None: """Unknown route key returns 400.""" event = _ws_event(route_key="$unknown") result = service.handle(event, None) diff --git a/lambda/tests/services/test_device_manager.py b/lambda/tests/services/test_device_manager.py index 5a81b285..ca99e9da 100644 --- a/lambda/tests/services/test_device_manager.py +++ b/lambda/tests/services/test_device_manager.py @@ -1,37 +1,42 @@ """Unit tests for DeviceManager with DynamoDB stubber""" -from unittest.mock import ANY, patch +from typing import TYPE_CHECKING, Generator +from unittest.mock import ANY, MagicMock, patch import pytest +from botocore.stub import Stubber from src.services.device_manager import DeviceManager +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + class TestDeviceManager: """Test DeviceManager DynamoDB operations""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> DeviceManager: """Create DeviceManager instance with stubbed table""" return DeviceManager(table=dynamodb_table) @pytest.fixture - def sample_uid(self): + def sample_uid(self) -> str: return "abcdef1234567890abcdef1234567890" @pytest.fixture - def sample_session_token_id(self): + def sample_session_token_id(self) -> str: return "session-token-id-abc123" @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: """Mock time.time() for device_manager""" with patch("src.services.device_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock @pytest.fixture - def mock_uuid(self): + def mock_uuid(self) -> Generator[MagicMock, None, None]: """Mock uuid.uuid4() for device_manager""" with patch("src.services.device_manager.uuid") as mock: mock_uuid4 = mock.uuid4.return_value @@ -42,14 +47,14 @@ def mock_uuid(self): def test_upsert_device_creates_new( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_session_token_id, - mock_time, - mock_uuid, - ): + manager: DeviceManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_session_token_id: str, + mock_time: MagicMock, + mock_uuid: MagicMock, + ) -> None: """upsert_device without id generates UUID and stores new device""" generated_id = "aabbccdd11223344aabbccdd11223344" @@ -94,13 +99,13 @@ def test_upsert_device_creates_new( def test_upsert_device_updates_existing( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - sample_session_token_id, - mock_time, - ): + manager: DeviceManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + sample_session_token_id: str, + mock_time: MagicMock, + ) -> None: """upsert_device with id merges fields into existing device""" device_id = "existing-device-id-00000000000000" @@ -160,11 +165,11 @@ def test_upsert_device_updates_existing( def test_get_devices_returns_all( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - ): + manager: DeviceManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + ) -> None: """get_devices returns all devices for a user""" device_id_1 = "device-1-00000000000000000000" device_id_2 = "device-2-00000000000000000000" @@ -208,11 +213,11 @@ def test_get_devices_returns_all( def test_get_devices_filters_idle( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - ): + manager: DeviceManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + ) -> None: """get_devices excludes devices with lastAccessTime below threshold""" device_id_active = "device-active-0000000000000000" device_id_idle = "device-idle-00000000000000000" @@ -251,11 +256,11 @@ def test_get_devices_filters_idle( def test_get_devices_empty( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - ): + manager: DeviceManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + ) -> None: """get_devices returns empty list when scan returns no items""" # Stub scan returning no items dynamodb_stubber.add_response( diff --git a/lambda/tests/services/test_fxa_crypto.py b/lambda/tests/services/test_fxa_crypto.py index affa796a..015e8c77 100644 --- a/lambda/tests/services/test_fxa_crypto.py +++ b/lambda/tests/services/test_fxa_crypto.py @@ -21,7 +21,7 @@ class TestNamespace: """Tests for the NAMESPACE constant.""" - def test_namespace_value(self): + def test_namespace_value(self) -> None: """Test that NAMESPACE matches the FxA protocol namespace.""" assert NAMESPACE == "identity.mozilla.com/picl/v1/" @@ -29,33 +29,33 @@ def test_namespace_value(self): class TestDeriveAuthPw: """Tests for derive_auth_pw function.""" - def test_returns_32_bytes(self): + def test_returns_32_bytes(self) -> None: """Test that derive_auth_pw returns exactly 32 bytes.""" qs_pw = b"\x00" * 32 result = derive_auth_pw(qs_pw) assert len(result) == 32 - def test_returns_bytes(self): + def test_returns_bytes(self) -> None: """Test that derive_auth_pw returns bytes type.""" qs_pw = b"\x00" * 32 result = derive_auth_pw(qs_pw) assert isinstance(result, bytes) - def test_deterministic(self): + def test_deterministic(self) -> None: """Test that the same input produces the same output.""" qs_pw = b"\xab\xcd" * 16 result1 = derive_auth_pw(qs_pw) result2 = derive_auth_pw(qs_pw) assert result1 == result2 - def test_differs_from_unwrap_bkey(self): + def test_differs_from_unwrap_bkey(self) -> None: """Test that authPW differs from unwrapBKey for the same input.""" qs_pw = b"\x01\x02\x03" * 11 # 33 bytes, arbitrary length auth_pw = derive_auth_pw(qs_pw) unwrap_bkey = derive_unwrap_bkey(qs_pw) assert auth_pw != unwrap_bkey - def test_different_inputs_produce_different_outputs(self): + def test_different_inputs_produce_different_outputs(self) -> None: """Test that different inputs produce different outputs.""" result1 = derive_auth_pw(b"\x00" * 32) result2 = derive_auth_pw(b"\x01" * 32) @@ -65,26 +65,26 @@ def test_different_inputs_produce_different_outputs(self): class TestDeriveUnwrapBkey: """Tests for derive_unwrap_bkey function.""" - def test_returns_32_bytes(self): + def test_returns_32_bytes(self) -> None: """Test that derive_unwrap_bkey returns exactly 32 bytes.""" qs_pw = b"\x00" * 32 result = derive_unwrap_bkey(qs_pw) assert len(result) == 32 - def test_returns_bytes(self): + def test_returns_bytes(self) -> None: """Test that derive_unwrap_bkey returns bytes type.""" qs_pw = b"\x00" * 32 result = derive_unwrap_bkey(qs_pw) assert isinstance(result, bytes) - def test_deterministic(self): + def test_deterministic(self) -> None: """Test that the same input produces the same output.""" qs_pw = b"\xab\xcd" * 16 result1 = derive_unwrap_bkey(qs_pw) result2 = derive_unwrap_bkey(qs_pw) assert result1 == result2 - def test_different_inputs_produce_different_outputs(self): + def test_different_inputs_produce_different_outputs(self) -> None: """Test that different inputs produce different outputs.""" result1 = derive_unwrap_bkey(b"\x00" * 32) result2 = derive_unwrap_bkey(b"\x01" * 32) @@ -94,32 +94,32 @@ def test_different_inputs_produce_different_outputs(self): class TestDeriveVerifyHash: """Tests for derive_verify_hash function.""" - def test_returns_32_bytes(self): + def test_returns_32_bytes(self) -> None: """Test that derive_verify_hash returns exactly 32 bytes.""" auth_pw = b"\x00" * 32 result = derive_verify_hash(auth_pw) assert len(result) == 32 - def test_returns_bytes(self): + def test_returns_bytes(self) -> None: """Test that derive_verify_hash returns bytes type.""" auth_pw = b"\x00" * 32 result = derive_verify_hash(auth_pw) assert isinstance(result, bytes) - def test_deterministic(self): + def test_deterministic(self) -> None: """Test that the same input produces the same output.""" auth_pw = b"\xab\xcd" * 16 result1 = derive_verify_hash(auth_pw) result2 = derive_verify_hash(auth_pw) assert result1 == result2 - def test_different_inputs_produce_different_outputs(self): + def test_different_inputs_produce_different_outputs(self) -> None: """Test that different inputs produce different outputs.""" result1 = derive_verify_hash(b"\x00" * 32) result2 = derive_verify_hash(b"\x01" * 32) assert result1 != result2 - def test_chained_derivation(self): + def test_chained_derivation(self) -> None: """Test that derive_verify_hash(derive_auth_pw(qsPW)) works correctly.""" qs_pw = b"\xaa" * 32 auth_pw = derive_auth_pw(qs_pw) @@ -131,19 +131,19 @@ def test_chained_derivation(self): class TestDeriveTokenId: """Tests for derive_token_id function.""" - def test_returns_32_bytes(self): + def test_returns_32_bytes(self) -> None: """Test that derive_token_id returns exactly 32 bytes.""" token = b"\x00" * 32 result = derive_token_id(token, "identity.mozilla.com/picl/v1/sessionToken") assert len(result) == 32 - def test_returns_bytes(self): + def test_returns_bytes(self) -> None: """Test that derive_token_id returns bytes type.""" token = b"\x00" * 32 result = derive_token_id(token, "identity.mozilla.com/picl/v1/sessionToken") assert isinstance(result, bytes) - def test_deterministic(self): + def test_deterministic(self) -> None: """Test that the same inputs produce the same output.""" token = b"\xab\xcd" * 16 info = "identity.mozilla.com/picl/v1/sessionToken" @@ -151,7 +151,7 @@ def test_deterministic(self): result2 = derive_token_id(token, info) assert result1 == result2 - def test_differs_from_req_hmac_key(self): + def test_differs_from_req_hmac_key(self) -> None: """Test that tokenId differs from reqHMACkey for the same input.""" token = b"\x00" * 32 info = "identity.mozilla.com/picl/v1/sessionToken" @@ -159,7 +159,7 @@ def test_differs_from_req_hmac_key(self): req_hmac_key = derive_req_hmac_key(token, info) assert token_id != req_hmac_key - def test_differs_from_key_request_key(self): + def test_differs_from_key_request_key(self) -> None: """Test that tokenId differs from keyRequestKey for the same input.""" token = b"\x00" * 32 info = "identity.mozilla.com/picl/v1/keyFetchToken" @@ -171,19 +171,19 @@ def test_differs_from_key_request_key(self): class TestDeriveReqHmacKey: """Tests for derive_req_hmac_key function.""" - def test_returns_32_bytes(self): + def test_returns_32_bytes(self) -> None: """Test that derive_req_hmac_key returns exactly 32 bytes.""" token = b"\x00" * 32 result = derive_req_hmac_key(token, "identity.mozilla.com/picl/v1/sessionToken") assert len(result) == 32 - def test_returns_bytes(self): + def test_returns_bytes(self) -> None: """Test that derive_req_hmac_key returns bytes type.""" token = b"\x00" * 32 result = derive_req_hmac_key(token, "identity.mozilla.com/picl/v1/sessionToken") assert isinstance(result, bytes) - def test_deterministic(self): + def test_deterministic(self) -> None: """Test that the same inputs produce the same output.""" token = b"\xab\xcd" * 16 info = "identity.mozilla.com/picl/v1/sessionToken" @@ -191,7 +191,7 @@ def test_deterministic(self): result2 = derive_req_hmac_key(token, info) assert result1 == result2 - def test_differs_from_key_request_key(self): + def test_differs_from_key_request_key(self) -> None: """Test that reqHMACkey differs from keyRequestKey for the same input.""" token = b"\x00" * 32 info = "identity.mozilla.com/picl/v1/keyFetchToken" @@ -203,19 +203,19 @@ def test_differs_from_key_request_key(self): class TestDeriveKeyRequestKey: """Tests for derive_key_request_key function.""" - def test_returns_32_bytes(self): + def test_returns_32_bytes(self) -> None: """Test that derive_key_request_key returns exactly 32 bytes.""" token = b"\x00" * 32 result = derive_key_request_key(token, "identity.mozilla.com/picl/v1/keyFetchToken") assert len(result) == 32 - def test_returns_bytes(self): + def test_returns_bytes(self) -> None: """Test that derive_key_request_key returns bytes type.""" token = b"\x00" * 32 result = derive_key_request_key(token, "identity.mozilla.com/picl/v1/keyFetchToken") assert isinstance(result, bytes) - def test_deterministic(self): + def test_deterministic(self) -> None: """Test that the same inputs produce the same output.""" token = b"\xab\xcd" * 16 info = "identity.mozilla.com/picl/v1/keyFetchToken" @@ -227,7 +227,7 @@ def test_deterministic(self): class TestTokenDerivedKeysAllDifferent: """Tests that all three token-derived keys differ from each other.""" - def test_all_three_keys_differ(self): + def test_all_three_keys_differ(self) -> None: """Test that tokenId, reqHMACkey, and keyRequestKey are all distinct.""" token = b"\x42" * 32 info = "identity.mozilla.com/picl/v1/keyFetchToken" @@ -240,7 +240,7 @@ def test_all_three_keys_differ(self): assert token_id != key_request_key assert req_hmac_key != key_request_key - def test_different_info_produces_different_keys(self): + def test_different_info_produces_different_keys(self) -> None: """Test that different info strings produce different derived keys.""" token = b"\x42" * 32 info1 = "identity.mozilla.com/picl/v1/sessionToken" @@ -254,7 +254,7 @@ def test_different_info_produces_different_keys(self): class TestEncryptKeyBundle: """Tests for encrypt_key_bundle function.""" - def test_returns_96_bytes(self): + def test_returns_96_bytes(self) -> None: """Test that encrypt_key_bundle returns exactly 96 bytes (64 ciphertext + 32 HMAC).""" key_request_key = b"\x00" * 32 k_a = b"\x11" * 32 @@ -262,7 +262,7 @@ def test_returns_96_bytes(self): result = encrypt_key_bundle(key_request_key, k_a, wrap_kb) assert len(result) == 96 - def test_returns_bytes(self): + def test_returns_bytes(self) -> None: """Test that encrypt_key_bundle returns bytes type.""" key_request_key = b"\x00" * 32 k_a = b"\x11" * 32 @@ -270,7 +270,7 @@ def test_returns_bytes(self): result = encrypt_key_bundle(key_request_key, k_a, wrap_kb) assert isinstance(result, bytes) - def test_ciphertext_differs_from_plaintext(self): + def test_ciphertext_differs_from_plaintext(self) -> None: """Test that the ciphertext portion differs from the plaintext (kA || wrapKB).""" key_request_key = b"\xaa" * 32 k_a = b"\x11" * 32 @@ -280,7 +280,7 @@ def test_ciphertext_differs_from_plaintext(self): plaintext = k_a + wrap_kb assert ciphertext != plaintext - def test_deterministic(self): + def test_deterministic(self) -> None: """Test that the same inputs produce the same output.""" key_request_key = b"\xaa" * 32 k_a = b"\x11" * 32 @@ -289,7 +289,7 @@ def test_deterministic(self): result2 = encrypt_key_bundle(key_request_key, k_a, wrap_kb) assert result1 == result2 - def test_different_keys_produce_different_output(self): + def test_different_keys_produce_different_output(self) -> None: """Test that different keyRequestKeys produce different bundles.""" k_a = b"\x11" * 32 wrap_kb = b"\x22" * 32 @@ -297,7 +297,7 @@ def test_different_keys_produce_different_output(self): result2 = encrypt_key_bundle(b"\xbb" * 32, k_a, wrap_kb) assert result1 != result2 - def test_mac_is_last_32_bytes(self): + def test_mac_is_last_32_bytes(self) -> None: """Test that the MAC (last 32 bytes) is a valid HMAC-SHA256 digest length.""" key_request_key = b"\xaa" * 32 k_a = b"\x11" * 32 @@ -310,22 +310,22 @@ def test_mac_is_last_32_bytes(self): class TestGenerateRandomBytes: """Tests for generate_random_bytes function.""" - def test_default_length(self): + def test_default_length(self) -> None: """Test that default length is 32 bytes.""" result = generate_random_bytes() assert len(result) == 32 - def test_custom_length(self): + def test_custom_length(self) -> None: """Test generation with custom length.""" result = generate_random_bytes(64) assert len(result) == 64 - def test_returns_bytes(self): + def test_returns_bytes(self) -> None: """Test that generate_random_bytes returns bytes type.""" result = generate_random_bytes() assert isinstance(result, bytes) - def test_different_each_call(self): + def test_different_each_call(self) -> None: """Test that successive calls produce different values.""" result1 = generate_random_bytes() result2 = generate_random_bytes() @@ -335,30 +335,30 @@ def test_different_each_call(self): class TestConstantTimeCompare: """Tests for constant_time_compare function.""" - def test_equal_bytes_returns_true(self): + def test_equal_bytes_returns_true(self) -> None: """Test that equal byte strings return True.""" a = b"\x01\x02\x03" assert constant_time_compare(a, a) is True - def test_equal_values_returns_true(self): + def test_equal_values_returns_true(self) -> None: """Test that equal-valued byte strings return True.""" a = b"\x01\x02\x03" b = b"\x01\x02\x03" assert constant_time_compare(a, b) is True - def test_different_bytes_returns_false(self): + def test_different_bytes_returns_false(self) -> None: """Test that different byte strings return False.""" a = b"\x01\x02\x03" b = b"\x04\x05\x06" assert constant_time_compare(a, b) is False - def test_different_lengths_returns_false(self): + def test_different_lengths_returns_false(self) -> None: """Test that byte strings of different lengths return False.""" a = b"\x01\x02\x03" b = b"\x01\x02" assert constant_time_compare(a, b) is False - def test_empty_bytes_returns_true(self): + def test_empty_bytes_returns_true(self) -> None: """Test that two empty byte strings return True.""" assert constant_time_compare(b"", b"") is True @@ -366,25 +366,25 @@ def test_empty_bytes_returns_true(self): class TestKnownVectors: """Regression tests with pinned HKDF outputs to catch protocol-breaking changes.""" - def test_derive_auth_pw_known_vector(self): + def test_derive_auth_pw_known_vector(self) -> None: """Test derive_auth_pw against a pinned output for bytes(32).""" qs_pw = bytes(32) result = derive_auth_pw(qs_pw) assert result.hex() == "addd287a170e5d4ab0a06a143a64fe3c6ab805ad0be1a38bd1ba5093c8fe124d" - def test_derive_verify_hash_chained_known_vector(self): + def test_derive_verify_hash_chained_known_vector(self) -> None: """Test derive_verify_hash(derive_auth_pw(bytes(32))) against a pinned output.""" auth_pw = derive_auth_pw(bytes(32)) result = derive_verify_hash(auth_pw) assert result.hex() == "b1873a935b3c91146743a9292107634b314a3ae6daf859f7fc0f986da557c27e" - def test_derive_unwrap_bkey_known_vector(self): + def test_derive_unwrap_bkey_known_vector(self) -> None: """Test derive_unwrap_bkey against a pinned output for bytes(32).""" qs_pw = bytes(32) result = derive_unwrap_bkey(qs_pw) assert result.hex() == "ad0e1de4f2362227e01eba2764d8d97c38ee1886bc13bcaa5d98690f0dee7781" - def test_encrypt_key_bundle_known_vector(self): + def test_encrypt_key_bundle_known_vector(self) -> None: """Test encrypt_key_bundle against a pinned 96-byte output for bytes(32) inputs.""" result = encrypt_key_bundle(bytes(32), bytes(32), bytes(32)) assert ( @@ -397,7 +397,7 @@ def test_encrypt_key_bundle_known_vector(self): class TestDeriveTokenKeys: """Tests for derive_token_keys function.""" - def test_returns_three_32_byte_keys(self): + def test_returns_three_32_byte_keys(self) -> None: """Test that derive_token_keys returns a tuple of three 32-byte keys.""" token = b"\x00" * 32 info = "identity.mozilla.com/picl/v1/sessionToken" @@ -406,7 +406,7 @@ def test_returns_three_32_byte_keys(self): assert len(req_hmac_key) == 32 assert len(key_request_key) == 32 - def test_returns_tuple(self): + def test_returns_tuple(self) -> None: """Test that derive_token_keys returns a tuple.""" token = b"\x00" * 32 info = "identity.mozilla.com/picl/v1/sessionToken" @@ -414,7 +414,7 @@ def test_returns_tuple(self): assert isinstance(result, tuple) assert len(result) == 3 - def test_matches_individual_derivations(self): + def test_matches_individual_derivations(self) -> None: """Test that derive_token_keys matches the individual derivation functions.""" token = b"\x42" * 32 info = "identity.mozilla.com/picl/v1/keyFetchToken" @@ -423,7 +423,7 @@ def test_matches_individual_derivations(self): assert req_hmac_key == derive_req_hmac_key(token, info) assert key_request_key == derive_key_request_key(token, info) - def test_deterministic(self): + def test_deterministic(self) -> None: """Test that the same inputs produce the same outputs.""" token = b"\xab\xcd" * 16 info = "identity.mozilla.com/picl/v1/sessionToken" @@ -431,7 +431,7 @@ def test_deterministic(self): result2 = derive_token_keys(token, info) assert result1 == result2 - def test_all_three_keys_differ(self): + def test_all_three_keys_differ(self) -> None: """Test that all three returned keys are distinct.""" token = b"\x42" * 32 info = "identity.mozilla.com/picl/v1/keyFetchToken" diff --git a/lambda/tests/services/test_fxa_token_manager.py b/lambda/tests/services/test_fxa_token_manager.py index a32ab22a..1c789c0d 100644 --- a/lambda/tests/services/test_fxa_token_manager.py +++ b/lambda/tests/services/test_fxa_token_manager.py @@ -1,9 +1,10 @@ """Unit tests for FxATokenManager with DynamoDB stubber""" +from typing import TYPE_CHECKING, Generator from unittest.mock import MagicMock, patch import pytest -from botocore.stub import ANY +from botocore.stub import ANY, Stubber from src.services import fxa_crypto from src.services.fxa_token_manager import ( @@ -13,37 +14,40 @@ ) from tests.fixtures.integration import build_hawk_auth_header as build_hawk_header +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + class TestCreateSessionToken: """Tests for create_session_token method""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) @pytest.fixture - def sample_uid(self): + def sample_uid(self) -> str: return "abcdef1234567890abcdef1234567890" @pytest.fixture - def fixed_token(self): + def fixed_token(self) -> bytes: return b"\xaa" * 32 @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: with patch("src.services.fxa_token_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock def test_returns_32_byte_raw_token( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - fixed_token, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + fixed_token: bytes, + mock_time: MagicMock, + ) -> None: """create_session_token returns the 32-byte raw token""" with patch("src.services.fxa_token_manager.fxa_crypto") as mock_crypto: mock_crypto.generate_random_bytes.return_value = fixed_token @@ -77,13 +81,13 @@ def test_returns_32_byte_raw_token( def test_stores_session_record_in_dynamodb( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - fixed_token, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + fixed_token: bytes, + mock_time: MagicMock, + ) -> None: """create_session_token stores a SESSION# record with correct fields""" with patch("src.services.fxa_token_manager.fxa_crypto") as mock_crypto: mock_crypto.generate_random_bytes.return_value = fixed_token @@ -118,27 +122,27 @@ class TestVerifySessionTokenId: """Tests for verify_session_token_id method""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) @pytest.fixture - def sample_uid(self): + def sample_uid(self) -> str: return "abcdef1234567890abcdef1234567890" @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: with patch("src.services.fxa_token_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock def test_returns_uid_for_valid_token( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_token_id returns uid when SESSION# record exists and not expired""" token_id_hex = "aa" * 32 @@ -165,10 +169,10 @@ def test_returns_uid_for_valid_token( def test_returns_none_for_unknown_token( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """verify_session_token_id returns None when token not found""" token_id_hex = "bb" * 32 @@ -187,12 +191,12 @@ def test_returns_none_for_unknown_token( def test_returns_none_for_expired_token( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_token_id returns None when token is expired""" token_id_hex = "aa" * 32 @@ -222,32 +226,32 @@ class TestCreateKeyFetchToken: """Tests for create_key_fetch_token method""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) @pytest.fixture - def sample_uid(self): + def sample_uid(self) -> str: return "abcdef1234567890abcdef1234567890" @pytest.fixture - def fixed_token(self): + def fixed_token(self) -> bytes: return b"\xcc" * 32 @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: with patch("src.services.fxa_token_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock def test_returns_32_byte_raw_token( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - fixed_token, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + fixed_token: bytes, + mock_time: MagicMock, + ) -> None: """create_key_fetch_token returns the 32-byte raw token""" with patch("src.services.fxa_token_manager.fxa_crypto") as mock_crypto: mock_crypto.generate_random_bytes.return_value = fixed_token @@ -280,13 +284,13 @@ def test_returns_32_byte_raw_token( def test_stores_keyfetch_record_with_raw_token_hex( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - fixed_token, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + fixed_token: bytes, + mock_time: MagicMock, + ) -> None: """create_key_fetch_token stores a KEYFETCH# record with the raw token hex""" with patch("src.services.fxa_token_manager.fxa_crypto") as mock_crypto: mock_crypto.generate_random_bytes.return_value = fixed_token @@ -320,27 +324,27 @@ class TestConsumeKeyFetchToken: """Tests for consume_key_fetch_token method (atomic delete)""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) @pytest.fixture - def sample_uid(self): + def sample_uid(self) -> str: return "abcdef1234567890abcdef1234567890" @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: with patch("src.services.fxa_token_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock def test_returns_uid_and_key_fetch_token( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """consume_key_fetch_token returns dict with uid and keyFetchToken""" token_id_hex = "dd" * 32 raw_token_hex = "cc" * 32 @@ -372,10 +376,10 @@ def test_returns_uid_and_key_fetch_token( def test_returns_none_for_unknown_token( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """consume_key_fetch_token returns None for unknown token (ConditionalCheckFailed)""" token_id_hex = "ee" * 32 @@ -391,12 +395,12 @@ def test_returns_none_for_unknown_token( def test_returns_none_for_expired_token( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """consume_key_fetch_token returns None when token is expired""" token_id_hex = "dd" * 32 raw_token_hex = "cc" * 32 @@ -428,15 +432,15 @@ class TestConsumeKeyFetchTokenEdgeCases: """Edge case tests for consume_key_fetch_token""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) def test_reraises_non_conditional_error( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """consume_key_fetch_token re-raises non-ConditionalCheckFailed errors""" from botocore.exceptions import ClientError @@ -453,18 +457,18 @@ def test_reraises_non_conditional_error( assert exc_info.value.response["Error"]["Code"] == "InternalServerError" @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: with patch("src.services.fxa_token_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock def test_returns_none_for_empty_attributes( self, - manager, - dynamodb_stubber, - storage_table_name, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_time: MagicMock, + ) -> None: """consume_key_fetch_token returns None when Attributes is empty""" token_id_hex = "ee" * 32 @@ -487,27 +491,27 @@ class TestVerifySessionHawk: """Tests for verify_session_hawk method""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) @pytest.fixture - def sample_uid(self): + def sample_uid(self) -> str: return "abcdef1234567890abcdef1234567890" @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: with patch("src.services.fxa_token_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock def test_returns_uid_for_valid_hawk( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns uid when HMAC is valid""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -518,7 +522,7 @@ def test_returns_uid_for_valid_hawk( "GET", "/v1/session/status", "localhost", - "443", + 443, ) dynamodb_stubber.add_response( @@ -550,19 +554,19 @@ def test_returns_uid_for_valid_hawk( ) result = manager.verify_session_hawk( - auth_header, "GET", "/v1/session/status", "localhost", "443" + auth_header, "GET", "/v1/session/status", "localhost", 443 ) assert result == sample_uid def test_returns_none_for_invalid_mac( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns None when HMAC does not match""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -588,24 +592,24 @@ def test_returns_none_for_invalid_mac( ) result = manager.verify_session_hawk( - auth_header, "GET", "/v1/session/status", "localhost", "443" + auth_header, "GET", "/v1/session/status", "localhost", 443 ) assert result is None - def test_returns_none_for_missing_header(self, manager): + def test_returns_none_for_missing_header(self, manager: FxATokenManager) -> None: """verify_session_hawk returns None when header is empty""" - result = manager.verify_session_hawk("", "GET", "/path", "host", "443") + result = manager.verify_session_hawk("", "GET", "/path", "host", 443) assert result is None def test_returns_none_for_expired_session( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns None when session is expired""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -616,7 +620,7 @@ def test_returns_none_for_expired_session( "GET", "/v1/session/status", "localhost", - "443", + 443, ) dynamodb_stubber.add_response( @@ -637,18 +641,18 @@ def test_returns_none_for_expired_session( ) result = manager.verify_session_hawk( - auth_header, "GET", "/v1/session/status", "localhost", "443" + auth_header, "GET", "/v1/session/status", "localhost", 443 ) assert result is None def test_returns_none_for_unknown_session( self, - manager, - dynamodb_stubber, - storage_table_name, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns None when session record not found""" token_id_hex = "aa" * 32 auth_header = f'Hawk id="{token_id_hex}", ts="1000000", nonce="abc", mac="AAAA"' @@ -662,17 +666,17 @@ def test_returns_none_for_unknown_session( }, ) - result = manager.verify_session_hawk(auth_header, "GET", "/path", "host", "443") + result = manager.verify_session_hawk(auth_header, "GET", "/path", "host", 443) assert result is None def test_returns_none_for_missing_hmac_key( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns None when reqHMACkey is missing""" token_id_hex = "aa" * 32 auth_header = f'Hawk id="{token_id_hex}", ts="1000000", nonce="abc", mac="AAAA"' @@ -692,17 +696,17 @@ def test_returns_none_for_missing_hmac_key( }, ) - result = manager.verify_session_hawk(auth_header, "GET", "/path", "host", "443") + result = manager.verify_session_hawk(auth_header, "GET", "/path", "host", 443) assert result is None def test_returns_uid_for_post_request( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns uid for POST requests (no content hash needed)""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -715,7 +719,7 @@ def test_returns_uid_for_post_request( "POST", "/v1/oauth/token", "localhost", - "443", + 443, ) dynamodb_stubber.add_response( @@ -747,19 +751,19 @@ def test_returns_uid_for_post_request( ) result = manager.verify_session_hawk( - auth_header, "POST", "/v1/oauth/token", "localhost", "443" + auth_header, "POST", "/v1/oauth/token", "localhost", 443 ) assert result == sample_uid def test_returns_uid_when_client_sends_content_hash( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns uid even when client includes content hash""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -771,7 +775,7 @@ def test_returns_uid_when_client_sends_content_hash( "POST", "/v1/oauth/token", "localhost", - "443", + 443, content='{"grant_type":"fxa-credentials"}', content_type="application/json", ) @@ -804,37 +808,37 @@ def test_returns_uid_when_client_sends_content_hash( ) result = manager.verify_session_hawk( - auth_header, "POST", "/v1/oauth/token", "localhost", "443" + auth_header, "POST", "/v1/oauth/token", "localhost", 443 ) assert result == sample_uid def test_returns_none_for_malformed_hawk_header( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns None for completely malformed header""" # Not a valid Hawk header at all auth_header = "Bearer some-token" result = manager.verify_session_hawk( - auth_header, "GET", "/v1/session/status", "localhost", "443" + auth_header, "GET", "/v1/session/status", "localhost", 443 ) assert result is None def test_returns_none_for_replayed_nonce( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk returns None when nonce has been replayed""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -845,7 +849,7 @@ def test_returns_none_for_replayed_nonce( "GET", "/v1/session/status", "localhost", - "443", + 443, ) # Stub: credential lookup succeeds @@ -874,18 +878,18 @@ def test_returns_none_for_replayed_nonce( ) result = manager.verify_session_hawk( - auth_header, "GET", "/v1/session/status", "localhost", "443" + auth_header, "GET", "/v1/session/status", "localhost", 443 ) assert result is None def test_reraises_non_conditional_nonce_error( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_session_hawk re-raises non-ConditionalCheckFailed errors from nonce check""" from botocore.exceptions import ClientError @@ -898,7 +902,7 @@ def test_reraises_non_conditional_nonce_error( "GET", "/v1/session/status", "localhost", - "443", + 443, ) # Stub: credential lookup succeeds @@ -927,9 +931,7 @@ def test_reraises_non_conditional_nonce_error( ) with pytest.raises(ClientError) as exc_info: - manager.verify_session_hawk( - auth_header, "GET", "/v1/session/status", "localhost", "443" - ) + manager.verify_session_hawk(auth_header, "GET", "/v1/session/status", "localhost", 443) assert exc_info.value.response["Error"]["Code"] == "InternalServerError" @@ -937,27 +939,27 @@ class TestVerifyKeyfetchHawk: """Tests for verify_keyfetch_hawk method""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) @pytest.fixture - def sample_uid(self): + def sample_uid(self) -> str: return "abcdef1234567890abcdef1234567890" @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: with patch("src.services.fxa_token_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock def test_returns_token_data_for_valid_hawk( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_keyfetch_hawk returns uid and keyFetchToken when HMAC is valid""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -969,7 +971,7 @@ def test_returns_token_data_for_valid_hawk( "GET", "/v1/account/keys", "localhost", - "443", + 443, ) dynamodb_stubber.add_response( @@ -1003,7 +1005,7 @@ def test_returns_token_data_for_valid_hawk( ) result = manager.verify_keyfetch_hawk( - auth_header, "GET", "/v1/account/keys", "localhost", "443" + auth_header, "GET", "/v1/account/keys", "localhost", 443 ) assert result is not None @@ -1012,10 +1014,10 @@ def test_returns_token_data_for_valid_hawk( def test_returns_none_for_unknown_token( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """verify_keyfetch_hawk returns None when token doesn't exist""" token_id_hex = "aa" * 32 auth_header = f'Hawk id="{token_id_hex}", ts="1000000", nonce="abc", mac="AAAA"' @@ -1027,24 +1029,24 @@ def test_returns_none_for_unknown_token( ) result = manager.verify_keyfetch_hawk( - auth_header, "GET", "/v1/account/keys", "localhost", "443" + auth_header, "GET", "/v1/account/keys", "localhost", 443 ) assert result is None - def test_returns_none_for_missing_header(self, manager): + def test_returns_none_for_missing_header(self, manager: FxATokenManager) -> None: """verify_keyfetch_hawk returns None when header is empty""" - result = manager.verify_keyfetch_hawk("", "GET", "/path", "host", "443") + result = manager.verify_keyfetch_hawk("", "GET", "/path", "host", 443) assert result is None def test_returns_none_for_expired_keyfetch( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_keyfetch_hawk returns None when token is expired""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -1055,7 +1057,7 @@ def test_returns_none_for_expired_keyfetch( "GET", "/v1/account/keys", "localhost", - "443", + 443, ) dynamodb_stubber.add_response( @@ -1078,18 +1080,18 @@ def test_returns_none_for_expired_keyfetch( ) result = manager.verify_keyfetch_hawk( - auth_header, "GET", "/v1/account/keys", "localhost", "443" + auth_header, "GET", "/v1/account/keys", "localhost", 443 ) assert result is None def test_returns_none_for_invalid_mac( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_keyfetch_hawk returns None when HMAC does not match""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -1116,18 +1118,18 @@ def test_returns_none_for_invalid_mac( ) result = manager.verify_keyfetch_hawk( - auth_header, "GET", "/v1/account/keys", "localhost", "443" + auth_header, "GET", "/v1/account/keys", "localhost", 443 ) assert result is None def test_returns_none_for_missing_hmac_key( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_keyfetch_hawk returns None when reqHMACkey is absent""" token_id_hex = "aa" * 32 auth_header = f'Hawk id="{token_id_hex}", ts="1000000", nonce="abc", mac="AAAA"' @@ -1151,16 +1153,16 @@ def test_returns_none_for_missing_hmac_key( ) result = manager.verify_keyfetch_hawk( - auth_header, "GET", "/v1/account/keys", "localhost", "443" + auth_header, "GET", "/v1/account/keys", "localhost", 443 ) assert result is None def test_reraises_non_conditional_error( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """verify_keyfetch_hawk re-raises non-ConditionalCheckFailed errors""" from botocore.exceptions import ClientError @@ -1174,16 +1176,16 @@ def test_reraises_non_conditional_error( ) with pytest.raises(ClientError) as exc_info: - manager.verify_keyfetch_hawk(auth_header, "GET", "/v1/account/keys", "localhost", "443") + manager.verify_keyfetch_hawk(auth_header, "GET", "/v1/account/keys", "localhost", 443) assert exc_info.value.response["Error"]["Code"] == "InternalServerError" def test_returns_none_for_empty_attributes( self, - manager, - dynamodb_stubber, - storage_table_name, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_time: MagicMock, + ) -> None: """verify_keyfetch_hawk returns None when Attributes is empty""" token_id_hex = "aa" * 32 auth_header = f'Hawk id="{token_id_hex}", ts="1000000", nonce="abc", mac="AAAA"' @@ -1200,18 +1202,18 @@ def test_returns_none_for_empty_attributes( ) result = manager.verify_keyfetch_hawk( - auth_header, "GET", "/v1/account/keys", "localhost", "443" + auth_header, "GET", "/v1/account/keys", "localhost", 443 ) assert result is None def test_returns_token_data_for_post_request( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_keyfetch_hawk returns uid and keyFetchToken for POST requests""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -1224,7 +1226,7 @@ def test_returns_token_data_for_post_request( "POST", "/v1/account/keys", "localhost", - "443", + 443, ) dynamodb_stubber.add_response( @@ -1258,7 +1260,7 @@ def test_returns_token_data_for_post_request( ) result = manager.verify_keyfetch_hawk( - auth_header, "POST", "/v1/account/keys", "localhost", "443" + auth_header, "POST", "/v1/account/keys", "localhost", 443 ) assert result is not None @@ -1267,12 +1269,12 @@ def test_returns_token_data_for_post_request( def test_returns_none_for_replayed_nonce( self, - manager, - dynamodb_stubber, - storage_table_name, - sample_uid, - mock_time, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + sample_uid: str, + mock_time: MagicMock, + ) -> None: """verify_keyfetch_hawk returns None when nonce has been replayed""" token_id_hex = "aa" * 32 req_hmac_key_hex = "bb" * 32 @@ -1284,7 +1286,7 @@ def test_returns_none_for_replayed_nonce( "GET", "/v1/account/keys", "localhost", - "443", + 443, ) # Stub: credential lookup succeeds (atomic delete) @@ -1315,7 +1317,7 @@ def test_returns_none_for_replayed_nonce( ) result = manager.verify_keyfetch_hawk( - auth_header, "GET", "/v1/account/keys", "localhost", "443" + auth_header, "GET", "/v1/account/keys", "localhost", 443 ) assert result is None @@ -1324,10 +1326,12 @@ class TestVerifyKeyfetchHawkEdgeCases: """Edge case tests for verify_keyfetch_hawk""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) - def test_returns_none_when_receiver_skips_credentials_map(self, manager): + def test_returns_none_when_receiver_skips_credentials_map( + self, manager: FxATokenManager + ) -> None: """verify_keyfetch_hawk returns None when result_holder has no uid (edge case: mohawk.Receiver completes but credentials_map was not invoked)""" token_id_hex = "aa" * 32 @@ -1337,7 +1341,7 @@ def test_returns_none_when_receiver_skips_credentials_map(self, manager): # Mock Receiver to succeed without calling credentials_map, # so result_holder stays empty result = manager.verify_keyfetch_hawk( - auth_header, "GET", "/v1/account/keys", "localhost", "443" + auth_header, "GET", "/v1/account/keys", "localhost", 443 ) assert result is None @@ -1346,15 +1350,15 @@ class TestDeleteSession: """Tests for delete_session method""" @pytest.fixture - def manager(self, dynamodb_table): + def manager(self, dynamodb_table: "Table") -> FxATokenManager: return FxATokenManager(table=dynamodb_table, metrics=MagicMock()) def test_deletes_session_record( self, - manager, - dynamodb_stubber, - storage_table_name, - ): + manager: FxATokenManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """delete_session deletes the SESSION# record""" token_id_hex = "ff" * 32 @@ -1375,10 +1379,10 @@ def test_deletes_session_record( class TestConstants: """Tests for module-level constants""" - def test_session_token_info(self): + def test_session_token_info(self) -> None: assert SESSION_TOKEN_INFO == "identity.mozilla.com/picl/v1/sessionToken" - def test_key_fetch_token_info(self): + def test_key_fetch_token_info(self) -> None: assert KEY_FETCH_TOKEN_INFO == "identity.mozilla.com/picl/v1/keyFetchToken" @@ -1386,23 +1390,23 @@ class TestCustomTTL: """Tests for custom TTL values""" @pytest.fixture - def fixed_token(self): + def fixed_token(self) -> bytes: return b"\xaa" * 32 @pytest.fixture - def mock_time(self): + def mock_time(self) -> Generator[MagicMock, None, None]: with patch("src.services.fxa_token_manager.time") as mock: mock.time.return_value = 1000000.0 yield mock def test_custom_session_ttl( self, - dynamodb_table, - dynamodb_stubber, - storage_table_name, - fixed_token, - mock_time, - ): + dynamodb_table: "Table", + dynamodb_stubber: Stubber, + storage_table_name: str, + fixed_token: bytes, + mock_time: MagicMock, + ) -> None: """Custom session_ttl_seconds is used for session token expiry""" custom_ttl = 3600 # 1 hour manager = FxATokenManager( @@ -1439,12 +1443,12 @@ def test_custom_session_ttl( def test_custom_keyfetch_ttl( self, - dynamodb_table, - dynamodb_stubber, - storage_table_name, - fixed_token, - mock_time, - ): + dynamodb_table: "Table", + dynamodb_stubber: Stubber, + storage_table_name: str, + fixed_token: bytes, + mock_time: MagicMock, + ) -> None: """Custom keyfetch_ttl_seconds is used for key-fetch token expiry""" custom_ttl = 60 # 1 minute manager = FxATokenManager( @@ -1480,14 +1484,14 @@ def test_custom_keyfetch_ttl( class TestExtractTokenIdFromHawkHeader: - def test_valid_header(self): + def test_valid_header(self) -> None: result = FxATokenManager.extract_token_id_from_hawk_header( 'Hawk id="abc123", ts="1234567890", nonce="xyz"' ) assert result == "abc123" - def test_empty_header(self): + def test_empty_header(self) -> None: assert FxATokenManager.extract_token_id_from_hawk_header("") is None - def test_no_match(self): + def test_no_match(self) -> None: assert FxATokenManager.extract_token_id_from_hawk_header("Bearer token") is None diff --git a/lambda/tests/services/test_hawk_service.py b/lambda/tests/services/test_hawk_service.py index eb7ad38a..3dbe1c25 100644 --- a/lambda/tests/services/test_hawk_service.py +++ b/lambda/tests/services/test_hawk_service.py @@ -21,14 +21,14 @@ @pytest.fixture -def mock_dynamodb_table(): +def mock_dynamodb_table() -> MagicMock: """Mock DynamoDB table""" table = MagicMock() return table @pytest.fixture -def hawk_service(mock_dynamodb_table): +def hawk_service(mock_dynamodb_table: MagicMock) -> HawkService: """Create HawkService with mocked DynamoDB""" service = HawkService(token_cache_table=mock_dynamodb_table) return service @@ -37,7 +37,9 @@ def hawk_service(mock_dynamodb_table): class TestHawkServiceInit: """Tests for HawkService initialization""" - def test_init_stores_table(self, hawk_service, mock_dynamodb_table): + def test_init_stores_table( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test that initialization stores the DynamoDB table""" assert hawk_service.token_cache_table is mock_dynamodb_table @@ -45,28 +47,28 @@ def test_init_stores_table(self, hawk_service, mock_dynamodb_table): class TestExtractHawkId: """Tests for _extract_hawk_id method""" - def test_extract_hawk_id_valid(self, hawk_service): + def test_extract_hawk_id_valid(self, hawk_service: HawkService) -> None: """Test extracting id from a valid Hawk header""" header = 'Hawk id="abc123", ts="1234567890", nonce="xyz", mac="sig=="' result = hawk_service._extract_hawk_id(header) assert result == "abc123" - def test_extract_hawk_id_missing_prefix(self, hawk_service): + def test_extract_hawk_id_missing_prefix(self, hawk_service: HawkService) -> None: """Test extraction fails when header doesn't start with 'Hawk '""" with pytest.raises(InvalidHawkHeaderException, match="must start with 'Hawk '"): hawk_service._extract_hawk_id('id="abc123", ts="1234567890"') - def test_extract_hawk_id_empty_header(self, hawk_service): + def test_extract_hawk_id_empty_header(self, hawk_service: HawkService) -> None: """Test extraction fails for empty header""" with pytest.raises(InvalidHawkHeaderException, match="must start with 'Hawk '"): hawk_service._extract_hawk_id("") - def test_extract_hawk_id_none_header(self, hawk_service): + def test_extract_hawk_id_none_header(self, hawk_service: HawkService) -> None: """Test extraction fails for None header""" with pytest.raises(InvalidHawkHeaderException, match="must start with 'Hawk '"): - hawk_service._extract_hawk_id(None) + hawk_service._extract_hawk_id(None) # type: ignore[arg-type] - def test_extract_hawk_id_missing_id_field(self, hawk_service): + def test_extract_hawk_id_missing_id_field(self, hawk_service: HawkService) -> None: """Test extraction fails when id field is missing""" with pytest.raises(InvalidHawkHeaderException, match="Missing id"): hawk_service._extract_hawk_id('Hawk ts="1234567890", nonce="xyz"') @@ -75,7 +77,7 @@ def test_extract_hawk_id_missing_id_field(self, hawk_service): class TestDecodeHawkId: """Tests for decode_hawk_id method""" - def test_decode_valid_hawk_id(self, hawk_service): + def test_decode_valid_hawk_id(self, hawk_service: HawkService) -> None: """Test decoding a valid HAWK ID""" # Create a valid HAWK ID: user123:5:1234567890 hawk_id = base64.urlsafe_b64encode(b"user123:5:1234567890").decode("utf-8").rstrip("=") @@ -86,7 +88,7 @@ def test_decode_valid_hawk_id(self, hawk_service): assert generation == 5 assert expiry == 1234567890 - def test_decode_hawk_id_with_padding(self, hawk_service): + def test_decode_hawk_id_with_padding(self, hawk_service: HawkService) -> None: """Test decoding HAWK ID that needs padding""" # Create HAWK ID with padding hawk_id = base64.urlsafe_b64encode(b"user:1:9999").decode("utf-8").rstrip("=") @@ -97,7 +99,7 @@ def test_decode_hawk_id_with_padding(self, hawk_service): assert generation == 1 assert expiry == 9999 - def test_decode_hawk_id_invalid_parts(self, hawk_service): + def test_decode_hawk_id_invalid_parts(self, hawk_service: HawkService) -> None: """Test decoding HAWK ID with wrong number of parts""" # Create invalid HAWK ID with only 2 parts hawk_id = base64.urlsafe_b64encode(b"user123:5").decode("utf-8").rstrip("=") @@ -107,7 +109,7 @@ def test_decode_hawk_id_invalid_parts(self, hawk_service): assert "HAWK ID must have 3 parts" in str(exc_info.value) - def test_decode_hawk_id_invalid_base64(self, hawk_service): + def test_decode_hawk_id_invalid_base64(self, hawk_service: HawkService) -> None: """Test decoding invalid base64 HAWK ID""" hawk_id = "not-valid-base64!!!" @@ -116,7 +118,7 @@ def test_decode_hawk_id_invalid_base64(self, hawk_service): assert "Invalid HAWK ID format" in str(exc_info.value) - def test_decode_hawk_id_invalid_generation(self, hawk_service): + def test_decode_hawk_id_invalid_generation(self, hawk_service: HawkService) -> None: """Test decoding HAWK ID with non-integer generation""" hawk_id = base64.urlsafe_b64encode(b"user:abc:1234567890").decode("utf-8").rstrip("=") @@ -129,7 +131,7 @@ def test_decode_hawk_id_invalid_generation(self, hawk_service): class TestValidateHawkIdExpiry: """Tests for validate_hawk_id_expiry method""" - def test_validate_expiry_not_expired(self, hawk_service): + def test_validate_expiry_not_expired(self, hawk_service: HawkService) -> None: """Test validation of non-expired token""" future_expiry = int(time.time()) + 300 # 5 minutes in future @@ -137,7 +139,7 @@ def test_validate_expiry_not_expired(self, hawk_service): assert result is True - def test_validate_expiry_expired(self, hawk_service): + def test_validate_expiry_expired(self, hawk_service: HawkService) -> None: """Test validation of expired token""" past_expiry = int(time.time()) - 300 # 5 minutes in past @@ -145,7 +147,7 @@ def test_validate_expiry_expired(self, hawk_service): assert result is False - def test_validate_expiry_at_boundary(self, hawk_service): + def test_validate_expiry_at_boundary(self, hawk_service: HawkService) -> None: """Test validation at expiry boundary""" current_time = int(time.time()) @@ -158,7 +160,9 @@ def test_validate_expiry_at_boundary(self, hawk_service): class TestGetHawkKeyFromCache: """Tests for get_hawk_key_from_cache method""" - def test_get_hawk_key_success(self, hawk_service, mock_dynamodb_table): + def test_get_hawk_key_success( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test successful retrieval of HAWK key from cache""" hawk_id = "test_hawk_id" mock_dynamodb_table.get_item.return_value = { @@ -176,7 +180,9 @@ def test_get_hawk_key_success(self, hawk_service, mock_dynamodb_table): assert generation == 5 mock_dynamodb_table.get_item.assert_called_once_with(Key={"PK": f"TOKEN#{hawk_id}"}) - def test_get_hawk_key_not_found(self, hawk_service, mock_dynamodb_table): + def test_get_hawk_key_not_found( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test retrieval when token not found in cache""" hawk_id = "nonexistent_hawk_id" mock_dynamodb_table.get_item.return_value = {} @@ -186,7 +192,9 @@ def test_get_hawk_key_not_found(self, hawk_service, mock_dynamodb_table): assert "HAWK token not found" in str(exc_info.value) - def test_get_hawk_key_dynamodb_error(self, hawk_service, mock_dynamodb_table): + def test_get_hawk_key_dynamodb_error( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test retrieval when DynamoDB error occurs""" hawk_id = "test_hawk_id" mock_dynamodb_table.get_item.side_effect = ClientError( @@ -203,7 +211,9 @@ def test_get_hawk_key_dynamodb_error(self, hawk_service, mock_dynamodb_table): class TestValidate: """Tests for validate method (full validation flow)""" - def test_validate_success(self, hawk_service, mock_dynamodb_table): + def test_validate_success( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test successful HAWK validation""" # Create valid HAWK credentials user_id = "user123" @@ -240,7 +250,7 @@ def test_validate_success(self, hawk_service, mock_dynamodb_table): assert credentials.expiry == expiry assert credentials.hawk_id == hawk_id - def test_validate_expired_token(self, hawk_service): + def test_validate_expired_token(self, hawk_service: HawkService) -> None: """Test validation with expired token""" user_id = "user123" generation = 5 @@ -263,7 +273,9 @@ def test_validate_expired_token(self, hawk_service): authorization_header, "GET", "/storage/bookmarks", "api.example.com", 443 ) - def test_validate_generation_mismatch(self, hawk_service, mock_dynamodb_table): + def test_validate_generation_mismatch( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test validation with generation mismatch""" user_id = "user123" generation = 5 @@ -295,7 +307,9 @@ def test_validate_generation_mismatch(self, hawk_service, mock_dynamodb_table): with pytest.raises(InvalidGenerationException): hawk_service.validate(authorization_header, method, path, host, port) - def test_validate_invalid_mac(self, hawk_service, mock_dynamodb_table): + def test_validate_invalid_mac( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test validation with invalid MAC (wrong key)""" user_id = "user123" generation = 5 @@ -327,7 +341,9 @@ def test_validate_invalid_mac(self, hawk_service, mock_dynamodb_table): with pytest.raises(InvalidHawkSignatureException, match="MAC verification failed"): hawk_service.validate(authorization_header, method, path, host, port) - def test_validate_timestamp_outside_skew(self, hawk_service, mock_dynamodb_table): + def test_validate_timestamp_outside_skew( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test validation with timestamp outside acceptable window""" user_id = "user123" generation = 5 @@ -361,7 +377,9 @@ def test_validate_timestamp_outside_skew(self, hawk_service, mock_dynamodb_table with pytest.raises(InvalidHawkSignatureException, match="outside acceptable window"): hawk_service.validate(authorization_header, method, path, host, port) - def test_validate_token_not_in_cache(self, hawk_service, mock_dynamodb_table): + def test_validate_token_not_in_cache( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test validation when token is not found in DynamoDB cache""" user_id = "user123" generation = 5 @@ -385,7 +403,9 @@ def test_validate_token_not_in_cache(self, hawk_service, mock_dynamodb_table): with pytest.raises(AuthenticationException, match="HAWK token not found"): hawk_service.validate(authorization_header, method, path, host, port) - def test_validate_dynamodb_error(self, hawk_service, mock_dynamodb_table): + def test_validate_dynamodb_error( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test validation when DynamoDB raises an error""" user_id = "user123" generation = 5 @@ -412,19 +432,21 @@ def test_validate_dynamodb_error(self, hawk_service, mock_dynamodb_table): with pytest.raises(AuthenticationException, match="Failed to retrieve HAWK token"): hawk_service.validate(authorization_header, method, path, host, port) - def test_validate_missing_header(self, hawk_service): + def test_validate_missing_header(self, hawk_service: HawkService) -> None: """Test validation with missing authorization header""" with pytest.raises(InvalidHawkHeaderException, match="must start with 'Hawk '"): hawk_service.validate("", "GET", "/storage/bookmarks", "api.example.com", 443) - def test_validate_malformed_header(self, hawk_service): + def test_validate_malformed_header(self, hawk_service: HawkService) -> None: """Test validation with malformed header (no Hawk prefix)""" with pytest.raises(InvalidHawkHeaderException, match="must start with 'Hawk '"): hawk_service.validate( "Bearer token123", "GET", "/storage/bookmarks", "api.example.com", 443 ) - def test_validate_bad_header_value(self, hawk_service, mock_dynamodb_table): + def test_validate_bad_header_value( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test validation with malformed Hawk header content (BadHeaderValue)""" user_id = "user123" generation = 5 @@ -452,7 +474,9 @@ def test_validate_bad_header_value(self, hawk_service, mock_dynamodb_table): authorization_header, "GET", "/storage/bookmarks", "api.example.com", 443 ) - def test_validate_missing_authorization_from_mohawk(self, hawk_service, mock_dynamodb_table): + def test_validate_missing_authorization_from_mohawk( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test MissingAuthorization exception path from mohawk.Receiver""" user_id = "user123" generation = 5 @@ -483,7 +507,9 @@ def test_validate_missing_authorization_from_mohawk(self, hawk_service, mock_dyn authorization_header, "GET", "/storage/bookmarks", "api.example.com", 443 ) - def test_validate_generic_hawk_fail(self, hawk_service, mock_dynamodb_table): + def test_validate_generic_hawk_fail( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test generic HawkFail catch-all exception path""" user_id = "user123" generation = 5 @@ -514,7 +540,9 @@ def test_validate_generic_hawk_fail(self, hawk_service, mock_dynamodb_table): authorization_header, "GET", "/storage/bookmarks", "api.example.com", 443 ) - def test_validate_rejects_replayed_nonce(self, hawk_service, mock_dynamodb_table): + def test_validate_rejects_replayed_nonce( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Replayed nonce is rejected""" user_id = "user1" generation = 0 @@ -537,7 +565,9 @@ def test_validate_rejects_replayed_nonce(self, hawk_service, mock_dynamodb_table class TestValidateQueryParamCorrection: """Tests for query parameter reordering correction in validate()""" - def test_validate_corrects_reordered_query_params(self, hawk_service, mock_dynamodb_table): + def test_validate_corrects_reordered_query_params( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """validate() succeeds when API Gateway alphabetizes query params.""" user_id = "user123" generation = 5 @@ -573,8 +603,8 @@ def test_validate_corrects_reordered_query_params(self, hawk_service, mock_dynam assert creds.user_id == user_id def test_validate_correction_returns_none_falls_through( - self, hawk_service, mock_dynamodb_table - ): + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """When no permutation matches, validate proceeds with original path (and fails).""" user_id = "user1" generation = 1 @@ -611,8 +641,8 @@ def test_validate_correction_returns_none_falls_through( ) def test_validate_no_correction_needed_for_single_param( - self, hawk_service, mock_dynamodb_table - ): + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Single query param doesn't trigger permutation logic.""" user_id = "user1" generation = 1 @@ -643,14 +673,18 @@ def test_validate_no_correction_needed_for_single_param( class TestSeenNonce: """Tests for _seen_nonce method""" - def test_seen_nonce_new_nonce(self, hawk_service, mock_dynamodb_table): + def test_seen_nonce_new_nonce( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """New nonce returns False (not seen before)""" mock_dynamodb_table.put_item.return_value = {} result = hawk_service._seen_nonce("sender1", "nonce1", "12345") assert result is False mock_dynamodb_table.put_item.assert_called_once() - def test_seen_nonce_replay_detected(self, hawk_service, mock_dynamodb_table): + def test_seen_nonce_replay_detected( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Replayed nonce returns True""" mock_dynamodb_table.put_item.side_effect = ClientError( {"Error": {"Code": "ConditionalCheckFailedException", "Message": ""}}, @@ -659,7 +693,9 @@ def test_seen_nonce_replay_detected(self, hawk_service, mock_dynamodb_table): result = hawk_service._seen_nonce("sender1", "nonce1", "12345") assert result is True - def test_seen_nonce_reraises_other_errors(self, hawk_service, mock_dynamodb_table): + def test_seen_nonce_reraises_other_errors( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Non-conditional errors are re-raised""" mock_dynamodb_table.put_item.side_effect = ClientError( {"Error": {"Code": "InternalServerError", "Message": "Server error"}}, @@ -673,7 +709,7 @@ def test_seen_nonce_reraises_other_errors(self, hawk_service, mock_dynamodb_tabl class TestGenerateHawkCredentials: """Tests for generate_hawk_credentials method""" - def test_generate_hawk_credentials(self, hawk_service): + def test_generate_hawk_credentials(self, hawk_service: HawkService) -> None: """Test generation of HAWK credentials""" user_id = "user123" generation = 5 @@ -687,7 +723,7 @@ def test_generate_hawk_credentials(self, hawk_service): assert len(credentials.hawk_key) == 64 # 32 bytes as hex = 64 chars assert credentials.expiry > int(time.time()) - def test_generate_hawk_credentials_expiry(self, hawk_service): + def test_generate_hawk_credentials_expiry(self, hawk_service: HawkService) -> None: """Test that generated credentials have correct expiry""" user_id = "user123" generation = 5 @@ -705,7 +741,7 @@ def test_generate_hawk_credentials_expiry(self, hawk_service): class TestGenerateHawkId: """Tests for generate_hawk_id method""" - def test_generate_hawk_id_format(self, hawk_service): + def test_generate_hawk_id_format(self, hawk_service: HawkService) -> None: """Test HAWK ID generation format""" user_id = "user123" generation = 5 @@ -721,7 +757,7 @@ def test_generate_hawk_id_format(self, hawk_service): decoded = base64.urlsafe_b64decode(hawk_id + "==").decode("utf-8") assert decoded == f"{user_id}:{generation}:{expiry}" - def test_generate_hawk_id_different_inputs(self, hawk_service): + def test_generate_hawk_id_different_inputs(self, hawk_service: HawkService) -> None: """Test that different inputs produce different HAWK IDs""" hawk_id1 = hawk_service.generate_hawk_id("user1", 1, 1000) hawk_id2 = hawk_service.generate_hawk_id("user2", 1, 1000) @@ -734,7 +770,7 @@ def test_generate_hawk_id_different_inputs(self, hawk_service): class TestGenerateHawkKey: """Tests for generate_hawk_key method""" - def test_generate_hawk_key_format(self, hawk_service): + def test_generate_hawk_key_format(self, hawk_service: HawkService) -> None: """Test HAWK key generation format""" hawk_key = hawk_service.generate_hawk_key() @@ -742,7 +778,7 @@ def test_generate_hawk_key_format(self, hawk_service): assert len(hawk_key) == 64 assert all(c in "0123456789abcdef" for c in hawk_key) - def test_generate_hawk_key_uniqueness(self, hawk_service): + def test_generate_hawk_key_uniqueness(self, hawk_service: HawkService) -> None: """Test that generated keys are unique""" key1 = hawk_service.generate_hawk_key() key2 = hawk_service.generate_hawk_key() @@ -753,7 +789,9 @@ def test_generate_hawk_key_uniqueness(self, hawk_service): class TestStoreTokenInCache: """Tests for store_token_in_cache method""" - def test_store_token_in_cache(self, hawk_service, mock_dynamodb_table): + def test_store_token_in_cache( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Test storing token in cache""" credentials = HawkCredentials( user_id="user123", @@ -781,7 +819,7 @@ def test_store_token_in_cache(self, hawk_service, mock_dynamodb_table): class TestHawkCredentialsDataclass: """Tests for HawkCredentials dataclass""" - def test_hawk_credentials_creation(self): + def test_hawk_credentials_creation(self) -> None: """Test creating HawkCredentials""" credentials = HawkCredentials( user_id="user123", @@ -797,7 +835,7 @@ def test_hawk_credentials_creation(self): assert credentials.hawk_id == "test_hawk_id" assert credentials.hawk_key == "a" * 64 - def test_hawk_credentials_optional_key(self): + def test_hawk_credentials_optional_key(self) -> None: """Test HawkCredentials with optional hawk_key""" credentials = HawkCredentials( user_id="user123", @@ -812,7 +850,7 @@ def test_hawk_credentials_optional_key(self): class TestParseHawkFields: """Tests for _parse_hawk_fields""" - def test_parses_all_fields(self, hawk_service): + def test_parses_all_fields(self, hawk_service: HawkService) -> None: header = 'Hawk id="abc", ts="123", nonce="xyz", mac="sig=", hash="h", ext="e"' fields = hawk_service._parse_hawk_fields(header) assert fields["id"] == "abc" @@ -822,14 +860,16 @@ def test_parses_all_fields(self, hawk_service): assert fields["hash"] == "h" assert fields["ext"] == "e" - def test_empty_header(self, hawk_service): + def test_empty_header(self, hawk_service: HawkService) -> None: assert hawk_service._parse_hawk_fields("") == {} class TestCorrectQueryOrder: """Tests for _correct_query_order — pre-computes MAC to find client's param ordering""" - def test_finds_correct_order(self, hawk_service, mock_dynamodb_table): + def test_finds_correct_order( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """When API Gateway reorders params, finds the client's original order.""" import hashlib import hmac as hmac_mod @@ -867,7 +907,9 @@ def test_finds_correct_order(self, hawk_service, mock_dynamodb_table): ) assert result == "/storage/prefs?newer=1.09&full=1&limit=1000" - def test_returns_none_when_already_correct(self, hawk_service, mock_dynamodb_table): + def test_returns_none_when_already_correct( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Returns the current path when the order already matches.""" import hashlib import hmac as hmac_mod @@ -900,14 +942,16 @@ def test_returns_none_when_already_correct(self, hawk_service, mock_dynamodb_tab # Returns the resource (it matched on the current order) assert result == resource - def test_returns_none_on_missing_header_fields(self, hawk_service): + def test_returns_none_on_missing_header_fields(self, hawk_service: HawkService) -> None: """Returns None when header can't be parsed.""" result = hawk_service._correct_query_order( "BadHeader", "hid", "GET", "/path?a=1&b=2", {"a": "1", "b": "2"}, "h", 443 ) assert result is None - def test_returns_none_on_cache_miss(self, hawk_service, mock_dynamodb_table): + def test_returns_none_on_cache_miss( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Returns None when hawk key not in cache (lets mohawk handle it).""" mock_dynamodb_table.get_item.return_value = {} # No Item hawk_id = base64.urlsafe_b64encode(b"user1:1:9999999999").decode().rstrip("=") @@ -917,7 +961,9 @@ def test_returns_none_on_cache_miss(self, hawk_service, mock_dynamodb_table): ) assert result is None - def test_returns_none_when_no_permutation_matches(self, hawk_service, mock_dynamodb_table): + def test_returns_none_when_no_permutation_matches( + self, hawk_service: HawkService, mock_dynamodb_table: MagicMock + ) -> None: """Returns None when MAC doesn't match any permutation (tampered request).""" hawk_key = "c" * 64 hawk_id = base64.urlsafe_b64encode(b"user1:1:9999999999").decode().rstrip("=") diff --git a/lambda/tests/services/test_jwt_service.py b/lambda/tests/services/test_jwt_service.py index f568c23a..14a43962 100644 --- a/lambda/tests/services/test_jwt_service.py +++ b/lambda/tests/services/test_jwt_service.py @@ -14,7 +14,7 @@ from src.services.jwt_service import JWTService -def _generate_test_rsa_key(): +def _generate_test_rsa_key() -> tuple[rsa.RSAPrivateKey, bytes]: """Generate a test RSA key pair for round-trip testing.""" private_key = rsa.generate_private_key( public_exponent=65537, @@ -29,12 +29,12 @@ def _generate_test_rsa_key(): @pytest.fixture -def rsa_keys(): +def rsa_keys() -> tuple[rsa.RSAPrivateKey, bytes]: return _generate_test_rsa_key() @pytest.fixture -def mock_kms(rsa_keys): +def mock_kms(rsa_keys: tuple[rsa.RSAPrivateKey, bytes]) -> MagicMock: _, public_key_der = rsa_keys client = MagicMock() client.get_public_key.return_value = { @@ -51,7 +51,7 @@ def mock_kms(rsa_keys): @pytest.fixture -def service(mock_kms): +def service(mock_kms: MagicMock) -> JWTService: return JWTService( kms_client=mock_kms, signing_key_id="key-123", @@ -60,12 +60,12 @@ def service(mock_kms): class TestIssuerProperty: - def test_returns_issuer(self, service): + def test_returns_issuer(self, service: JWTService) -> None: assert service.issuer == "https://auth.prod.ffsync.layertwo.dev" class TestSignJWT: - def test_returns_three_part_jwt(self, service): + def test_returns_three_part_jwt(self, service: JWTService) -> None: token = service.sign_jwt( sub="user1", scope="https://identity.mozilla.com/apps/oldsync", @@ -74,7 +74,7 @@ def test_returns_three_part_jwt(self, service): parts = token.split(".") assert len(parts) == 3 - def test_header_specifies_rs256(self, service): + def test_header_specifies_rs256(self, service: JWTService) -> None: token = service.sign_jwt(sub="user1", scope="openid", ttl=300) # Add padding for base64 decode header_b64 = token.split(".")[0] @@ -84,7 +84,7 @@ def test_header_specifies_rs256(self, service): assert header["typ"] == "JWT" assert "kid" in header - def test_payload_contains_claims(self, service): + def test_payload_contains_claims(self, service: JWTService) -> None: token = service.sign_jwt(sub="user1", scope="openid", ttl=300) payload_b64 = token.split(".")[1] payload_b64 += "=" * (4 - len(payload_b64) % 4) @@ -96,35 +96,35 @@ def test_payload_contains_claims(self, service): assert "iat" in payload assert payload["exp"] - payload["iat"] == 300 - def test_payload_contains_client_id(self, service): + def test_payload_contains_client_id(self, service: JWTService) -> None: token = service.sign_jwt(sub="user1", scope="openid", ttl=300, client_id="test-client") payload_b64 = token.split(".")[1] payload_b64 += "=" * (4 - len(payload_b64) % 4) payload = json.loads(base64.urlsafe_b64decode(payload_b64)) assert payload["client_id"] == "test-client" - def test_payload_omits_client_id_when_none(self, service): + def test_payload_omits_client_id_when_none(self, service: JWTService) -> None: token = service.sign_jwt(sub="user1", scope="openid", ttl=300) payload_b64 = token.split(".")[1] payload_b64 += "=" * (4 - len(payload_b64) % 4) payload = json.loads(base64.urlsafe_b64decode(payload_b64)) assert "client_id" not in payload - def test_payload_contains_fxa_uid(self, service): + def test_payload_contains_fxa_uid(self, service: JWTService) -> None: token = service.sign_jwt(sub="oidc-sub", scope="openid", ttl=300, fxa_uid="uid1") payload_b64 = token.split(".")[1] payload_b64 += "=" * (4 - len(payload_b64) % 4) payload = json.loads(base64.urlsafe_b64decode(payload_b64)) assert payload["fxa_uid"] == "uid1" - def test_payload_omits_fxa_uid_when_none(self, service): + def test_payload_omits_fxa_uid_when_none(self, service: JWTService) -> None: token = service.sign_jwt(sub="user1", scope="openid", ttl=300) payload_b64 = token.split(".")[1] payload_b64 += "=" * (4 - len(payload_b64) % 4) payload = json.loads(base64.urlsafe_b64decode(payload_b64)) assert "fxa_uid" not in payload - def test_calls_kms_sign(self, service, mock_kms): + def test_calls_kms_sign(self, service: JWTService, mock_kms: MagicMock) -> None: service.sign_jwt(sub="user1", scope="openid", ttl=300) mock_kms.sign.assert_called_once() call_kwargs = mock_kms.sign.call_args.kwargs @@ -132,7 +132,7 @@ def test_calls_kms_sign(self, service, mock_kms): assert call_kwargs["SigningAlgorithm"] == "RSASSA_PKCS1_V1_5_SHA_256" assert call_kwargs["MessageType"] == "RAW" - def test_different_subs_produce_different_tokens(self, service): + def test_different_subs_produce_different_tokens(self, service: JWTService) -> None: token1 = service.sign_jwt(sub="user1", scope="openid", ttl=300) token2 = service.sign_jwt(sub="user2", scope="openid", ttl=300) # Headers may be the same but payloads should differ @@ -140,7 +140,7 @@ def test_different_subs_produce_different_tokens(self, service): class TestGetPublicKeyJWK: - def test_returns_jwk_dict(self, service): + def test_returns_jwk_dict(self, service: JWTService) -> None: jwk = service.get_public_key_jwk() assert isinstance(jwk, dict) assert jwk["kty"] == "RSA" @@ -150,13 +150,13 @@ def test_returns_jwk_dict(self, service): assert "e" in jwk assert "kid" in jwk - def test_caches_public_key(self, service, mock_kms): + def test_caches_public_key(self, service: JWTService, mock_kms: MagicMock) -> None: service.get_public_key_jwk() service.get_public_key_jwk() # Should only call KMS once due to caching mock_kms.get_public_key.assert_called_once() - def test_kid_matches_header(self, service): + def test_kid_matches_header(self, service: JWTService) -> None: jwk = service.get_public_key_jwk() token = service.sign_jwt(sub="user1", scope="openid", ttl=300) header_b64 = token.split(".")[0] diff --git a/lambda/tests/services/test_jwt_verifier.py b/lambda/tests/services/test_jwt_verifier.py index 913b7b47..953e9617 100644 --- a/lambda/tests/services/test_jwt_verifier.py +++ b/lambda/tests/services/test_jwt_verifier.py @@ -18,7 +18,7 @@ def _b64url(data: bytes) -> str: return base64.urlsafe_b64encode(data).rstrip(b"=").decode("ascii") -def _generate_test_keypair(): +def _generate_test_keypair() -> tuple[rsa.RSAPrivateKey, dict[str, str]]: """Generate an RSA keypair for testing.""" private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) public_key = private_key.public_key() @@ -38,7 +38,7 @@ def _generate_test_keypair(): return private_key, jwk -def _sign_jwt(private_key, payload: dict, header: dict | None = None) -> str: +def _sign_jwt(private_key: rsa.RSAPrivateKey, payload: dict, header: dict | None = None) -> str: """Sign a JWT with the given private key.""" if header is None: header = {"alg": "RS256", "typ": "JWT", "kid": "test-kid"} @@ -54,12 +54,12 @@ def _sign_jwt(private_key, payload: dict, header: dict | None = None) -> str: @pytest.fixture -def keypair(): +def keypair() -> tuple[rsa.RSAPrivateKey, dict[str, str]]: return _generate_test_keypair() @pytest.fixture -def mock_jwt_service(keypair): +def mock_jwt_service(keypair: tuple[rsa.RSAPrivateKey, dict[str, str]]) -> MagicMock: _, jwk = keypair svc = MagicMock() svc.get_public_key_jwk.return_value = jwk @@ -68,12 +68,14 @@ def mock_jwt_service(keypair): @pytest.fixture -def verifier(mock_jwt_service): +def verifier(mock_jwt_service: MagicMock) -> JWTVerifier: return JWTVerifier(jwt_service=mock_jwt_service) class TestJWTVerifier: - def test_valid_token(self, verifier, keypair): + def test_valid_token( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: private_key, _ = keypair now = int(time.time()) payload = { @@ -90,7 +92,9 @@ def test_valid_token(self, verifier, keypair): assert claims.exp == now + 900 assert claims.fxa_uid is None - def test_fxa_uid_extracted_from_token(self, verifier, keypair): + def test_fxa_uid_extracted_from_token( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: private_key, _ = keypair now = int(time.time()) payload = { @@ -105,7 +109,9 @@ def test_fxa_uid_extracted_from_token(self, verifier, keypair): assert claims.sub == "oidc-sub-123" assert claims.fxa_uid == "uid-abc123" - def test_expired_token_raises(self, verifier, keypair): + def test_expired_token_raises( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: private_key, _ = keypair now = int(time.time()) payload = { @@ -118,7 +124,7 @@ def test_expired_token_raises(self, verifier, keypair): with pytest.raises(InvalidTokenError, match="expired"): verifier.validate_token(token) - def test_invalid_signature_raises(self, verifier): + def test_invalid_signature_raises(self, verifier: JWTVerifier) -> None: # Sign with a different key other_private_key, _ = _generate_test_keypair() now = int(time.time()) @@ -132,7 +138,9 @@ def test_invalid_signature_raises(self, verifier): with pytest.raises(InvalidTokenError, match="signature"): verifier.validate_token(token) - def test_missing_sub_raises(self, verifier, keypair): + def test_missing_sub_raises( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: private_key, _ = keypair now = int(time.time()) payload = {"iss": "https://auth.example.com", "iat": now, "exp": now + 900} @@ -140,7 +148,9 @@ def test_missing_sub_raises(self, verifier, keypair): with pytest.raises(InvalidTokenError, match="sub"): verifier.validate_token(token) - def test_missing_exp_raises(self, verifier, keypair): + def test_missing_exp_raises( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: private_key, _ = keypair now = int(time.time()) payload = {"sub": "user123", "iss": "https://auth.example.com", "iat": now} @@ -148,11 +158,13 @@ def test_missing_exp_raises(self, verifier, keypair): with pytest.raises(InvalidTokenError, match="exp"): verifier.validate_token(token) - def test_invalid_jwt_format_raises(self, verifier): + def test_invalid_jwt_format_raises(self, verifier: JWTVerifier) -> None: with pytest.raises(InvalidTokenError, match="format"): verifier.validate_token("not-a-jwt") - def test_unsupported_algorithm_raises(self, verifier, keypair): + def test_unsupported_algorithm_raises( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: private_key, _ = keypair now = int(time.time()) payload = { @@ -165,20 +177,24 @@ def test_unsupported_algorithm_raises(self, verifier, keypair): with pytest.raises(InvalidTokenError, match="algorithm"): verifier.validate_token(token) - def test_invalid_header_base64_raises(self, verifier): + def test_invalid_header_base64_raises(self, verifier: JWTVerifier) -> None: # Create a token with invalid base64 header token = "!!!invalid!!!.eyJzdWIiOiJ0ZXN0In0.signature" with pytest.raises(InvalidTokenError, match="header"): verifier.validate_token(token) - def test_invalid_payload_base64_raises(self, verifier, keypair): + def test_invalid_payload_base64_raises( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: # Create a token with valid header but invalid payload header_b64 = _b64url(json.dumps({"alg": "RS256", "typ": "JWT"}).encode()) token = f"{header_b64}.!!!invalid!!!.signature" with pytest.raises(InvalidTokenError, match="payload"): verifier.validate_token(token) - def test_invalid_issuer_raises(self, verifier, keypair): + def test_invalid_issuer_raises( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: private_key, _ = keypair now = int(time.time()) payload = { @@ -191,7 +207,9 @@ def test_invalid_issuer_raises(self, verifier, keypair): with pytest.raises(InvalidTokenError, match="issuer"): verifier.validate_token(token) - def test_client_id_mapped_to_aud(self, verifier, keypair): + def test_client_id_mapped_to_aud( + self, verifier: JWTVerifier, keypair: tuple[rsa.RSAPrivateKey, dict[str, str]] + ) -> None: private_key, _ = keypair now = int(time.time()) payload = { diff --git a/lambda/tests/services/test_oauth_code_manager.py b/lambda/tests/services/test_oauth_code_manager.py index 9859f92d..a2014a6c 100644 --- a/lambda/tests/services/test_oauth_code_manager.py +++ b/lambda/tests/services/test_oauth_code_manager.py @@ -10,17 +10,17 @@ @pytest.fixture -def mock_table(): +def mock_table() -> MagicMock: return MagicMock() @pytest.fixture -def manager(mock_table): +def manager(mock_table: MagicMock) -> OAuthCodeManager: return OAuthCodeManager(table=mock_table, code_ttl_seconds=600, refresh_ttl_seconds=86400) class TestCreateAuthorizationCode: - def test_returns_code_string(self, manager): + def test_returns_code_string(self, manager: OAuthCodeManager) -> None: code = manager.create_authorization_code( uid="uid1", client_id="client1", @@ -31,7 +31,7 @@ def test_returns_code_string(self, manager): assert isinstance(code, str) assert len(code) > 0 - def test_stores_code_in_dynamo(self, manager, mock_table): + def test_stores_code_in_dynamo(self, manager: OAuthCodeManager, mock_table: MagicMock) -> None: manager.create_authorization_code( uid="uid1", client_id="client1", @@ -50,7 +50,9 @@ def test_stores_code_in_dynamo(self, manager, mock_table): assert item["keysJwe"] == "" assert "expiry" in item - def test_stores_keys_jwe_in_dynamo(self, manager, mock_table): + def test_stores_keys_jwe_in_dynamo( + self, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: manager.create_authorization_code( uid="uid1", client_id="client1", @@ -63,7 +65,7 @@ def test_stores_keys_jwe_in_dynamo(self, manager, mock_table): item = mock_table.put_item.call_args.kwargs["Item"] assert item["keysJwe"] == "some-jwe-value" - def test_different_calls_produce_different_codes(self, manager): + def test_different_calls_produce_different_codes(self, manager: OAuthCodeManager) -> None: code1 = manager.create_authorization_code( uid="uid1", client_id="client1", @@ -83,7 +85,9 @@ def test_different_calls_produce_different_codes(self, manager): class TestConsumeAuthorizationCode: @patch("src.services.oauth_code_manager.time") - def test_returns_code_data_atomically(self, mock_time, manager, mock_table): + def test_returns_code_data_atomically( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 mock_table.delete_item.return_value = { "Attributes": { @@ -111,7 +115,9 @@ def test_returns_code_data_atomically(self, mock_time, manager, mock_table): assert call_kwargs["ConditionExpression"] == "attribute_exists(PK)" @patch("src.services.oauth_code_manager.time") - def test_returns_empty_keys_jwe_when_missing(self, mock_time, manager, mock_table): + def test_returns_empty_keys_jwe_when_missing( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 mock_table.delete_item.return_value = { "Attributes": { @@ -128,7 +134,9 @@ def test_returns_empty_keys_jwe_when_missing(self, mock_time, manager, mock_tabl assert result is not None assert result["keysJwe"] == "" - def test_returns_none_for_unknown_code(self, manager, mock_table): + def test_returns_none_for_unknown_code( + self, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_table.delete_item.side_effect = ClientError( {"Error": {"Code": "ConditionalCheckFailedException", "Message": ""}}, "DeleteItem", @@ -137,7 +145,9 @@ def test_returns_none_for_unknown_code(self, manager, mock_table): assert result is None @patch("src.services.oauth_code_manager.time") - def test_returns_none_for_expired_code(self, mock_time, manager, mock_table): + def test_returns_none_for_expired_code( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 mock_table.delete_item.return_value = { "Attributes": { @@ -155,12 +165,14 @@ def test_returns_none_for_expired_code(self, mock_time, manager, mock_table): class TestCreateRefreshToken: - def test_returns_token_string(self, manager): + def test_returns_token_string(self, manager: OAuthCodeManager) -> None: token = manager.create_refresh_token(uid="uid1", client_id="client1", scope="openid") assert isinstance(token, str) assert len(token) > 0 - def test_stores_refresh_in_dynamo(self, manager, mock_table): + def test_stores_refresh_in_dynamo( + self, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: manager.create_refresh_token(uid="uid1", client_id="client1", scope="openid") mock_table.put_item.assert_called_once() item = mock_table.put_item.call_args.kwargs["Item"] @@ -173,7 +185,9 @@ def test_stores_refresh_in_dynamo(self, manager, mock_table): class TestConsumeRefreshToken: @patch("src.services.oauth_code_manager.time") - def test_returns_data_atomically(self, mock_time, manager, mock_table): + def test_returns_data_atomically( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 token_hash = hashlib.sha256(b"token123").hexdigest() mock_table.delete_item.return_value = { @@ -193,7 +207,9 @@ def test_returns_data_atomically(self, mock_time, manager, mock_table): assert call_kwargs["ReturnValues"] == "ALL_OLD" assert call_kwargs["ConditionExpression"] == "attribute_exists(PK)" - def test_returns_none_for_unknown_token(self, manager, mock_table): + def test_returns_none_for_unknown_token( + self, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_table.delete_item.side_effect = ClientError( {"Error": {"Code": "ConditionalCheckFailedException", "Message": ""}}, "DeleteItem", @@ -202,7 +218,9 @@ def test_returns_none_for_unknown_token(self, manager, mock_table): assert result is None @patch("src.services.oauth_code_manager.time") - def test_returns_none_for_expired_token(self, mock_time, manager, mock_table): + def test_returns_none_for_expired_token( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 token_hash = hashlib.sha256(b"token123").hexdigest() mock_table.delete_item.return_value = { @@ -220,7 +238,9 @@ def test_returns_none_for_expired_token(self, mock_time, manager, mock_table): class TestConsumeAuthorizationCodeEdgeCases: @patch("src.services.oauth_code_manager.time") - def test_reraises_non_conditional_error(self, mock_time, manager, mock_table): + def test_reraises_non_conditional_error( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 mock_table.delete_item.side_effect = ClientError( {"Error": {"Code": "InternalServerError", "Message": ""}}, @@ -231,7 +251,9 @@ def test_reraises_non_conditional_error(self, mock_time, manager, mock_table): assert exc_info.value.response["Error"]["Code"] == "InternalServerError" @patch("src.services.oauth_code_manager.time") - def test_returns_none_for_empty_attributes(self, mock_time, manager, mock_table): + def test_returns_none_for_empty_attributes( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 mock_table.delete_item.return_value = {} result = manager.consume_authorization_code("abc") @@ -240,7 +262,9 @@ def test_returns_none_for_empty_attributes(self, mock_time, manager, mock_table) class TestConsumeRefreshTokenEdgeCases: @patch("src.services.oauth_code_manager.time") - def test_reraises_non_conditional_error(self, mock_time, manager, mock_table): + def test_reraises_non_conditional_error( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 mock_table.delete_item.side_effect = ClientError( {"Error": {"Code": "InternalServerError", "Message": ""}}, @@ -251,7 +275,9 @@ def test_reraises_non_conditional_error(self, mock_time, manager, mock_table): assert exc_info.value.response["Error"]["Code"] == "InternalServerError" @patch("src.services.oauth_code_manager.time") - def test_returns_none_for_empty_attributes(self, mock_time, manager, mock_table): + def test_returns_none_for_empty_attributes( + self, mock_time: MagicMock, manager: OAuthCodeManager, mock_table: MagicMock + ) -> None: mock_time.time.return_value = 1000000.0 mock_table.delete_item.return_value = {} result = manager.consume_refresh_token("hash") @@ -259,7 +285,7 @@ def test_returns_none_for_empty_attributes(self, mock_time, manager, mock_table) class TestVerifyCodeChallenge: - def test_valid_s256_challenge(self, manager): + def test_valid_s256_challenge(self, manager: OAuthCodeManager) -> None: # verifier -> SHA256 -> base64url = challenge import base64 @@ -268,21 +294,21 @@ def test_valid_s256_challenge(self, manager): challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode("ascii") assert manager.verify_code_challenge(verifier, challenge, "S256") is True - def test_invalid_s256_challenge(self, manager): + def test_invalid_s256_challenge(self, manager: OAuthCodeManager) -> None: assert manager.verify_code_challenge("wrong", "invalid_challenge", "S256") is False - def test_plain_challenge(self, manager): + def test_plain_challenge(self, manager: OAuthCodeManager) -> None: verifier = "plain-challenge-value" assert manager.verify_code_challenge(verifier, verifier, "plain") is True - def test_plain_challenge_mismatch(self, manager): + def test_plain_challenge_mismatch(self, manager: OAuthCodeManager) -> None: assert manager.verify_code_challenge("a", "b", "plain") is False - def test_unsupported_method_returns_false(self, manager): + def test_unsupported_method_returns_false(self, manager: OAuthCodeManager) -> None: assert manager.verify_code_challenge("v", "c", "unsupported") is False class TestDeleteRefreshToken: - def test_deletes_by_hash(self, manager, mock_table): + def test_deletes_by_hash(self, manager: OAuthCodeManager, mock_table: MagicMock) -> None: manager.delete_refresh_token("somehash") mock_table.delete_item.assert_called_once() diff --git a/lambda/tests/services/test_oidc_validator.py b/lambda/tests/services/test_oidc_validator.py index 15877909..923f1806 100644 --- a/lambda/tests/services/test_oidc_validator.py +++ b/lambda/tests/services/test_oidc_validator.py @@ -1,6 +1,7 @@ """Unit tests for OIDCValidator""" from datetime import datetime, timezone +from typing import Dict from unittest.mock import MagicMock, patch import jwt @@ -21,22 +22,22 @@ def current_timestamp() -> int: @pytest.fixture -def provider_url(): +def provider_url() -> str: return "https://auth.example.com" @pytest.fixture -def client_id(): +def client_id() -> str: return "test-client-id" @pytest.fixture -def validator(provider_url, client_id): +def validator(provider_url: str, client_id: str) -> OIDCValidator: return OIDCValidator(provider_url, client_id, user_agent="foobar", metrics=MagicMock()) @pytest.fixture -def mock_provider_config(): +def mock_provider_config() -> Dict[str, str]: return { "issuer": "https://auth.example.com", "jwks_uri": "https://auth.example.com/.well-known/jwks.json", @@ -49,24 +50,24 @@ def mock_provider_config(): class TestOIDCValidatorInit: """Test OIDCValidator initialization""" - def test_init_strips_trailing_slash(self, client_id): + def test_init_strips_trailing_slash(self, client_id: str) -> None: """Test that trailing slash is stripped from provider URL""" validator = OIDCValidator( "https://auth.example.com/", client_id, user_agent="foobar", metrics=MagicMock() ) assert validator.provider_url == "https://auth.example.com" - def test_init_stores_client_id(self, provider_url, client_id): + def test_init_stores_client_id(self, provider_url: str, client_id: str) -> None: """Test that client_id is stored correctly""" validator = OIDCValidator(provider_url, client_id, user_agent="foobar", metrics=MagicMock()) assert validator.client_id == client_id - def test_init_default_clock_skew_tolerance(self, provider_url, client_id): + def test_init_default_clock_skew_tolerance(self, provider_url: str, client_id: str) -> None: """Test that default clock_skew_tolerance is 300 seconds""" validator = OIDCValidator(provider_url, client_id, user_agent="foobar", metrics=MagicMock()) assert validator.clock_skew_tolerance == 300 - def test_init_custom_clock_skew_tolerance(self, provider_url, client_id): + def test_init_custom_clock_skew_tolerance(self, provider_url: str, client_id: str) -> None: """Test that custom clock_skew_tolerance is stored correctly""" validator = OIDCValidator( provider_url, @@ -77,12 +78,12 @@ def test_init_custom_clock_skew_tolerance(self, provider_url, client_id): ) assert validator.clock_skew_tolerance == 600 - def test_init_default_cache_ttl_seconds(self, provider_url, client_id): + def test_init_default_cache_ttl_seconds(self, provider_url: str, client_id: str) -> None: """Test that default cache_ttl_seconds is 3600 seconds""" validator = OIDCValidator(provider_url, client_id, user_agent="foobar", metrics=MagicMock()) assert validator.cache_ttl_seconds == 3600 - def test_init_custom_cache_ttl_seconds(self, provider_url, client_id): + def test_init_custom_cache_ttl_seconds(self, provider_url: str, client_id: str) -> None: """Test that custom cache_ttl_seconds is stored correctly""" validator = OIDCValidator( provider_url, @@ -97,7 +98,9 @@ def test_init_custom_cache_ttl_seconds(self, provider_url, client_id): class TestDiscoverProviderConfig: """Test discover_provider_config method""" - def test_discover_provider_config_success(self, validator, mock_provider_config): + def test_discover_provider_config_success( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test successful provider config discovery""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -115,7 +118,9 @@ def test_discover_provider_config_success(self, validator, mock_provider_config) headers={"User-Agent": "foobar"}, ) - def test_discover_provider_config_caching(self, validator, mock_provider_config): + def test_discover_provider_config_caching( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test that provider config is cached""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -132,7 +137,9 @@ def test_discover_provider_config_caching(self, validator, mock_provider_config) # Should only call once due to caching assert mock_get.call_count == 1 - def test_discover_provider_config_cache_expiry(self, validator, mock_provider_config): + def test_discover_provider_config_cache_expiry( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test that cache expires after TTL""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -154,8 +161,8 @@ def test_discover_provider_config_cache_expiry(self, validator, mock_provider_co assert mock_get.call_count == 2 def test_discover_provider_config_custom_cache_ttl( - self, provider_url, client_id, mock_provider_config - ): + self, provider_url: str, client_id: str, mock_provider_config: Dict[str, str] + ) -> None: """Test that custom cache TTL is respected""" validator = OIDCValidator( provider_url, @@ -191,9 +198,9 @@ def test_discover_provider_config_custom_cache_ttl( assert mock_get.call_count == 2 - def test_discover_provider_config_timeout(self, validator): + def test_discover_provider_config_timeout(self, validator: OIDCValidator) -> None: """Test ServiceUnavailableError on timeout""" - import requests # type: ignore[import-untyped] + import requests with patch("src.services.oidc_validator.requests.get") as mock_get: mock_get.side_effect = requests.exceptions.Timeout() @@ -203,9 +210,9 @@ def test_discover_provider_config_timeout(self, validator): assert "timed out" in str(exc_info.value.message) - def test_discover_provider_config_connection_error(self, validator): + def test_discover_provider_config_connection_error(self, validator: OIDCValidator) -> None: """Test ServiceUnavailableError on connection error""" - import requests # type: ignore[import-untyped] + import requests with patch("src.services.oidc_validator.requests.get") as mock_get: mock_get.side_effect = requests.exceptions.ConnectionError() @@ -215,9 +222,9 @@ def test_discover_provider_config_connection_error(self, validator): assert "unreachable" in str(exc_info.value.message) - def test_discover_provider_config_http_error(self, validator): + def test_discover_provider_config_http_error(self, validator: OIDCValidator) -> None: """Test ServiceUnavailableError on HTTP error""" - import requests # type: ignore[import-untyped] + import requests with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -232,7 +239,7 @@ def test_discover_provider_config_http_error(self, validator): assert "returned error" in str(exc_info.value.message) - def test_discover_provider_config_invalid_json(self, validator): + def test_discover_provider_config_invalid_json(self, validator: OIDCValidator) -> None: """Test ServiceUnavailableError on invalid config""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -249,7 +256,9 @@ def test_discover_provider_config_invalid_json(self, validator): class TestValidateToken: """Test validate_token method""" - def test_validate_token_success(self, validator, mock_provider_config): + def test_validate_token_success( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test successful token validation""" mock_claims = { "sub": "user123", @@ -281,7 +290,9 @@ def test_validate_token_success(self, validator, mock_provider_config): assert claims.aud == "test-client-id" assert claims.email == "user@example.com" - def test_validate_token_expired(self, validator, mock_provider_config): + def test_validate_token_expired( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidCredentialsError on expired token""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -302,7 +313,9 @@ def test_validate_token_expired(self, validator, mock_provider_config): assert "expired" in str(exc_info.value.message) - def test_validate_token_invalid_audience(self, validator, mock_provider_config): + def test_validate_token_invalid_audience( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidCredentialsError on invalid audience""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -323,7 +336,9 @@ def test_validate_token_invalid_audience(self, validator, mock_provider_config): assert "audience" in str(exc_info.value.message) - def test_validate_token_invalid_issuer(self, validator, mock_provider_config): + def test_validate_token_invalid_issuer( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidCredentialsError on invalid issuer""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -344,7 +359,9 @@ def test_validate_token_invalid_issuer(self, validator, mock_provider_config): assert "issuer" in str(exc_info.value.message) - def test_validate_token_missing_sub_claim(self, validator, mock_provider_config): + def test_validate_token_missing_sub_claim( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidCredentialsError when sub claim is missing""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -365,7 +382,9 @@ def test_validate_token_missing_sub_claim(self, validator, mock_provider_config) assert "missing required claim" in str(exc_info.value.message).lower() - def test_validate_token_invalid_signature(self, validator, mock_provider_config): + def test_validate_token_invalid_signature( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidTokenError on invalid signature""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -386,7 +405,9 @@ def test_validate_token_invalid_signature(self, validator, mock_provider_config) assert "Invalid token" in str(exc_info.value.message) - def test_validate_token_jwk_client_error(self, validator, mock_provider_config): + def test_validate_token_jwk_client_error( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidTokenError when JWK client fails""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -404,7 +425,9 @@ def test_validate_token_jwk_client_error(self, validator, mock_provider_config): assert "signing key" in str(exc_info.value.message) - def test_validate_token_audience_as_list(self, validator, mock_provider_config): + def test_validate_token_audience_as_list( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test handling of audience claim as list""" mock_claims = { "sub": "user123", @@ -433,7 +456,9 @@ def test_validate_token_audience_as_list(self, validator, mock_provider_config): # Should take first audience from list assert claims.aud == "test-client-id" - def test_validate_token_empty_sub_claim(self, validator, mock_provider_config): + def test_validate_token_empty_sub_claim( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidCredentialsError when sub claim is empty string""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -460,7 +485,9 @@ def test_validate_token_empty_sub_claim(self, validator, mock_provider_config): assert "sub claim" in str(exc_info.value.message) - def test_validate_token_empty_audience_list(self, validator, mock_provider_config): + def test_validate_token_empty_audience_list( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test handling of empty audience list""" mock_claims = { "sub": "user123", @@ -489,7 +516,9 @@ def test_validate_token_empty_audience_list(self, validator, mock_provider_confi # Should return empty string for empty audience list assert claims.aud == "" - def test_validate_token_unexpected_exception(self, validator, mock_provider_config): + def test_validate_token_unexpected_exception( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidTokenError on unexpected exception""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -511,7 +540,7 @@ def test_validate_token_unexpected_exception(self, validator, mock_provider_conf assert "Token validation failed" in str(exc_info.value.message) - def test_validate_token_service_unavailable_reraise(self, validator): + def test_validate_token_service_unavailable_reraise(self, validator: OIDCValidator) -> None: """Test ServiceUnavailableError is re-raised during token validation""" with patch.object(validator, "discover_provider_config") as mock_discover: mock_discover.side_effect = ServiceUnavailableError("Provider unreachable") @@ -521,7 +550,9 @@ def test_validate_token_service_unavailable_reraise(self, validator): assert "Provider unreachable" in str(exc_info.value.message) - def test_validate_token_timestamp_within_tolerance(self, validator, mock_provider_config): + def test_validate_token_timestamp_within_tolerance( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test successful validation when timestamp is within tolerance""" current_time = int(datetime.now(timezone.utc).timestamp()) mock_claims = { @@ -551,7 +582,9 @@ def test_validate_token_timestamp_within_tolerance(self, validator, mock_provide assert claims.sub == "user123" assert claims.iat == current_time - 100 - def test_validate_token_timestamp_exceeds_tolerance(self, validator, mock_provider_config): + def test_validate_token_timestamp_exceeds_tolerance( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidTimestampError when timestamp exceeds tolerance""" from src.shared.exceptions import InvalidTimestampError @@ -585,8 +618,8 @@ def test_validate_token_timestamp_exceeds_tolerance(self, validator, mock_provid assert "300 seconds" in str(exc_info.value.message) def test_validate_token_timestamp_future_exceeds_tolerance( - self, validator, mock_provider_config - ): + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test InvalidTimestampError when future timestamp exceeds tolerance""" from src.shared.exceptions import InvalidTimestampError @@ -618,7 +651,9 @@ def test_validate_token_timestamp_future_exceeds_tolerance( assert "400 seconds" in str(exc_info.value.message) - def test_validate_token_custom_tolerance(self, provider_url, client_id, mock_provider_config): + def test_validate_token_custom_tolerance( + self, provider_url: str, client_id: str, mock_provider_config: Dict[str, str] + ) -> None: """Test timestamp validation with custom tolerance""" validator = OIDCValidator( provider_url, @@ -654,7 +689,9 @@ def test_validate_token_custom_tolerance(self, provider_url, client_id, mock_pro assert claims.sub == "user123" - def test_validate_token_no_iat_claim(self, validator, mock_provider_config): + def test_validate_token_no_iat_claim( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test validation succeeds when iat claim is missing (optional validation)""" current_time = int(datetime.now(timezone.utc).timestamp()) mock_claims = { @@ -684,7 +721,9 @@ def test_validate_token_no_iat_claim(self, validator, mock_provider_config): assert claims.sub == "user123" - def test_validate_token_uses_current_time(self, validator, mock_provider_config): + def test_validate_token_uses_current_time( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test that validation uses current time internally""" mock_claims = { "sub": "user123", @@ -716,7 +755,9 @@ def test_validate_token_uses_current_time(self, validator, mock_provider_config) class TestGetJwkClient: """Test _get_jwk_client method""" - def test_get_jwk_client_creates_client(self, validator, mock_provider_config): + def test_get_jwk_client_creates_client( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test that _get_jwk_client creates PyJWKClient on first call""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -738,7 +779,9 @@ def test_get_jwk_client_creates_client(self, validator, mock_provider_config): headers={"User-Agent": "foobar"}, ) - def test_get_jwk_client_custom_cache_ttl(self, provider_url, client_id, mock_provider_config): + def test_get_jwk_client_custom_cache_ttl( + self, provider_url: str, client_id: str, mock_provider_config: Dict[str, str] + ) -> None: """Test that _get_jwk_client uses custom cache TTL""" validator = OIDCValidator( provider_url, @@ -768,7 +811,9 @@ def test_get_jwk_client_custom_cache_ttl(self, provider_url, client_id, mock_pro headers={"User-Agent": "foobar"}, ) - def test_get_jwk_client_caches_client(self, validator, mock_provider_config): + def test_get_jwk_client_caches_client( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test that _get_jwk_client returns cached client on subsequent calls""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() @@ -793,7 +838,9 @@ def test_get_jwk_client_caches_client(self, validator, mock_provider_config): class TestClearCache: """Test clear_cache method""" - def test_clear_cache(self, validator, mock_provider_config): + def test_clear_cache( + self, validator: OIDCValidator, mock_provider_config: Dict[str, str] + ) -> None: """Test that clear_cache resets all cached data""" with patch("src.services.oidc_validator.requests.get") as mock_get: mock_response = MagicMock() diff --git a/lambda/tests/services/test_storage_hawk_middleware.py b/lambda/tests/services/test_storage_hawk_middleware.py index cd68e24d..80f10c69 100644 --- a/lambda/tests/services/test_storage_hawk_middleware.py +++ b/lambda/tests/services/test_storage_hawk_middleware.py @@ -1,5 +1,6 @@ """Tests for HawkAuthMiddleware (storage mode)""" +from typing import Dict, Optional from unittest.mock import MagicMock import pytest @@ -10,13 +11,13 @@ def _make_app( - auth_header=None, - method="GET", - path="/1.5/123/storage/bookmarks", - query_params=None, - domain_name="storage.example.com", - path_params=None, -): + auth_header: Optional[str] = None, + method: str = "GET", + path: str = "/1.5/123/storage/bookmarks", + query_params: Optional[Dict[str, str]] = None, + domain_name: str = "storage.example.com", + path_params: Optional[Dict[str, str]] = None, +) -> MagicMock: """Build a mock APIGatewayRestResolver app with current_event.""" app = MagicMock() headers = {} @@ -42,7 +43,7 @@ def _make_app( class TestHawkAuthMiddlewareSuccess: - def test_success_injects_hawk_uid(self): + def test_success_injects_hawk_uid(self) -> None: """Successful Hawk validation injects hawk_uid and calls next.""" hawk_service = MagicMock() creds = HawkCredentials( @@ -77,7 +78,7 @@ def test_success_injects_hawk_uid(self): mock_next.assert_called_once_with(app) assert result.status_code == 200 - def test_lowercase_authorization_header(self): + def test_lowercase_authorization_header(self) -> None: """Middleware finds lowercase 'authorization' header.""" hawk_service = MagicMock() creds = HawkCredentials( @@ -112,7 +113,7 @@ def test_lowercase_authorization_header(self): class TestHawkAuthMiddlewareFailure: - def test_missing_auth_header_raises_error(self): + def test_missing_auth_header_raises_error(self) -> None: """Missing Authorization header raises HawkAuthenticationError.""" hawk_service = MagicMock() middleware = HawkAuthMiddleware(hawk_service=hawk_service, metrics=MagicMock()) @@ -126,7 +127,7 @@ def test_missing_auth_header_raises_error(self): hawk_service.validate.assert_not_called() mock_next.assert_not_called() - def test_hawk_validation_exception_raises_error(self): + def test_hawk_validation_exception_raises_error(self) -> None: """Exception from hawk_service.validate raises HawkAuthenticationError.""" hawk_service = MagicMock() hawk_service.validate.side_effect = Exception("MacMismatch") @@ -142,7 +143,7 @@ def test_hawk_validation_exception_raises_error(self): class TestHawkAuthMiddlewareQueryString: - def test_query_string_included_in_path(self): + def test_query_string_included_in_path(self) -> None: """Query string parameters are appended to path for MAC validation.""" hawk_service = MagicMock() creds = HawkCredentials( @@ -173,7 +174,7 @@ def test_query_string_included_in_path(self): class TestHawkAuthMiddlewareHostFallback: - def test_domain_name_attribute_error_falls_back_to_host_header(self): + def test_domain_name_attribute_error_falls_back_to_host_header(self) -> None: """When request_context.domain_name raises, falls back to host header.""" hawk_service = MagicMock() creds = HawkCredentials(user_id="user1", generation=0, expiry=9999999999, hawk_id="hid") @@ -206,19 +207,19 @@ def test_domain_name_attribute_error_falls_back_to_host_header(self): class TestHawkAuthMiddlewareInit: - def test_requires_hawk_service_or_token_manager(self): + def test_requires_hawk_service_or_token_manager(self) -> None: """Middleware requires at least one of hawk_service or token_manager.""" with pytest.raises(ValueError, match="Either hawk_service or token_manager"): HawkAuthMiddleware(metrics=MagicMock()) - def test_session_mode_with_token_manager(self): + def test_session_mode_with_token_manager(self) -> None: """Middleware can be initialized with token_manager for session auth.""" token_manager = MagicMock() middleware = HawkAuthMiddleware(token_manager=token_manager, metrics=MagicMock()) assert middleware._token_manager is token_manager assert middleware._hawk_service is None - def test_storage_mode_with_hawk_service(self): + def test_storage_mode_with_hawk_service(self) -> None: """Middleware can be initialized with hawk_service for storage auth.""" hawk_service = MagicMock() middleware = HawkAuthMiddleware(hawk_service=hawk_service, metrics=MagicMock()) @@ -227,7 +228,7 @@ def test_storage_mode_with_hawk_service(self): class TestHawkAuthMiddlewareSessionMode: - def test_session_hawk_success_injects_hawk_uid(self): + def test_session_hawk_success_injects_hawk_uid(self) -> None: """Session Hawk validation injects hawk_uid and calls next.""" token_manager = MagicMock() token_manager.verify_session_hawk.return_value = "uid123" @@ -255,7 +256,7 @@ def test_session_hawk_success_injects_hawk_uid(self): mock_next.assert_called_once_with(app) assert result.status_code == 200 - def test_session_hawk_invalid_token_raises_error(self): + def test_session_hawk_invalid_token_raises_error(self) -> None: """Invalid session token raises HawkAuthenticationError.""" token_manager = MagicMock() token_manager.verify_session_hawk.return_value = None diff --git a/lambda/tests/services/test_storage_manager.py b/lambda/tests/services/test_storage_manager.py index 4b1b34ab..e7c704eb 100644 --- a/lambda/tests/services/test_storage_manager.py +++ b/lambda/tests/services/test_storage_manager.py @@ -1,8 +1,10 @@ """Unit tests for StorageManager with DynamoDB stubber""" from decimal import Decimal +from typing import TYPE_CHECKING import pytest +from botocore.stub import Stubber from src.services.storage_manager import StorageManager from src.shared.exceptions import ( @@ -11,16 +13,21 @@ ) from src.shared.models import BasicStorageObject +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + class TestStorageManager: """Test StorageManager DynamoDB operations""" @pytest.fixture - def storage_manager(self, dynamodb_table): + def storage_manager(self, dynamodb_table: "Table") -> StorageManager: """Create StorageManager instance with stubbed table""" return StorageManager(table=dynamodb_table) - def test_get_collection_success(self, storage_manager, dynamodb_stubber, storage_table_name): + def test_get_collection_success( + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test successful collection retrieval""" dynamodb_stubber.add_response( "get_item", @@ -50,7 +57,9 @@ def test_get_collection_success(self, storage_manager, dynamodb_stubber, storage assert collection.count == 5 assert collection.usage == 1024 - def test_get_collection_not_found(self, storage_manager, dynamodb_stubber, storage_table_name): + def test_get_collection_not_found( + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test collection not found""" dynamodb_stubber.add_response( "get_item", @@ -68,8 +77,8 @@ def test_get_collection_not_found(self, storage_manager, dynamodb_stubber, stora storage_manager.get_collection("test-user-123", "nonexistent") def test_get_storage_object_success( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test successful storage object retrieval""" dynamodb_stubber.add_response( "get_item", @@ -101,8 +110,8 @@ def test_get_storage_object_success( assert obj.sortindex == 100 def test_get_storage_object_not_found( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test storage object not found""" dynamodb_stubber.add_response( "get_item", @@ -120,8 +129,8 @@ def test_get_storage_object_not_found( storage_manager.get_storage_object("test-user-123", "bookmarks", "nonexistent") def test_get_storage_object_without_optional_fields( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test retrieval of storage object without sortindex and ttl""" dynamodb_stubber.add_response( "get_item", @@ -152,13 +161,12 @@ def test_get_storage_object_without_optional_fields( def test_create_or_update_collection_without_objects( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_timestamp_datetime, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test creating collection without objects""" # Collection existence check — not found, so this is a new collection dynamodb_stubber.add_response( @@ -196,7 +204,7 @@ def test_create_or_update_collection_without_objects( ) assert collection.name == "bookmarks" - assert collection.modified == mock_timestamp_datetime + assert collection.modified == mock_timestamp assert collection.count == 0 assert collection.usage == 0 assert batch_result.model_dump()["success"] == [] @@ -204,12 +212,12 @@ def test_create_or_update_collection_without_objects( def test_create_or_update_collection_with_objects( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test creating collection with objects""" objects = [ BasicStorageObject( @@ -271,12 +279,12 @@ def test_create_or_update_collection_with_objects( def test_update_collection( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating collection""" # Stub get_collection dynamodb_stubber.add_response( @@ -329,8 +337,8 @@ def test_update_collection( assert batch_result.model_dump()["success"] == ["obj1"] def test_update_collection_not_found( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test updating non-existent collection raises error""" dynamodb_stubber.add_response( "get_item", @@ -357,12 +365,12 @@ def test_update_collection_not_found( def test_delete_collection( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test deleting collection""" # Stub get_collection to verify it exists @@ -411,7 +419,9 @@ def test_delete_collection( modified = storage_manager.delete_collection("test-user-123", "bookmarks") assert modified == mock_timestamp - def test_list_collections(self, storage_manager, dynamodb_stubber, storage_table_name): + def test_list_collections( + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test listing all collections""" dynamodb_stubber.add_response( "query", @@ -453,7 +463,9 @@ def test_list_collections(self, storage_manager, dynamodb_stubber, storage_table assert collections[1].name == "history" assert collections[1].count == 10 - def test_list_collections_empty(self, storage_manager, dynamodb_stubber, storage_table_name): + def test_list_collections_empty( + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test listing collections when none exist""" dynamodb_stubber.add_response( "query", @@ -470,8 +482,8 @@ def test_list_collections_empty(self, storage_manager, dynamodb_stubber, storage assert collections == [] def test_list_collections_with_pagination( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test listing collections with pagination""" # Stub first page with LastEvaluatedKey dynamodb_stubber.add_response( @@ -550,7 +562,9 @@ def test_list_collections_with_pagination( assert collections[2].name == "passwords" assert collections[2].count == 15 - def test_get_collection_objects(self, storage_manager, dynamodb_stubber, storage_table_name): + def test_get_collection_objects( + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects from collection""" dynamodb_stubber.add_response( "query", @@ -592,8 +606,8 @@ def test_get_collection_objects(self, storage_manager, dynamodb_stubber, storage assert result["last_modified"] == 1234567891.00 def test_get_collection_objects_with_filters( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects with ID filter""" # With ids parameter, uses batch_get_item instead of query dynamodb_stubber.add_response( @@ -622,8 +636,8 @@ def test_get_collection_objects_with_filters( assert result["items"][0].id == "obj1" def test_get_collection_objects_pagination( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test pagination of collection objects""" dynamodb_stubber.add_response( "query", @@ -672,12 +686,12 @@ def test_get_collection_objects_pagination( def test_update_storage_object( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating storage object""" # Stub get_storage_object to verify it exists @@ -732,8 +746,12 @@ def test_update_storage_object( assert updated_obj.sortindex == 100 def test_update_storage_object_not_found( - self, storage_manager, dynamodb_stubber, storage_table_name, mock_get_current_timestamp - ): + self, + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_get_current_timestamp: None, + ) -> None: """Test updating non-existent object creates it (PUT semantics)""" dynamodb_stubber.add_response( "get_item", @@ -765,12 +783,12 @@ def test_update_storage_object_not_found( def test_delete_storage_object( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test deleting storage object""" # Stub get_storage_object to verify it exists @@ -811,8 +829,8 @@ def test_delete_storage_object( assert modified == mock_timestamp def test_delete_storage_object_not_found( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test deleting non-existent object raises error""" dynamodb_stubber.add_response( "get_item", @@ -830,8 +848,8 @@ def test_delete_storage_object_not_found( storage_manager.delete_storage_object("test-user-123", "bookmarks", "nonexistent") def test_get_collection_client_error( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test get_collection with ClientError for ResourceNotFoundException""" dynamodb_stubber.add_client_error( "get_item", @@ -843,8 +861,8 @@ def test_get_collection_client_error( storage_manager.get_collection("test-user-123", "bookmarks") def test_get_storage_object_client_error( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test get_storage_object with ClientError for ResourceNotFoundException""" dynamodb_stubber.add_client_error( "get_item", @@ -857,12 +875,12 @@ def test_get_storage_object_client_error( def test_create_collection_with_failed_object( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test creating collection when batch_write_item fails entirely. With batch_writer(), per-item error tracking no longer works the same way. @@ -903,12 +921,12 @@ def test_create_collection_with_failed_object( def test_update_collection_with_failed_object( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating collection when batch_write_item fails entirely. With batch_writer(), per-item error tracking no longer works. @@ -963,8 +981,8 @@ def test_update_collection_with_failed_object( storage_manager.update_collection("test-user-123", "bookmarks", objects) def test_get_collection_objects_with_newer_filter( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects with newer timestamp filter""" # newer/older without ids now pushes FilterExpression to DynamoDB dynamodb_stubber.add_response( @@ -991,8 +1009,8 @@ def test_get_collection_objects_with_newer_filter( assert result["items"][0].id == "obj1" def test_get_collection_objects_with_older_filter( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects with older timestamp filter""" # older without ids now pushes FilterExpression to DynamoDB dynamodb_stubber.add_response( @@ -1019,8 +1037,8 @@ def test_get_collection_objects_with_older_filter( assert result["items"][0].id == "obj2" def test_get_collection_objects_sort_oldest( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects sorted by oldest first""" dynamodb_stubber.add_response( "query", @@ -1058,8 +1076,8 @@ def test_get_collection_objects_sort_oldest( assert result["items"][1].id == "obj1" def test_get_collection_objects_sort_index( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects sorted by sortindex""" dynamodb_stubber.add_response( "query", @@ -1100,12 +1118,12 @@ def test_get_collection_objects_sort_index( def test_update_storage_object_with_all_fields( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating storage object with all optional fields""" # Stub get_storage_object to verify it exists dynamodb_stubber.add_response( @@ -1146,8 +1164,8 @@ def test_update_storage_object_with_all_fields( assert updated_obj.sortindex == 150 def test_get_collection_client_error_other( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test get_collection with other ClientError""" dynamodb_stubber.add_client_error( "get_item", @@ -1159,8 +1177,8 @@ def test_get_collection_client_error_other( storage_manager.get_collection("test-user-123", "bookmarks") def test_get_storage_object_client_error_other( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test get_storage_object with other ClientError""" dynamodb_stubber.add_client_error( "get_item", @@ -1172,8 +1190,8 @@ def test_get_storage_object_client_error_other( storage_manager.get_storage_object("test-user-123", "bookmarks", "obj123") def test_get_collection_objects_empty_result( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects when collection is empty""" dynamodb_stubber.add_response( "query", @@ -1196,12 +1214,12 @@ def test_get_collection_objects_empty_result( def test_update_collection_with_mixed_success_fail( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating collection with multiple objects succeeding. With batch_writer(), per-item error tracking is no longer possible. @@ -1266,12 +1284,12 @@ def test_update_collection_with_mixed_success_fail( def test_update_collection_with_sortindex_and_ttl( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating collection with objects that have sortindex and ttl""" # Stub get_collection dynamodb_stubber.add_response( @@ -1325,12 +1343,12 @@ def test_update_collection_with_sortindex_and_ttl( def test_update_storage_object_preserves_ttl( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating storage object preserves ttl when not provided""" # Stub get_storage_object with ttl dynamodb_stubber.add_response( @@ -1368,12 +1386,12 @@ def test_update_storage_object_preserves_ttl( def test_update_storage_object_without_sortindex( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating object without providing sortindex to test branch""" # Stub get_storage_object - object without sortindex but with ttl dynamodb_stubber.add_response( @@ -1412,8 +1430,8 @@ def test_update_storage_object_without_sortindex( assert updated_obj.sortindex is None def test_get_collection_objects_invalid_sort( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects with invalid sort parameter (should not sort)""" dynamodb_stubber.add_response( "query", @@ -1454,12 +1472,12 @@ def test_get_collection_objects_invalid_sort( def test_update_storage_object_with_only_sortindex( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating object with only sortindex to test branch""" # Stub get_storage_object dynamodb_stubber.add_response( @@ -1510,12 +1528,12 @@ def test_update_storage_object_with_only_sortindex( def test_update_storage_object_with_sortindex_no_ttl( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating object with sortindex but no ttl to cover branch""" # Stub get_storage_object - has sortindex and ttl dynamodb_stubber.add_response( @@ -1558,12 +1576,12 @@ def test_update_storage_object_with_sortindex_no_ttl( def test_update_storage_object_sortindex_without_ttl( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test updating object that has sortindex but no ttl - covering branch 401->405""" # Stub get_storage_object - object WITH sortindex but NO ttl dynamodb_stubber.add_response( @@ -1612,8 +1630,8 @@ def test_update_storage_object_sortindex_without_ttl( assert updated_obj.sortindex == 100 def test_create_collection_batch_limit_exceeded( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test creating collection with too many objects raises ServerLimitExceededException""" from src.shared.exceptions import ServerLimitExceededException @@ -1631,8 +1649,8 @@ def test_create_collection_batch_limit_exceeded( storage_manager.create_or_update_collection("test-user-123", "bookmarks", objects) def test_create_collection_batch_size_exceeded( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test creating collection with too large payload raises ServerLimitExceededException""" from src.shared.exceptions import ServerLimitExceededException @@ -1651,11 +1669,11 @@ def test_create_collection_batch_size_exceeded( def test_create_collection_precondition_create_only_fails( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_get_current_timestamp: None, + ) -> None: """Test create-only mode fails when collection exists""" from src.shared.exceptions import PreconditionFailedException @@ -1688,11 +1706,11 @@ def test_create_collection_precondition_create_only_fails( def test_create_collection_precondition_modified_since( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_get_current_timestamp: None, + ) -> None: """Test precondition fails when collection modified since timestamp""" from src.shared.exceptions import PreconditionFailedException @@ -1724,8 +1742,8 @@ def test_create_collection_precondition_modified_since( ) def test_update_collection_batch_limit_exceeded( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test updating collection with too many objects raises ServerLimitExceededException""" from src.shared.exceptions import ServerLimitExceededException @@ -1743,8 +1761,8 @@ def test_update_collection_batch_limit_exceeded( storage_manager.update_collection("test-user-123", "bookmarks", objects) def test_update_collection_batch_size_exceeded( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test updating collection with too large payload raises ServerLimitExceededException""" from src.shared.exceptions import ServerLimitExceededException @@ -1763,11 +1781,11 @@ def test_update_collection_batch_size_exceeded( def test_update_collection_precondition_modified_since( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_get_current_timestamp: None, + ) -> None: """Test update precondition fails when collection modified since timestamp""" from src.shared.exceptions import PreconditionFailedException @@ -1799,8 +1817,8 @@ def test_update_collection_precondition_modified_since( ) def test_get_collection_objects_ids_limit_exceeded( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting objects with too many IDs raises ValidationException""" from src.shared.exceptions import ValidationException @@ -1812,11 +1830,11 @@ def test_get_collection_objects_ids_limit_exceeded( def test_update_storage_object_precondition_create_only_fails( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_get_current_timestamp: None, + ) -> None: """Test create-only mode fails when object exists""" from src.shared.exceptions import PreconditionFailedException @@ -1852,11 +1870,11 @@ def test_update_storage_object_precondition_create_only_fails( def test_update_storage_object_precondition_modified_since( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_get_current_timestamp: None, + ) -> None: """Test update precondition fails when object modified since timestamp""" from src.shared.exceptions import PreconditionFailedException @@ -1892,11 +1910,11 @@ def test_update_storage_object_precondition_modified_since( def test_update_storage_object_precondition_nonexistent_object( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_get_current_timestamp: None, + ) -> None: """Test precondition fails when checking non-existent object with non-zero timestamp""" from src.shared.exceptions import PreconditionFailedException @@ -1924,12 +1942,12 @@ def test_update_storage_object_precondition_nonexistent_object( def test_delete_collection_objects( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test batch deleting multiple objects""" # Stub get_collection to verify collection exists (called once, result reused) dynamodb_stubber.add_response( @@ -1966,8 +1984,8 @@ def test_delete_collection_objects( assert modified == mock_timestamp def test_delete_collection_objects_limit_exceeded( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test batch delete with too many IDs raises ValidationException""" from src.shared.exceptions import ValidationException @@ -1979,12 +1997,12 @@ def test_delete_collection_objects_limit_exceeded( def test_delete_collection_objects_with_error( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test batch delete when batch_write_item fails entirely. With batch_writer(), per-item error tracking is no longer possible. @@ -2028,12 +2046,12 @@ def test_delete_collection_objects_with_error( def test_delete_all_storage( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test deleting all storage for a user via list_collections + delete_collection (no scan)""" # list_collections: GSI query returns one collection dynamodb_stubber.add_response( @@ -2108,12 +2126,12 @@ def test_delete_all_storage( def test_delete_all_storage_with_pagination( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test deleting all storage for a user with multiple collections via paginated GSI query""" # list_collections page 1: returns bookmarks, with LastEvaluatedKey dynamodb_stubber.add_response( @@ -2256,11 +2274,8 @@ def test_delete_all_storage_with_pagination( assert modified == mock_timestamp def test_get_quota( - self, - storage_manager, - dynamodb_stubber, - storage_table_name, - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test getting quota information""" # list_collections uses query with GSI dynamodb_stubber.add_response( @@ -2302,12 +2317,12 @@ def test_get_quota( def test_create_collection_precondition_passes( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test precondition passes when collection not modified since timestamp""" # Stub get_collection to return collection modified before the precondition timestamp dynamodb_stubber.add_response( @@ -2342,12 +2357,12 @@ def test_create_collection_precondition_passes( def test_update_collection_precondition_passes( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test update precondition passes when collection not modified since timestamp""" # Stub get_collection to return collection modified before the precondition timestamp dynamodb_stubber.add_response( @@ -2382,12 +2397,12 @@ def test_update_collection_precondition_passes( def test_update_storage_object_precondition_passes( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test update precondition passes when object not modified since timestamp""" # Stub get_storage_object to return object modified before the precondition timestamp dynamodb_stubber.add_response( @@ -2428,12 +2443,12 @@ def test_update_storage_object_precondition_passes( def test_update_collection_overwrites_existing_bso_usage_delta( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test that overwriting an existing BSO uses net delta (new_size - old_size), not new_size""" # Existing collection has usage=100, count=1 # Existing BSO "obj1" has payload "old" (3 bytes) @@ -2519,12 +2534,12 @@ def test_update_collection_overwrites_existing_bso_usage_delta( def test_update_collection_count_not_incremented_for_existing_bso( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """TDD: updating an existing BSO must not increment the collection count. Before fix: new_count = collection.count + len(success) → wrong @@ -2598,12 +2613,12 @@ def test_update_collection_count_not_incremented_for_existing_bso( def test_delete_collection_with_pagination( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """TDD: delete_collection must paginate when items span multiple query pages. Before fix: only first query page is deleted @@ -2675,12 +2690,12 @@ def test_delete_collection_with_pagination( def test_delete_all_storage_via_list_and_delete( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """TDD: delete_all_storage must use list_collections + delete_collection, not table.scan. Before fix: calls table.scan (full table read) @@ -2758,12 +2773,12 @@ def test_delete_all_storage_via_list_and_delete( def test_delete_all_storage_skips_concurrently_deleted_collection( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test that delete_all_storage silently skips collections deleted concurrently. If a collection is returned by list_collections but has already been deleted @@ -2813,12 +2828,12 @@ def test_delete_all_storage_skips_concurrently_deleted_collection( def test_create_or_update_collection_updates_existing_bso( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test create_or_update_collection when updating an already-existing BSO. When the collection already exists and the incoming object ID matches an @@ -2896,12 +2911,12 @@ def test_create_or_update_collection_updates_existing_bso( def test_create_or_update_collection_adds_new_bso_to_existing_collection( self, - storage_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_get_current_timestamp, - ): + storage_manager: StorageManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_get_current_timestamp: None, + ) -> None: """Test create_or_update_collection when adding a brand-new BSO to an existing collection. When collection_exists=True but the incoming object ID does not exist yet, @@ -2962,8 +2977,8 @@ def test_create_or_update_collection_adds_new_bso_to_existing_collection( assert batch_result.failed == {} def test_get_collection_objects_not_full( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test get_collection_objects with full=False uses ProjectionExpression""" dynamodb_stubber.add_response( "query", @@ -2989,11 +3004,8 @@ def test_get_collection_objects_not_full( assert result["items"][0].id == "obj1" def test_get_collection_objects_with_ids_and_newer( - self, - storage_manager, - dynamodb_stubber, - storage_table_name, - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test get_collection_objects with ids + newer applies client-side filter""" dynamodb_stubber.add_response( "batch_get_item", @@ -3033,11 +3045,8 @@ def test_get_collection_objects_with_ids_and_newer( assert result["items"][0].id == "obj2" def test_get_collection_objects_with_ids_and_older( - self, - storage_manager, - dynamodb_stubber, - storage_table_name, - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test get_collection_objects with ids + older applies client-side filter""" dynamodb_stubber.add_response( "batch_get_item", @@ -3077,8 +3086,8 @@ def test_get_collection_objects_with_ids_and_older( assert result["items"][0].id == "obj1" def test_get_collection_objects_query_pagination( - self, storage_manager, dynamodb_stubber, storage_table_name - ): + self, storage_manager: StorageManager, dynamodb_stubber: Stubber, storage_table_name: str + ) -> None: """Test get_collection_objects handles DynamoDB query pagination""" # First page with LastEvaluatedKey dynamodb_stubber.add_response( diff --git a/lambda/tests/services/test_token_generator.py b/lambda/tests/services/test_token_generator.py index 4fdbe654..e7c1c4cf 100644 --- a/lambda/tests/services/test_token_generator.py +++ b/lambda/tests/services/test_token_generator.py @@ -22,19 +22,19 @@ class TestTokenGenerator: """Test TokenGenerator orchestration logic with mocked HawkService""" @pytest.fixture - def mock_hawk_service(self): + def mock_hawk_service(self) -> MagicMock: """Mock HawkService for testing orchestration""" service = MagicMock() service.token_duration = 300 return service @pytest.fixture - def token_generator(self, storage_domain, mock_hawk_service): + def token_generator(self, storage_domain: str, mock_hawk_service: MagicMock) -> TokenGenerator: """TokenGenerator instance with mocked HawkService""" return TokenGenerator(storage_domain=storage_domain, hawk_service=mock_hawk_service) @pytest.fixture - def mock_hawk_credentials(self): + def mock_hawk_credentials(self) -> HawkCredentials: """Mock HawkCredentials returned by HawkService""" return HawkCredentials( user_id="user123", @@ -46,7 +46,7 @@ def mock_hawk_credentials(self): # ========== UID Generation Tests ========== - def test_generate_uid_consistency(self, token_generator): + def test_generate_uid_consistency(self, token_generator: TokenGenerator) -> None: """Test UID is consistent for same user_id and generation Validates: Requirements 4.1, 4.2 @@ -59,7 +59,7 @@ def test_generate_uid_consistency(self, token_generator): assert uid1 == uid2 - def test_generate_uid_different_users(self, token_generator): + def test_generate_uid_different_users(self, token_generator: TokenGenerator) -> None: """Test UID differs for different user_ids Validates: Requirements 4.1, 4.2 @@ -69,7 +69,7 @@ def test_generate_uid_different_users(self, token_generator): assert uid1 != uid2 - def test_generate_uid_changes_with_generation(self, token_generator): + def test_generate_uid_changes_with_generation(self, token_generator: TokenGenerator) -> None: """Test UID changes when generation changes (node reset) Validates: Requirements 2.4, 4.1 @@ -85,7 +85,7 @@ def test_generate_uid_changes_with_generation(self, token_generator): assert uid_gen1 != uid_gen2 assert uid_gen0 != uid_gen2 - def test_generate_uid_positive(self, token_generator): + def test_generate_uid_positive(self, token_generator: TokenGenerator) -> None: """Test UID is always positive Validates: Requirements 4.1, 4.2 @@ -101,8 +101,11 @@ def test_generate_uid_positive(self, token_generator): # ========== Token Generation Orchestration Tests ========== def test_generate_token_calls_hawk_service_correctly( - self, token_generator, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test generate_token calls HawkService.generate_hawk_credentials with correct params Validates: Requirements 4.1, 4.2, 4.3, 4.4 @@ -119,8 +122,11 @@ def test_generate_token_calls_hawk_service_correctly( mock_hawk_service.generate_hawk_credentials.assert_called_once_with(user_id, generation) def test_generate_token_stores_credentials_in_cache( - self, token_generator, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test generate_token stores credentials via HawkService.store_token_in_cache Validates: Requirements 4.1, 4.2, 4.5 @@ -137,8 +143,11 @@ def test_generate_token_stores_credentials_in_cache( mock_hawk_service.store_token_in_cache.assert_called_once_with(mock_hawk_credentials) def test_generate_token_returns_complete_response( - self, token_generator, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test generate_token returns TokenResponse with all required fields Validates: Requirements 1.1, 4.1, 4.2 @@ -161,8 +170,11 @@ def test_generate_token_returns_complete_response( assert token.hashalg is not None def test_generate_token_uses_hawk_credentials( - self, token_generator, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test generate_token uses HAWK credentials from HawkService Validates: Requirements 4.1, 4.2, 4.3, 4.4 @@ -180,8 +192,12 @@ def test_generate_token_uses_hawk_credentials( assert token.key == mock_hawk_credentials.hawk_key def test_generate_token_constructs_api_endpoint_correctly( - self, token_generator, storage_url, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + storage_url: str, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test generate_token constructs api_endpoint with correct format Validates: Requirements 2.3, 2.5 @@ -198,8 +214,11 @@ def test_generate_token_constructs_api_endpoint_correctly( assert token.api_endpoint == f"{storage_url}/1.5/{uid}" def test_generate_token_uses_provided_uid( - self, token_generator, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test generate_token uses the uid parameter provided (not generating its own) Validates: Requirements 2.1, 4.1 @@ -216,8 +235,12 @@ def test_generate_token_uses_provided_uid( assert token.uid == uid def test_generate_token_different_uids_different_endpoints( - self, token_generator, storage_url, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + storage_url: str, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test different uids result in different api_endpoints Validates: Requirements 2.2, 2.3 @@ -236,8 +259,11 @@ def test_generate_token_different_uids_different_endpoints( assert token2.api_endpoint == f"{storage_url}/1.5/222222" def test_generate_token_duration_from_hawk_service( - self, token_generator, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test generate_token uses duration from HawkService Validates: Requirements 1.5, 4.1 @@ -254,8 +280,11 @@ def test_generate_token_duration_from_hawk_service( assert token.duration == 300 def test_generate_token_hashalg_constant( - self, token_generator, mock_hawk_service, mock_hawk_credentials - ): + self, + token_generator: TokenGenerator, + mock_hawk_service: MagicMock, + mock_hawk_credentials: HawkCredentials, + ) -> None: """Test generate_token always uses sha256 hash algorithm Validates: Requirements 4.1, 4.2 @@ -273,7 +302,7 @@ def test_generate_token_hashalg_constant( # ========== Class Constants Tests ========== - def test_hash_algorithm_constant(self): + def test_hash_algorithm_constant(self) -> None: """Test HASH_ALGORITHM class constant is sha256 Validates: Requirements 4.1, 4.2 diff --git a/lambda/tests/services/test_user_manager.py b/lambda/tests/services/test_user_manager.py index 046d69d3..970d1b4d 100644 --- a/lambda/tests/services/test_user_manager.py +++ b/lambda/tests/services/test_user_manager.py @@ -1,31 +1,35 @@ """Unit tests for UserManager with DynamoDB stubber""" from decimal import Decimal +from typing import TYPE_CHECKING import pytest from botocore.exceptions import ClientError -from botocore.stub import ANY +from botocore.stub import ANY, Stubber from src.services.user_manager import UserManager from src.shared.exceptions import InvalidClientStateError, ServiceUnavailableError from src.shared.user import UserRecord +if TYPE_CHECKING: + from types_boto3_dynamodb.service_resource import Table + class TestUserManager: """Test UserManager DynamoDB operations""" @pytest.fixture - def user_manager(self, dynamodb_table): + def user_manager(self, dynamodb_table: "Table") -> UserManager: """Create UserManager instance with stubbed table""" return UserManager(table=dynamodb_table) def test_get_or_create_user_new_user( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + ) -> None: """Test creating a new user when user doesn't exist""" user_id = "user123456789" @@ -57,12 +61,12 @@ def test_get_or_create_user_new_user( def test_get_or_create_user_existing_user( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test retrieving existing user when conditional write fails""" user_id = "user123456789" existing_timestamp = 1234567800.00 @@ -102,10 +106,10 @@ def test_get_or_create_user_existing_user( def test_get_or_create_user_dynamodb_unavailable( self, - user_manager, - dynamodb_stubber, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + mock_datetime_now: None, + ) -> None: """Test ServiceUnavailableError when DynamoDB is unavailable""" dynamodb_stubber.add_client_error( "put_item", @@ -120,12 +124,12 @@ def test_get_or_create_user_dynamodb_unavailable( def test_increment_generation_success( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test successful generation increment""" user_id = "user123456789" @@ -159,9 +163,9 @@ def test_increment_generation_success( def test_increment_generation_dynamodb_unavailable( self, - user_manager, - dynamodb_stubber, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + ) -> None: """Test ServiceUnavailableError when DynamoDB is unavailable during increment""" dynamodb_stubber.add_client_error( "update_item", @@ -176,10 +180,10 @@ def test_increment_generation_dynamodb_unavailable( def test_validate_generation_valid( self, - user_manager, - dynamodb_stubber, - storage_table_name, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """Test generation validation when generation matches""" user_id = "user123456789" current_generation = 5 @@ -206,10 +210,10 @@ def test_validate_generation_valid( def test_validate_generation_invalid( self, - user_manager, - dynamodb_stubber, - storage_table_name, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """Test generation validation when generation doesn't match""" user_id = "user123456789" stored_generation = 5 @@ -237,10 +241,10 @@ def test_validate_generation_invalid( def test_validate_generation_user_not_found( self, - user_manager, - dynamodb_stubber, - storage_table_name, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """Test generation validation when user doesn't exist""" user_id = "user999999999" @@ -257,9 +261,9 @@ def test_validate_generation_user_not_found( def test_validate_generation_dynamodb_unavailable( self, - user_manager, - dynamodb_stubber, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + ) -> None: """Test ServiceUnavailableError when DynamoDB is unavailable during validation""" dynamodb_stubber.add_client_error( "get_item", @@ -274,12 +278,11 @@ def test_validate_generation_dynamodb_unavailable( def test_get_user_success( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_timestamp_datetime, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + ) -> None: """Test successful user retrieval via get_user""" user_id = "user123456789" @@ -307,15 +310,15 @@ def test_get_user_success( assert user.user_id == user_id assert user.generation == 3 assert user.client_state == "deadbeef" - assert user.created_at == mock_timestamp_datetime - assert user.updated_at == mock_timestamp_datetime + assert user.created_at == mock_timestamp + assert user.updated_at == mock_timestamp def test_get_user_not_found( self, - user_manager, - dynamodb_stubber, - storage_table_name, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """Test get_user returns None when user doesn't exist""" user_id = "user999999999" @@ -332,10 +335,10 @@ def test_get_user_not_found( def test_get_or_create_user_unexpected_error( self, - user_manager, - dynamodb_stubber, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + mock_datetime_now: None, + ) -> None: """Test that unexpected ClientErrors are re-raised in get_or_create_user""" dynamodb_stubber.add_client_error( @@ -351,9 +354,9 @@ def test_get_or_create_user_unexpected_error( def test_get_user_unexpected_error( self, - user_manager, - dynamodb_stubber, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + ) -> None: """Test that unexpected ClientErrors are re-raised in get_user""" dynamodb_stubber.add_client_error( @@ -369,9 +372,9 @@ def test_get_user_unexpected_error( def test_increment_generation_unexpected_error( self, - user_manager, - dynamodb_stubber, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + ) -> None: """Test that unexpected ClientErrors are re-raised in increment_generation""" dynamodb_stubber.add_client_error( @@ -387,13 +390,12 @@ def test_increment_generation_unexpected_error( def test_get_or_create_user_new_user_with_client_state( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - mock_timestamp_datetime, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test creating a new user with client_state""" user_id = "user123456789" client_state = "abc123def456" @@ -423,17 +425,17 @@ def test_get_or_create_user_new_user_with_client_state( assert user.generation == 0 assert user.client_state == client_state assert user.client_state_history == [] - assert user.created_at == mock_timestamp_datetime - assert user.updated_at == mock_timestamp_datetime + assert user.created_at == mock_timestamp + assert user.updated_at == mock_timestamp def test_get_or_create_user_existing_user_same_client_state( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test existing user with same client_state does not increment generation""" user_id = "user123456789" client_state = "abc123" @@ -474,12 +476,12 @@ def test_get_or_create_user_existing_user_same_client_state( def test_get_or_create_user_existing_user_different_client_state( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test existing user with different client_state increments generation and updates history""" user_id = "user123456789" old_client_state = "old_state" @@ -557,12 +559,12 @@ def test_get_or_create_user_existing_user_different_client_state( def test_get_or_create_user_client_state_change_dynamodb_unavailable( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test ServiceUnavailableError when DynamoDB fails during client_state update""" user_id = "user123456789" old_client_state = "old_state" @@ -610,10 +612,10 @@ def test_get_or_create_user_client_state_change_dynamodb_unavailable( def test_get_user_missing_client_state_defaults_to_empty( self, - user_manager, - dynamodb_stubber, - storage_table_name, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + ) -> None: """Test get_user returns empty string for missing client_state (legacy records)""" user_id = "user123456789" @@ -644,11 +646,11 @@ def test_get_user_missing_client_state_defaults_to_empty( def test_update_user_client_state_unexpected_error( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + ) -> None: """Test unexpected ClientError is re-raised in update_user_client_state""" user_id = "user123456789" @@ -667,12 +669,12 @@ def test_update_user_client_state_unexpected_error( def test_get_or_create_user_exists_but_cannot_retrieve( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test ServiceUnavailableError when user exists but cannot be retrieved""" user_id = "user123456789" @@ -700,9 +702,9 @@ def test_get_or_create_user_exists_but_cannot_retrieve( def test_validate_client_state_rejects_previously_seen_state( self, - user_manager, - mock_timestamp_datetime, - ): + user_manager: UserManager, + mock_timestamp: float, + ) -> None: """Test rejection of previously-seen client state""" user_id = "user123456789" client_state = "previously_seen_state" @@ -712,8 +714,8 @@ def test_validate_client_state_rejects_previously_seen_state( user_id=user_id, generation=5, client_state="current_state", - created_at=mock_timestamp_datetime, - updated_at=mock_timestamp_datetime, + created_at=mock_timestamp, + updated_at=mock_timestamp, client_state_history=["old_state_1", client_state, "old_state_2"], ) @@ -725,9 +727,9 @@ def test_validate_client_state_rejects_previously_seen_state( def test_validate_client_state_rejects_empty_with_history( self, - user_manager, - mock_timestamp_datetime, - ): + user_manager: UserManager, + mock_timestamp: float, + ) -> None: """Test rejection of empty state when history contains non-empty values""" user_id = "user123456789" @@ -736,8 +738,8 @@ def test_validate_client_state_rejects_empty_with_history( user_id=user_id, generation=5, client_state="current_state", - created_at=mock_timestamp_datetime, - updated_at=mock_timestamp_datetime, + created_at=mock_timestamp, + updated_at=mock_timestamp, client_state_history=["old_state_1", "old_state_2"], ) @@ -749,9 +751,9 @@ def test_validate_client_state_rejects_empty_with_history( def test_validate_client_state_allows_new_state( self, - user_manager, - mock_timestamp_datetime, - ): + user_manager: UserManager, + mock_timestamp: float, + ) -> None: """Test that new client state not in history is allowed""" user_id = "user123456789" new_client_state = "brand_new_state" @@ -761,8 +763,8 @@ def test_validate_client_state_allows_new_state( user_id=user_id, generation=5, client_state="current_state", - created_at=mock_timestamp_datetime, - updated_at=mock_timestamp_datetime, + created_at=mock_timestamp, + updated_at=mock_timestamp, client_state_history=["old_state_1", "old_state_2"], ) @@ -771,9 +773,9 @@ def test_validate_client_state_allows_new_state( def test_validate_client_state_allows_empty_with_empty_history( self, - user_manager, - mock_timestamp_datetime, - ): + user_manager: UserManager, + mock_timestamp: float, + ) -> None: """Test that empty state is allowed when history is empty""" user_id = "user123456789" @@ -782,8 +784,8 @@ def test_validate_client_state_allows_empty_with_empty_history( user_id=user_id, generation=0, client_state="", - created_at=mock_timestamp_datetime, - updated_at=mock_timestamp_datetime, + created_at=mock_timestamp, + updated_at=mock_timestamp, client_state_history=[], ) @@ -792,12 +794,12 @@ def test_validate_client_state_allows_empty_with_empty_history( def test_get_or_create_user_rejects_previously_seen_client_state( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test that get_or_create_user rejects previously-seen client state""" user_id = "user123456789" old_client_state = "old_state" @@ -839,12 +841,12 @@ def test_get_or_create_user_rejects_previously_seen_client_state( def test_get_or_create_user_rejects_empty_state_with_history( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test that get_or_create_user rejects empty state when history exists""" user_id = "user123456789" current_state = "current_state" @@ -885,12 +887,12 @@ def test_get_or_create_user_rejects_empty_state_with_history( def test_update_user_client_state_caps_history_at_50_entries( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test that client_state_history is capped at 50 entries when updated""" user_id = "user123456789" old_client_state = "state_50" @@ -951,12 +953,12 @@ def test_update_user_client_state_caps_history_at_50_entries( def test_get_or_create_user_updates_history_on_state_change( self, - user_manager, - dynamodb_stubber, - storage_table_name, - mock_timestamp, - mock_datetime_now, - ): + user_manager: UserManager, + dynamodb_stubber: Stubber, + storage_table_name: str, + mock_timestamp: float, + mock_datetime_now: None, + ) -> None: """Test that history is updated when client state changes""" user_id = "user123456789" old_client_state = "old_state" diff --git a/lambda/tests/shared/test_base_route.py b/lambda/tests/shared/test_base_route.py index af5ff6cd..4ff61f49 100644 --- a/lambda/tests/shared/test_base_route.py +++ b/lambda/tests/shared/test_base_route.py @@ -8,108 +8,9 @@ class TestBaseRoute: """Tests for BaseRoute abstract class""" - def test_cannot_instantiate_directly(self): + def test_cannot_instantiate_directly(self) -> None: """Test that BaseRoute cannot be instantiated directly""" with pytest.raises(TypeError) as exc_info: BaseRoute() # type: ignore[abstract] assert "abstract" in str(exc_info.value).lower() - - def test_must_implement_bind(self): - """Test that subclasses must implement bind method""" - - class IncompleteRoute(BaseRoute): - def handle(self, event): - pass - - with pytest.raises(TypeError) as exc_info: - IncompleteRoute() # type: ignore[abstract] - - assert "abstract" in str(exc_info.value).lower() - - def test_must_implement_handle(self): - """Test that subclasses must implement handle method""" - - class IncompleteRoute(BaseRoute): - def bind(self, api): - pass - - with pytest.raises(TypeError) as exc_info: - IncompleteRoute() # type: ignore[abstract] - - assert "abstract" in str(exc_info.value).lower() - - def test_can_instantiate_complete_subclass(self): - """Test that subclass with both methods can be instantiated""" - - class CompleteRoute(BaseRoute): - def bind(self, api): - return "bound" - - def handle(self, event): - return "handled" - - route = CompleteRoute() - - assert route.bind("api") == "bound" - assert route.handle("event") == "handled" - - def test_subclass_can_have_additional_methods(self): - """Test that subclass can have additional methods""" - - class ExtendedRoute(BaseRoute): - def bind(self, api): - pass - - def handle(self, event): - return self.process(event) - - def process(self, event): - return f"processed: {event}" - - route = ExtendedRoute() - - assert route.handle("test") == "processed: test" - - def test_subclass_can_have_constructor(self): - """Test that subclass can have its own constructor""" - - class RouteWithConstructor(BaseRoute): - def __init__(self, storage_manager): - self.storage_manager = storage_manager - - def bind(self, api): - pass - - def handle(self, event): - return self.storage_manager - - mock_storage = "mock_storage" - route = RouteWithConstructor(mock_storage) - - assert route.handle(None) == mock_storage - - def test_multiple_subclasses_independent(self): - """Test that multiple subclasses are independent""" - - class Route1(BaseRoute): - def bind(self, api): - return "route1_bind" - - def handle(self, event): - return "route1_handle" - - class Route2(BaseRoute): - def bind(self, api): - return "route2_bind" - - def handle(self, event): - return "route2_handle" - - route1 = Route1() - route2 = Route2() - - assert route1.bind(None) == "route1_bind" - assert route2.bind(None) == "route2_bind" - assert route1.handle(None) == "route1_handle" - assert route2.handle(None) == "route2_handle" diff --git a/lambda/tests/shared/test_exceptions.py b/lambda/tests/shared/test_exceptions.py index beed5fbb..2a6f6b6e 100644 --- a/lambda/tests/shared/test_exceptions.py +++ b/lambda/tests/shared/test_exceptions.py @@ -1,6 +1,5 @@ """Tests for exception classes""" -import json from http import HTTPStatus import pytest @@ -35,12 +34,13 @@ UnsupportedMediaTypeException, ValidationException, ) +from tests.conftest import json_body class TestSyncStorageException: """Tests for SyncStorageException base class""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = SyncStorageException() @@ -49,14 +49,14 @@ def test_default_initialization(self): assert exc.error_code == "InternalServerError" assert str(exc) == "Internal server error" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = SyncStorageException("Custom error message") assert exc.message == "Custom error message" assert str(exc) == "Custom error message" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting exception to Response""" exc = SyncStorageException("Test error") response = exc.to_response() @@ -64,8 +64,7 @@ def test_to_response(self): assert response.status_code == HTTPStatus.INTERNAL_SERVER_ERROR assert response.content_type == "application/json" - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "InternalServerError" assert body["message"] == "Test error" @@ -73,7 +72,7 @@ def test_to_response(self): class TestValidationException: """Tests for ValidationException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = ValidationException() @@ -81,21 +80,20 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.BAD_REQUEST assert exc.error_code == "ValidationException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = ValidationException("Invalid collection name") assert exc.message == "Invalid collection name" assert str(exc) == "Invalid collection name" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to Response""" exc = ValidationException("Invalid input") response = exc.to_response() assert response.status_code == HTTPStatus.BAD_REQUEST - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "ValidationException" assert body["message"] == "Invalid input" @@ -103,7 +101,7 @@ def test_to_response(self): class TestConflictException: """Tests for ConflictException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = ConflictException() @@ -111,27 +109,26 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.CONFLICT assert exc.error_code == "ConflictException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = ConflictException("Collection already exists") assert exc.message == "Collection already exists" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to Response""" exc = ConflictException("Conflict detected") response = exc.to_response() assert response.status_code == HTTPStatus.CONFLICT - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "ConflictException" class TestPreconditionFailedException: """Tests for PreconditionFailedException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = PreconditionFailedException() @@ -139,27 +136,26 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.PRECONDITION_FAILED assert exc.error_code == "PreconditionFailedException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = PreconditionFailedException("Modified since check failed") assert exc.message == "Modified since check failed" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to Response""" exc = PreconditionFailedException("Precondition not met") response = exc.to_response() assert response.status_code == HTTPStatus.PRECONDITION_FAILED - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "PreconditionFailedException" class TestQuotaExceededException: """Tests for QuotaExceededException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = QuotaExceededException() @@ -167,13 +163,13 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.INSUFFICIENT_STORAGE assert exc.error_code == "QuotaExceededException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = QuotaExceededException("Maximum storage limit reached") assert exc.message == "Maximum storage limit reached" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to Response returns Mozilla code (Requirement 13.1, 13.5)""" exc = QuotaExceededException("Quota exceeded") response = exc.to_response() @@ -186,7 +182,7 @@ def test_to_response(self): class TestCollectionNotFoundException: """Tests for CollectionNotFoundException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = CollectionNotFoundException() @@ -194,27 +190,26 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.NOT_FOUND assert exc.error_code == "CollectionNotFoundException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = CollectionNotFoundException("Collection 'bookmarks' not found") assert exc.message == "Collection 'bookmarks' not found" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to Response""" exc = CollectionNotFoundException("Not found") response = exc.to_response() assert response.status_code == HTTPStatus.NOT_FOUND - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "CollectionNotFoundException" class TestStorageObjectNotFoundException: """Tests for StorageObjectNotFoundException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = StorageObjectNotFoundException() @@ -222,27 +217,26 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.NOT_FOUND assert exc.error_code == "StorageObjectNotFoundException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = StorageObjectNotFoundException("Object 'item123' not found") assert exc.message == "Object 'item123' not found" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to Response""" exc = StorageObjectNotFoundException("Object missing") response = exc.to_response() assert response.status_code == HTTPStatus.NOT_FOUND - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "StorageObjectNotFoundException" class TestAuthenticationException: """Tests for AuthenticationException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = AuthenticationException() @@ -250,27 +244,26 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.UNAUTHORIZED assert exc.error_code == "AuthenticationException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = AuthenticationException("Invalid credentials") assert exc.message == "Invalid credentials" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to Response""" exc = AuthenticationException("Auth failed") response = exc.to_response() assert response.status_code == HTTPStatus.UNAUTHORIZED - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "AuthenticationException" class TestExceptionInheritance: """Test exception inheritance""" - def test_all_exceptions_inherit_from_base(self): + def test_all_exceptions_inherit_from_base(self) -> None: """Test that all custom exceptions inherit from SyncStorageException""" exceptions = [ ValidationException(), @@ -286,7 +279,7 @@ def test_all_exceptions_inherit_from_base(self): assert isinstance(exc, SyncStorageException) assert isinstance(exc, Exception) - def test_exceptions_are_raisable(self): + def test_exceptions_are_raisable(self) -> None: """Test that exceptions can be raised and caught""" with pytest.raises(ValidationException) as exc_info: raise ValidationException("Test error") @@ -303,7 +296,7 @@ def test_exceptions_are_raisable(self): class TestInvalidTokenError: """Tests for InvalidTokenError exception""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = InvalidTokenError() @@ -313,14 +306,14 @@ def test_default_initialization(self): assert exc.error_code == "InvalidTokenError" assert str(exc) == "Invalid or expired token" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = InvalidTokenError("Token signature verification failed") assert exc.message == "Token signature verification failed" assert str(exc) == "Token signature verification failed" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting exception to Response""" exc = InvalidTokenError("Token expired") response = exc.to_response() @@ -328,8 +321,7 @@ def test_to_response(self): assert response.status_code == HTTPStatus.UNAUTHORIZED assert response.content_type == "application/json" - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "InvalidTokenError" assert body["message"] == "Token expired" @@ -337,7 +329,7 @@ def test_to_response(self): class TestInvalidCredentialsError: """Tests for InvalidCredentialsError exception""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = InvalidCredentialsError() @@ -346,27 +338,26 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.UNAUTHORIZED assert exc.error_code == "InvalidCredentialsError" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = InvalidCredentialsError("Authentication failed") assert exc.message == "Authentication failed" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting exception to Response""" exc = InvalidCredentialsError("Bad credentials") response = exc.to_response() assert response.status_code == HTTPStatus.UNAUTHORIZED - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "InvalidCredentialsError" class TestTokenValidationError: """Tests for TokenValidationError exception""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = TokenValidationError() @@ -375,34 +366,33 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.BAD_REQUEST assert exc.error_code == "ValidationException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = TokenValidationError("Invalid token format") assert exc.message == "Invalid token format" - def test_inherits_from_validation_exception(self): + def test_inherits_from_validation_exception(self) -> None: """Test that TokenValidationError inherits from ValidationException""" exc = TokenValidationError() assert isinstance(exc, ValidationException) assert isinstance(exc, SyncStorageException) - def test_to_response(self): + def test_to_response(self) -> None: """Test converting exception to Response""" exc = TokenValidationError("Malformed token") response = exc.to_response() assert response.status_code == HTTPStatus.BAD_REQUEST - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "ValidationException" class TestServiceUnavailableError: """Tests for ServiceUnavailableError exception""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = ServiceUnavailableError() @@ -411,13 +401,13 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.SERVICE_UNAVAILABLE assert exc.error_code == "ServiceUnavailableError" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = ServiceUnavailableError("OIDC provider unreachable") assert exc.message == "OIDC provider unreachable" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting exception to Response""" exc = ServiceUnavailableError("Database connection failed") response = exc.to_response() @@ -425,8 +415,7 @@ def test_to_response(self): assert response.status_code == HTTPStatus.SERVICE_UNAVAILABLE assert response.content_type == "application/json" - assert response.body is not None - body = json.loads(response.body) + body = json_body(response) assert body["error"] == "ServiceUnavailableError" assert body["message"] == "Database connection failed" @@ -434,7 +423,7 @@ def test_to_response(self): class TestRequestTooLargeException: """Tests for RequestTooLargeException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = RequestTooLargeException() @@ -442,13 +431,13 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.REQUEST_ENTITY_TOO_LARGE assert exc.error_code == "RequestTooLargeException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = RequestTooLargeException("Payload exceeds 2MB limit") assert exc.message == "Payload exceeds 2MB limit" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting exception to Response""" exc = RequestTooLargeException("Request too large") response = exc.to_response() @@ -460,7 +449,7 @@ def test_to_response(self): class TestMethodNotAllowedException: """Tests for MethodNotAllowedException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = MethodNotAllowedException() @@ -468,13 +457,13 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.METHOD_NOT_ALLOWED assert exc.error_code == "MethodNotAllowedException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = MethodNotAllowedException("POST not allowed on this resource") assert exc.message == "POST not allowed on this resource" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting exception to Response""" exc = MethodNotAllowedException("Method not allowed") response = exc.to_response() @@ -486,7 +475,7 @@ def test_to_response(self): class TestUnsupportedMediaTypeException: """Tests for UnsupportedMediaTypeException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = UnsupportedMediaTypeException() @@ -494,13 +483,13 @@ def test_default_initialization(self): assert exc.status_code == HTTPStatus.UNSUPPORTED_MEDIA_TYPE assert exc.error_code == "UnsupportedMediaTypeException" - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = UnsupportedMediaTypeException("Content-Type must be application/json") assert exc.message == "Content-Type must be application/json" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting exception to Response""" exc = UnsupportedMediaTypeException("Unsupported media type") response = exc.to_response() @@ -512,7 +501,7 @@ def test_to_response(self): class TestServerLimitExceededException: """Tests for ServerLimitExceededException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = ServerLimitExceededException() @@ -521,13 +510,13 @@ def test_default_initialization(self): assert exc.error_code == "ServerLimitExceededException" assert exc.mozilla_code == 17 - def test_custom_message(self): + def test_custom_message(self) -> None: """Test exception with custom message""" exc = ServerLimitExceededException("Batch size exceeds 100 records") assert exc.message == "Batch size exceeds 100 records" - def test_to_response_returns_mozilla_code(self): + def test_to_response_returns_mozilla_code(self) -> None: """Test converting exception to Response returns Mozilla code (Requirement 13.1, 13.7)""" exc = ServerLimitExceededException("Server limit exceeded") response = exc.to_response() @@ -541,7 +530,7 @@ def test_to_response_returns_mozilla_code(self): class TestQuotaExceededExceptionMozillaCode: """Tests for QuotaExceededException Mozilla response code""" - def test_to_response_returns_mozilla_code(self): + def test_to_response_returns_mozilla_code(self) -> None: """Test converting exception to Response returns Mozilla code (Requirement 13.1, 13.5)""" exc = QuotaExceededException("Quota exceeded") response = exc.to_response() @@ -555,7 +544,7 @@ def test_to_response_returns_mozilla_code(self): class TestInvalidBSOException: """Tests for InvalidBSOException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = InvalidBSOException() @@ -564,7 +553,7 @@ def test_default_initialization(self): assert exc.error_code == "InvalidBSOException" assert exc.mozilla_code == 8 - def test_to_response_returns_mozilla_code(self): + def test_to_response_returns_mozilla_code(self) -> None: """Test converting exception to Response returns Mozilla code (Requirement 13.1, 13.3)""" exc = InvalidBSOException("Invalid BSO payload") response = exc.to_response() @@ -578,7 +567,7 @@ def test_to_response_returns_mozilla_code(self): class TestInvalidCollectionException: """Tests for InvalidCollectionException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = InvalidCollectionException() @@ -587,7 +576,7 @@ def test_default_initialization(self): assert exc.error_code == "InvalidCollectionException" assert exc.mozilla_code == 13 - def test_to_response_returns_mozilla_code(self): + def test_to_response_returns_mozilla_code(self) -> None: """Test converting exception to Response returns Mozilla code (Requirement 13.1, 13.4)""" exc = InvalidCollectionException("Collection name too long") response = exc.to_response() @@ -601,7 +590,7 @@ def test_to_response_returns_mozilla_code(self): class TestJSONParseException: """Tests for JSONParseException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = JSONParseException() @@ -610,7 +599,7 @@ def test_default_initialization(self): assert exc.error_code == "JSONParseException" assert exc.mozilla_code == 6 - def test_to_response_returns_mozilla_code(self): + def test_to_response_returns_mozilla_code(self) -> None: """Test converting exception to Response returns Mozilla code (Requirement 13.1, 13.2)""" exc = JSONParseException("Malformed JSON") response = exc.to_response() @@ -624,7 +613,7 @@ def test_to_response_returns_mozilla_code(self): class TestIncompatibleClientException: """Tests for IncompatibleClientException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test exception with default message""" exc = IncompatibleClientException() @@ -633,7 +622,7 @@ def test_default_initialization(self): assert exc.error_code == "IncompatibleClientException" assert exc.mozilla_code == 16 - def test_to_response_returns_mozilla_code(self): + def test_to_response_returns_mozilla_code(self) -> None: """Test converting exception to Response returns Mozilla code (Requirement 13.1, 13.6)""" exc = IncompatibleClientException("Client version not supported") response = exc.to_response() @@ -647,7 +636,7 @@ def test_to_response_returns_mozilla_code(self): class TestOptionalResponseHeaders: """Tests for optional response headers (Requirements 5.7, 18.1-18.4)""" - def test_retry_after_header(self): + def test_retry_after_header(self) -> None: """Test Retry-After header on ConflictException (Requirement 5.7)""" exc = ConflictException("Resource conflict", retry_after=30) response = exc.to_response() @@ -656,7 +645,7 @@ def test_retry_after_header(self): assert response.headers is not None assert response.headers.get("Retry-After") == "30" - def test_x_weave_backoff_header(self): + def test_x_weave_backoff_header(self) -> None: """Test X-Weave-Backoff header (Requirement 18.1)""" exc = ServiceUnavailableError("Server under load", backoff=60) response = exc.to_response() @@ -665,7 +654,7 @@ def test_x_weave_backoff_header(self): assert response.headers is not None assert response.headers.get("X-Weave-Backoff") == "60" - def test_x_weave_alert_header(self): + def test_x_weave_alert_header(self) -> None: """Test X-Weave-Alert header (Requirement 18.3)""" exc = ServiceUnavailableError("Service decommissioned", alert="hard-eol") response = exc.to_response() @@ -674,7 +663,7 @@ def test_x_weave_alert_header(self): assert response.headers is not None assert response.headers.get("X-Weave-Alert") == "hard-eol" - def test_multiple_optional_headers(self): + def test_multiple_optional_headers(self) -> None: """Test multiple optional headers together""" exc = ConflictException( "Conflict detected", retry_after=15, backoff=30, alert="Please retry" @@ -687,7 +676,7 @@ def test_multiple_optional_headers(self): assert response.headers.get("X-Weave-Backoff") == "30" assert response.headers.get("X-Weave-Alert") == "Please retry" - def test_no_optional_headers_by_default(self): + def test_no_optional_headers_by_default(self) -> None: """Test that optional headers are not present by default""" exc = ValidationException("Invalid input") response = exc.to_response() @@ -703,7 +692,7 @@ def test_no_optional_headers_by_default(self): class TestTokenServerExceptionsWithKwargs: """Test Token Server exceptions accept **kwargs for optional headers""" - def test_invalid_timestamp_error_with_kwargs(self): + def test_invalid_timestamp_error_with_kwargs(self) -> None: """Test InvalidTimestampError accepts optional headers""" exc = InvalidTimestampError("Timestamp mismatch", retry_after=10) response = exc.to_response() @@ -712,7 +701,7 @@ def test_invalid_timestamp_error_with_kwargs(self): assert response.headers is not None assert response.headers.get("Retry-After") == "10" - def test_invalid_generation_error_with_kwargs(self): + def test_invalid_generation_error_with_kwargs(self) -> None: """Test InvalidGenerationError accepts optional headers""" exc = InvalidGenerationError("Generation outdated", alert="Please re-authenticate") response = exc.to_response() @@ -721,7 +710,7 @@ def test_invalid_generation_error_with_kwargs(self): assert response.headers is not None assert response.headers.get("X-Weave-Alert") == "Please re-authenticate" - def test_invalid_client_state_error_with_kwargs(self): + def test_invalid_client_state_error_with_kwargs(self) -> None: """Test InvalidClientStateError accepts optional headers""" exc = InvalidClientStateError("Invalid state", backoff=5) response = exc.to_response() @@ -730,7 +719,7 @@ def test_invalid_client_state_error_with_kwargs(self): assert response.headers is not None assert response.headers.get("X-Weave-Backoff") == "5" - def test_new_users_disabled_error_with_kwargs(self): + def test_new_users_disabled_error_with_kwargs(self) -> None: """Test NewUsersDisabledError accepts optional headers""" exc = NewUsersDisabledError("Registration disabled", alert="Service closed") response = exc.to_response() @@ -746,18 +735,18 @@ def test_new_users_disabled_error_with_kwargs(self): class TestInvalidHawkHeaderException: """Tests for InvalidHawkHeaderException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test default initialization""" exception = InvalidHawkHeaderException() assert exception.message == "Malformed HAWK Authorization header" assert exception.status_code == HTTPStatus.UNAUTHORIZED - def test_custom_message(self): + def test_custom_message(self) -> None: """Test custom message""" exception = InvalidHawkHeaderException("Custom HAWK header error") assert exception.message == "Custom HAWK header error" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to response""" exception = InvalidHawkHeaderException() response = exception.to_response() @@ -770,18 +759,18 @@ def test_to_response(self): class TestInvalidHawkSignatureException: """Tests for InvalidHawkSignatureException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test default initialization""" exception = InvalidHawkSignatureException() assert exception.message == "HAWK signature verification failed" assert exception.status_code == HTTPStatus.UNAUTHORIZED - def test_custom_message(self): + def test_custom_message(self) -> None: """Test custom message""" exception = InvalidHawkSignatureException("Signature mismatch") assert exception.message == "Signature mismatch" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to response""" exception = InvalidHawkSignatureException() response = exception.to_response() @@ -794,18 +783,18 @@ def test_to_response(self): class TestExpiredHawkTokenException: """Tests for ExpiredHawkTokenException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test default initialization""" exception = ExpiredHawkTokenException() assert exception.message == "HAWK token has expired" assert exception.status_code == HTTPStatus.UNAUTHORIZED - def test_custom_message(self): + def test_custom_message(self) -> None: """Test custom message""" exception = ExpiredHawkTokenException("Token expired at 1234567890") assert exception.message == "Token expired at 1234567890" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to response""" exception = ExpiredHawkTokenException() response = exception.to_response() @@ -818,18 +807,18 @@ def test_to_response(self): class TestInvalidGenerationException: """Tests for InvalidGenerationException""" - def test_default_initialization(self): + def test_default_initialization(self) -> None: """Test default initialization""" exception = InvalidGenerationException() assert exception.message == "HAWK token generation number is outdated" assert exception.status_code == HTTPStatus.UNAUTHORIZED - def test_custom_message(self): + def test_custom_message(self) -> None: """Test custom message""" exception = InvalidGenerationException("Generation mismatch: expected 5, got 3") assert exception.message == "Generation mismatch: expected 5, got 3" - def test_to_response(self): + def test_to_response(self) -> None: """Test converting to response""" exception = InvalidGenerationException() response = exception.to_response() @@ -842,14 +831,14 @@ def test_to_response(self): class TestHawkExceptionInheritance: """Tests for HAWK exception inheritance""" - def test_hawk_exceptions_inherit_from_authentication_exception(self): + def test_hawk_exceptions_inherit_from_authentication_exception(self) -> None: """Test that all HAWK exceptions inherit from AuthenticationException""" assert issubclass(InvalidHawkHeaderException, AuthenticationException) assert issubclass(InvalidHawkSignatureException, AuthenticationException) assert issubclass(ExpiredHawkTokenException, AuthenticationException) assert issubclass(InvalidGenerationException, AuthenticationException) - def test_hawk_exceptions_are_raisable(self): + def test_hawk_exceptions_are_raisable(self) -> None: """Test that HAWK exceptions can be raised and caught""" with pytest.raises(InvalidHawkHeaderException): diff --git a/lambda/tests/shared/test_models.py b/lambda/tests/shared/test_models.py index 748c6fac..2965f18e 100644 --- a/lambda/tests/shared/test_models.py +++ b/lambda/tests/shared/test_models.py @@ -21,138 +21,138 @@ class TestValidatePayloadSize: - def test_valid_payload(self): + def test_valid_payload(self) -> None: """Valid payload should not raise exception""" payload = "a" * 1000 validate_payload_size(payload) # Should not raise - def test_payload_at_max_size(self): + def test_payload_at_max_size(self) -> None: """Payload at exactly max size should be valid""" payload = "a" * MAX_PAYLOAD_BYTES validate_payload_size(payload) # Should not raise - def test_payload_exceeds_max_size(self): + def test_payload_exceeds_max_size(self) -> None: """Payload exceeding max size should raise ValidationError""" payload = "a" * (MAX_PAYLOAD_BYTES + 1) with pytest.raises(ValidationError, match="Payload size .* exceeds maximum"): validate_payload_size(payload) - def test_empty_payload(self): + def test_empty_payload(self) -> None: """Empty payload should be valid""" validate_payload_size("") # Should not raise class TestValidateBSOId: - def test_valid_bso_id(self): + def test_valid_bso_id(self) -> None: """Valid BSO ID should not raise exception""" validate_bso_id("valid-bso-id") # Should not raise - def test_bso_id_at_max_length(self): + def test_bso_id_at_max_length(self) -> None: """BSO ID at exactly max length should be valid""" validate_bso_id("a" * MAX_BSO_ID_LENGTH) # Should not raise - def test_bso_id_exceeds_max_length(self): + def test_bso_id_exceeds_max_length(self) -> None: """BSO ID exceeding max length should raise ValidationError""" with pytest.raises(ValidationError, match="BSO ID length .* exceeds maximum"): validate_bso_id("a" * (MAX_BSO_ID_LENGTH + 1)) - def test_bso_id_with_special_chars(self): + def test_bso_id_with_special_chars(self) -> None: """BSO ID with printable ASCII characters should be valid""" validate_bso_id("valid-bso-id_123.test") # Should not raise - def test_bso_id_with_non_printable_chars(self): + def test_bso_id_with_non_printable_chars(self) -> None: """BSO ID with non-printable ASCII should raise ValidationError""" with pytest.raises(ValidationError, match="non-printable ASCII"): validate_bso_id("invalid\x00id") - def test_bso_id_with_tab_char(self): + def test_bso_id_with_tab_char(self) -> None: """BSO ID with tab character should raise ValidationError""" with pytest.raises(ValidationError, match="non-printable ASCII"): validate_bso_id("invalid\tid") - def test_bso_id_with_del_char(self): + def test_bso_id_with_del_char(self) -> None: """BSO ID with DEL character (0x7F) should raise ValidationError""" with pytest.raises(ValidationError, match="non-printable ASCII"): validate_bso_id("invalid\x7fid") - def test_empty_bso_id(self): + def test_empty_bso_id(self) -> None: """Empty BSO ID should be rejected (smithy ObjectId requires min length 1)""" with pytest.raises(ValidationError): validate_bso_id("") class TestValidateCollectionName: - def test_empty_collection_name(self): + def test_empty_collection_name(self) -> None: """Empty collection name should be rejected (smithy CollectionName min length 1)""" with pytest.raises(ValidationError): validate_collection_name("") - def test_valid_collection_name(self): + def test_valid_collection_name(self) -> None: """Valid collection name should not raise exception""" validate_collection_name("bookmarks") # Should not raise - def test_collection_name_with_special_chars(self): + def test_collection_name_with_special_chars(self) -> None: """Collection name with allowed special characters""" validate_collection_name("my-collection_1.0") # Should not raise - def test_collection_name_with_invalid_chars(self): + def test_collection_name_with_invalid_chars(self) -> None: """Collection name with invalid characters should raise ValidationError""" with pytest.raises(ValidationError, match="invalid character"): validate_collection_name("invalid collection!") - def test_collection_name_with_space(self): + def test_collection_name_with_space(self) -> None: """Collection name with space should raise ValidationError""" with pytest.raises(ValidationError, match="invalid character"): validate_collection_name("invalid name") - def test_collection_name_at_max_length(self): + def test_collection_name_at_max_length(self) -> None: """Collection name at exactly max length should be valid""" validate_collection_name("a" * MAX_COLLECTION_NAME_LENGTH) # Should not raise - def test_collection_name_exceeds_max_length(self): + def test_collection_name_exceeds_max_length(self) -> None: """Collection name exceeding max length should raise ValidationError""" with pytest.raises(ValidationError, match="Collection name length .* exceeds maximum"): validate_collection_name("a" * (MAX_COLLECTION_NAME_LENGTH + 1)) class TestBSOInput: - def test_all_fields_optional(self): + def test_all_fields_optional(self) -> None: bso = BSOInput() assert bso.id is None assert bso.payload is None assert bso.sortindex is None assert bso.ttl is None - def test_sortindex_at_bounds(self): + def test_sortindex_at_bounds(self) -> None: BSOInput(sortindex=999999999) BSOInput(sortindex=-999999999) - def test_sortindex_out_of_range(self): + def test_sortindex_out_of_range(self) -> None: with pytest.raises(PydanticValidationError): BSOInput(sortindex=1000000000) with pytest.raises(PydanticValidationError): BSOInput(sortindex=-1000000000) - def test_ttl_must_be_positive(self): + def test_ttl_must_be_positive(self) -> None: with pytest.raises(PydanticValidationError): BSOInput(ttl=0) with pytest.raises(PydanticValidationError): BSOInput(ttl=-1) - def test_ttl_at_max(self): + def test_ttl_at_max(self) -> None: BSOInput(ttl=999999999) - def test_ttl_exceeds_max(self): + def test_ttl_exceeds_max(self) -> None: with pytest.raises(PydanticValidationError): BSOInput(ttl=1000000000) - def test_payload_accepts_large_string(self): + def test_payload_accepts_large_string(self) -> None: """Payload validation is byte-based (validate_payload_size), not char-based.""" BSOInput(payload="a" * 262144) # no Pydantic char limit class TestCamelModelAliasing: - def test_device_output_serializes_to_camel(self): + def test_device_output_serializes_to_camel(self) -> None: dev = DeviceOutput( id="d1", name="My Phone", @@ -169,7 +169,7 @@ def test_device_output_serializes_to_camel(self): # snake_case keys should NOT appear when by_alias=True assert "push_callback" not in d - def test_device_output_accepts_camel_input(self): + def test_device_output_accepts_camel_input(self) -> None: dev = DeviceOutput.model_validate( { "id": "d1", @@ -183,7 +183,7 @@ def test_device_output_accepts_camel_input(self): assert dev.push_callback == "https://push" assert dev.created_at == 100 - def test_device_output_accepts_snake_input(self): + def test_device_output_accepts_snake_input(self) -> None: dev = DeviceOutput( id="d1", name="Phone", @@ -196,7 +196,7 @@ def test_device_output_accepts_snake_input(self): class TestBatchResultOutput: - def test_basic_creation(self): + def test_basic_creation(self) -> None: br = BatchResultOutput( success=["a", "b"], failed={"c": ["error"]}, @@ -210,35 +210,35 @@ def test_basic_creation(self): class TestCollectionDataOutput: - def test_basic_creation(self): + def test_basic_creation(self) -> None: cd = CollectionDataOutput(name="bookmarks", modified=1.0, count=5, usage=1024) assert cd.name == "bookmarks" assert cd.count == 5 class TestModifiedOutput: - def test_basic_creation(self): + def test_basic_creation(self) -> None: m = ModifiedOutput(modified=1.23) assert m.modified == 1.23 class TestAccountCreateInput: - def test_valid(self): + def test_valid(self) -> None: pw = "a" * 64 a = AccountCreateInput(email="user@example.com", auth_pw=pw) assert a.auth_pw == pw - def test_auth_pw_too_short(self): + def test_auth_pw_too_short(self) -> None: with pytest.raises(PydanticValidationError): AccountCreateInput(email="user@example.com", auth_pw="short") - def test_auth_pw_too_long(self): + def test_auth_pw_too_long(self) -> None: with pytest.raises(PydanticValidationError): AccountCreateInput(email="user@example.com", auth_pw="a" * 65) class TestToDynamoDict: - def test_converts_float_to_decimal(self): + def test_converts_float_to_decimal(self) -> None: from src.shared.models import BasicStorageObject, to_dynamo_dict bso = BasicStorageObject(id="x", payload="p", modified=3.14) @@ -248,7 +248,7 @@ def test_converts_float_to_decimal(self): assert dumped["id"] == "x" assert dumped["payload"] == "p" - def test_recurses_into_dict_and_list(self): + def test_recurses_into_dict_and_list(self) -> None: from src.shared.models import _to_dynamo result = _to_dynamo({"a": 1.5, "b": [2.5, 3], "c": {"d": 4.0}}) @@ -260,7 +260,7 @@ def test_recurses_into_dict_and_list(self): class TestDeviceOutputDecimalFields: - def test_decimal_fields_convert_to_int(self): + def test_decimal_fields_convert_to_int(self) -> None: dev = DeviceOutput.model_validate( { "id": "d1", diff --git a/lambda/tests/shared/test_oidc.py b/lambda/tests/shared/test_oidc.py index c1218e9d..4099e5f6 100644 --- a/lambda/tests/shared/test_oidc.py +++ b/lambda/tests/shared/test_oidc.py @@ -8,7 +8,7 @@ class TestOIDCTokenClaims: """Tests for OIDCTokenClaims model""" - def test_creation_with_all_fields(self): + def test_creation_with_all_fields(self) -> None: claims = OIDCTokenClaims( sub="user123", iss="https://auth.example.com", @@ -25,7 +25,7 @@ def test_creation_with_all_fields(self): assert claims.iat == 1234567800 assert claims.email == "user@example.com" - def test_creation_without_email(self): + def test_creation_without_email(self) -> None: claims = OIDCTokenClaims( sub="user456", iss="https://auth.example.com", @@ -37,7 +37,7 @@ def test_creation_without_email(self): assert claims.sub == "user456" assert claims.email is None - def test_exp_greater_than_iat(self): + def test_exp_greater_than_iat(self) -> None: claims = OIDCTokenClaims( sub="user", iss="https://auth.example.com", @@ -48,7 +48,7 @@ def test_exp_greater_than_iat(self): assert claims.exp > claims.iat - def test_asdict(self): + def test_asdict(self) -> None: claims = OIDCTokenClaims( sub="dictuser", iss="https://auth.example.com", @@ -67,7 +67,7 @@ def test_asdict(self): class TestOIDCProviderConfig: """Tests for OIDCProviderConfig model""" - def test_creation_with_all_fields(self): + def test_creation_with_all_fields(self) -> None: config = OIDCProviderConfig( issuer="https://auth.example.com", jwks_uri="https://auth.example.com/jwks", @@ -86,7 +86,7 @@ def test_creation_with_all_fields(self): class TestErrorDetail: """Tests for ErrorDetail model""" - def test_creation_with_all_fields(self): + def test_creation_with_all_fields(self) -> None: error = ErrorDetail( location="header", name="Authorization", @@ -97,16 +97,16 @@ def test_creation_with_all_fields(self): assert error.name == "Authorization" assert error.description == "Missing authorization header" - def test_creation_with_body_location(self): + def test_creation_with_body_location(self) -> None: error = ErrorDetail(location="body", name="email", description="Invalid email format") assert error.location == "body" assert error.name == "email" - def test_creation_with_query_location(self): + def test_creation_with_query_location(self) -> None: error = ErrorDetail(location="query", name="limit", description="Limit must be positive") assert error.location == "query" - def test_asdict(self): + def test_asdict(self) -> None: error = ErrorDetail(location="header", name="Accept", description="Unsupported media type") data = error.model_dump() diff --git a/lambda/tests/shared/test_token.py b/lambda/tests/shared/test_token.py index 17c97eb8..c3efd44f 100644 --- a/lambda/tests/shared/test_token.py +++ b/lambda/tests/shared/test_token.py @@ -8,7 +8,7 @@ class TestTokenResponse: """Tests for TokenResponse model""" - def test_creation_with_all_fields(self): + def test_creation_with_all_fields(self) -> None: token = TokenResponse( id="hawk_id_base64", key="hawk_key_hex_64_chars", @@ -25,7 +25,7 @@ def test_creation_with_all_fields(self): assert token.duration == 300 assert token.hashalg == "sha256" - def test_duration_is_300_seconds(self): + def test_duration_is_300_seconds(self) -> None: token = TokenResponse( id="test_id", key="test_key", @@ -36,7 +36,7 @@ def test_duration_is_300_seconds(self): ) assert token.duration == 300 - def test_hashalg_is_sha256(self): + def test_hashalg_is_sha256(self) -> None: token = TokenResponse( id="test_id", key="test_key", @@ -47,7 +47,7 @@ def test_hashalg_is_sha256(self): ) assert token.hashalg == "sha256" - def test_api_endpoint_format(self): + def test_api_endpoint_format(self) -> None: token = TokenResponse( id="test_id", key="test_key", @@ -60,7 +60,7 @@ def test_api_endpoint_format(self): assert "/1.5/" in token.api_endpoint assert token.api_endpoint.endswith("user456") - def test_asdict(self): + def test_asdict(self) -> None: token = TokenResponse( id="dict_id", key="dict_key", @@ -80,7 +80,7 @@ def test_asdict(self): assert data["duration"] == 300 assert data["hashalg"] == "sha256" - def test_uid_is_numeric(self): + def test_uid_is_numeric(self) -> None: token = TokenResponse( id="test_id", key="test_key", @@ -92,7 +92,7 @@ def test_uid_is_numeric(self): assert isinstance(token.uid, int) assert token.uid > 0 - def test_different_uids_for_different_users(self): + def test_different_uids_for_different_users(self) -> None: token1 = TokenResponse( id="id1", key="key1", diff --git a/lambda/tests/shared/test_utils.py b/lambda/tests/shared/test_utils.py index 5f3e2b1f..b383afde 100644 --- a/lambda/tests/shared/test_utils.py +++ b/lambda/tests/shared/test_utils.py @@ -13,7 +13,7 @@ class TestWeaveTimestamp: """Test Weave timestamp generation (Requirements 9.1, 9.2)""" - def test_get_weave_timestamp_format(self): + def test_get_weave_timestamp_format(self) -> None: """Test that get_weave_timestamp returns correct format""" timestamp = get_weave_timestamp() @@ -27,7 +27,7 @@ def test_get_weave_timestamp_format(self): float_value = float(timestamp) assert float_value > 0 - def test_get_weave_timestamp_precision(self): + def test_get_weave_timestamp_precision(self) -> None: """Test that get_weave_timestamp has exactly 2 decimal places""" timestamp = get_weave_timestamp() @@ -40,7 +40,7 @@ def test_get_weave_timestamp_precision(self): class TestExtractHawkRequestParams: """Test extract_hawk_request_params helper""" - def test_extracts_domain_name_from_request_context(self): + def test_extracts_domain_name_from_request_context(self) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -55,7 +55,7 @@ def test_extracts_domain_name_from_request_context(self): assert host == "auth.prod.ffsync.layertwo.dev" assert port == 443 - def test_appends_query_string_to_path(self): + def test_appends_query_string_to_path(self) -> None: event = APIGatewayProxyEvent( { "httpMethod": "POST", @@ -68,7 +68,7 @@ def test_appends_query_string_to_path(self): method, path, host, port = extract_hawk_request_params(event) assert path == "/v1/session/destroy?service=sync" - def test_falls_back_to_host_header_when_no_request_context(self): + def test_falls_back_to_host_header_when_no_request_context(self) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -79,7 +79,7 @@ def test_falls_back_to_host_header_when_no_request_context(self): method, path, host, port = extract_hawk_request_params(event) assert host == "fallback.example.com" - def test_falls_back_to_localhost_when_no_host_or_context(self): + def test_falls_back_to_localhost_when_no_host_or_context(self) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", @@ -90,7 +90,7 @@ def test_falls_back_to_localhost_when_no_host_or_context(self): method, path, host, port = extract_hawk_request_params(event) assert host == "localhost" - def test_no_query_string_when_none(self): + def test_no_query_string_when_none(self) -> None: event = APIGatewayProxyEvent( { "httpMethod": "GET", diff --git a/lambda/uv.lock b/lambda/uv.lock index 075bde02..a03a0f01 100644 --- a/lambda/uv.lock +++ b/lambda/uv.lock @@ -90,15 +90,15 @@ wheels = [ [[package]] name = "aws-lambda-powertools" -version = "3.34.0" +version = "3.35.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "jmespath" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/6b/d2/ead05769a42baf2c01e9e8909b5b7a0a42b0798ab2d0eccacdef3d1f009d/aws_lambda_powertools-3.34.0.tar.gz", hash = "sha256:75f2c65a5997630666c9c3c495002c1f9d1e9defd3cb6d5d9fc7272b7dc9ef54", size = 800151, upload-time = "2026-08-10T11:26:10.742Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a9/5b/9b36aa010e686bbe0a8bfb7367430994254a68d1e8c2383af7b3264c6ebb/aws_lambda_powertools-3.35.0.tar.gz", hash = "sha256:27c1c7d7920f214a68a11505dca33ce3709a38c4b79f5fbc94141ec0af92c9f8", size = 803046, upload-time = "2026-09-15T12:18:56.22Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b3/d1/9e9433f99d8c9084d3bc831900f1ce7a99b0a2c153d9223e46839d8cc665/aws_lambda_powertools-3.34.0-py3-none-any.whl", hash = "sha256:ab1354c58085ccecf92e9c548b1336ad5bae72c7ae23e61d2aa0ff15520d679c", size = 957321, upload-time = "2026-08-10T11:26:09.002Z" }, + { url = "https://files.pythonhosted.org/packages/c7/6e/2a2764dd3fef758e8bde3ec5c0fb6e2ceeace2a873ca5def30b6c69a7bcf/aws_lambda_powertools-3.35.0-py3-none-any.whl", hash = "sha256:e8fbc27d74272fb99cbec67305250e825d1630dceb2dbfdcf5908dc9c4c82dcf", size = 960326, upload-time = "2026-09-15T12:18:54.031Z" }, ] [[package]] @@ -125,30 +125,30 @@ wheels = [ [[package]] name = "boto3" -version = "1.43.89" +version = "1.43.96" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "botocore" }, { name = "jmespath" }, { name = "s3transfer" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/50/26/48b3da85526a72a02df55e564481fc348e93699c15f0f502681b12ac2c8a/boto3-1.43.89.tar.gz", hash = "sha256:c28abbe472e9b7cad08807356311aeec51bde5218c18489da827045d2267bfd9", size = 112702, upload-time = "2026-09-04T19:24:57.143Z" } +sdist = { url = "https://files.pythonhosted.org/packages/16/b6/41173fa75983750c794e9b64017a3203407725a0e8c9c7f6de39686dc97b/boto3-1.43.96.tar.gz", hash = "sha256:30fb2b5467ef5175ed48f43c06c435eec5da841594a5d7653c4da679df0740fc", size = 112691, upload-time = "2026-09-16T19:57:09.113Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/cd/12/e1b5cb4a00a9bfd72cf2d3f982c5826757aacdfc90aa4bd61902dcc94856/boto3-1.43.89-py3-none-any.whl", hash = "sha256:fe4190afe63eb562b6ba6a3911cf4427473b35fa047adde093bf696d3ae09fc0", size = 140028, upload-time = "2026-09-04T19:24:55.929Z" }, + { url = "https://files.pythonhosted.org/packages/01/21/7629ebb023a4857102e57231be4332bf1c8b5c1dcb9103223184b0c91b77/boto3-1.43.96-py3-none-any.whl", hash = "sha256:73f0386fcc412ce5ad3ae4245e2a405e69e889988dd77e91b82e3b66404f3fee", size = 140031, upload-time = "2026-09-16T19:57:07.825Z" }, ] [[package]] name = "botocore" -version = "1.43.89" +version = "1.43.97" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "jmespath" }, { name = "python-dateutil" }, { name = "urllib3" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/53/06/f63fb1befdf77af18539fb24ea01f2da0f13965ed5de091061708ac96416/botocore-1.43.89.tar.gz", hash = "sha256:f0574942970742657b0e0716cf08c2dfe6bef8e6de5fbb7081c3424e262b4cca", size = 16074206, upload-time = "2026-09-04T19:24:52.464Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f9/c7/aff84828bc3cd05328c65650320be27bb13e63f37917f1aba3f6309498c8/botocore-1.43.97.tar.gz", hash = "sha256:7c0e18686367e98de49826d8ac9ae951ce0915c172969a5d96b768c9afa9e979", size = 16120115, upload-time = "2026-09-17T19:28:48.141Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/9e/9d/96f9dee6d12eedf1c2b4264eefd59c7ac8cac10daadb9a7bfccc9ee881c6/botocore-1.43.89-py3-none-any.whl", hash = "sha256:d7211220c815427fe71225acc6909e4ab5dfab3b03770e72fd16cf9eb86b3d1a", size = 15768272, upload-time = "2026-09-04T19:24:49.769Z" }, + { url = "https://files.pythonhosted.org/packages/a9/3b/ffb7780e5f4c6af60eee31b8badd2762bba1793ca05f926e9381919c7d85/botocore-1.43.97-py3-none-any.whl", hash = "sha256:6b2beefd4f515aa7bfc614c945493a19315577ff31b464abbc9e33b7aff5112c", size = 15813332, upload-time = "2026-09-17T19:28:45.248Z" }, ] [[package]] @@ -461,7 +461,7 @@ wheels = [ [[package]] name = "datamodel-code-generator" -version = "0.76.2" +version = "0.82.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "argcomplete" }, @@ -472,10 +472,11 @@ dependencies = [ { name = "jinja2" }, { name = "pydantic" }, { name = "pyyaml" }, + { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/c7/cd/97b5ce0bc7452199a06204e1252fe10ad63cc37c76c2e9f76146b766b708/datamodel_code_generator-0.76.2.tar.gz", hash = "sha256:c8c25e24b5b90c1c45fc40309fedd0a2f571425c7b6efae97efb317c03325c32", size = 2260230, upload-time = "2026-09-04T11:37:47.931Z" } +sdist = { url = "https://files.pythonhosted.org/packages/4a/21/82e94fe764c27e3764870eed81839c2c3c2779936637e4cbecef943aa2fd/datamodel_code_generator-0.82.0.tar.gz", hash = "sha256:ceb08a32a4358c74f92c47220e149041d03dc6470443ba9d0f3d9d406f38c5a3", size = 2941392, upload-time = "2026-09-16T08:57:43.626Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b3/7a/00b736585a6fac1e9c94275a07d8e470a3accdaab37b54d13fdc1c985bc0/datamodel_code_generator-0.76.2-py3-none-any.whl", hash = "sha256:8cc2bffa5a7e81a4b5a1cfce28530eaf9be465f3591bb7cb53030eb0c956f6d9", size = 661183, upload-time = "2026-09-04T11:37:45.73Z" }, + { url = "https://files.pythonhosted.org/packages/cd/3d/6c824b706e8b9be8dfd3e6f5024439b0880efeb71ae826326a60077624f3/datamodel_code_generator-0.82.0-py3-none-any.whl", hash = "sha256:e7a73068aa2f16d34907a09607fb7eb241f29ee1dcd9b2af2fdeb586ba6a8091", size = 724246, upload-time = "2026-09-16T08:57:41.728Z" }, ] [[package]] @@ -501,7 +502,7 @@ dependencies = [ { name = "requests" }, ] -[package.optional-dependencies] +[package.dev-dependencies] dev = [ { name = "black" }, { name = "datamodel-code-generator" }, @@ -513,33 +514,36 @@ dev = [ { name = "pytest" }, { name = "pytest-cov" }, { name = "pytest-xdist" }, - { name = "types-boto3" }, + { name = "types-boto3", extra = ["apigatewaymanagementapi", "dynamodb", "kms"] }, { name = "types-requests" }, ] [package.metadata] requires-dist = [ - { name = "aws-lambda-powertools", specifier = "==3.34.0" }, - { name = "black", marker = "extra == 'dev'", specifier = "==26.5.1" }, - { name = "boto3", specifier = "==1.43.89" }, + { name = "aws-lambda-powertools", specifier = "==3.35.0" }, + { name = "boto3", specifier = "==1.43.96" }, { name = "cryptography", specifier = "==50.0.1" }, - { name = "datamodel-code-generator", marker = "extra == 'dev'", specifier = "==0.76.2" }, - { name = "flake8", marker = "extra == 'dev'", specifier = "==7.3.0" }, - { name = "flake8-pyproject", marker = "extra == 'dev'", specifier = "==1.2.4" }, - { name = "hypothesis", marker = "extra == 'dev'", specifier = ">=6.0.0" }, - { name = "isort", marker = "extra == 'dev'", specifier = "==8.0.1" }, { name = "mohawk", specifier = "==1.1.0" }, - { name = "mypy", marker = "extra == 'dev'", specifier = "==2.3.1" }, { name = "pydantic", specifier = "==2.13.5" }, - { name = "pyjwt", specifier = "==2.13.0" }, - { name = "pytest", marker = "extra == 'dev'", specifier = "==9.1.1" }, - { name = "pytest-cov", marker = "extra == 'dev'", specifier = "==7.1.0" }, - { name = "pytest-xdist", marker = "extra == 'dev'", specifier = "==3.8.0" }, + { name = "pyjwt", specifier = "==2.14.0" }, { name = "requests", specifier = "==2.34.2" }, - { name = "types-boto3", marker = "extra == 'dev'", specifier = "==1.43.89" }, - { name = "types-requests", marker = "extra == 'dev'", specifier = "==2.33.0.20260712" }, ] -provides-extras = ["dev"] + +[package.metadata.requires-dev] +dev = [ + { name = "black", specifier = "==26.5.1" }, + { name = "datamodel-code-generator", specifier = "==0.82.0" }, + { name = "flake8", specifier = "==7.3.0" }, + { name = "flake8-pyproject", specifier = "==1.2.4" }, + { name = "hypothesis", specifier = ">=6.0.0" }, + { name = "isort", specifier = "==8.0.1" }, + { name = "mypy", specifier = "==2.3.1" }, + { name = "pytest", specifier = "==9.1.1" }, + { name = "pytest-cov", specifier = "==7.1.0" }, + { name = "pytest-xdist", specifier = "==3.8.0" }, + { name = "types-boto3", extras = ["apigatewaymanagementapi", "dynamodb", "kms"], specifier = "==1.43.96" }, + { name = "types-requests", specifier = "==2.33.0.20260906" }, +] [[package]] name = "flake8" @@ -1005,11 +1009,11 @@ wheels = [ [[package]] name = "pyjwt" -version = "2.13.0" +version = "2.14.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/3b/81/58d0ac84e1ef3a3843791d6954d94c0b33d526c75eeb1efbce9d0a4c4077/pyjwt-2.13.0.tar.gz", hash = "sha256:41571c89ca91598c79e8ef18a2d07367d4810fbbd6f637794879baf1b7703423", size = 107515, upload-time = "2026-05-21T19:54:36.618Z" } +sdist = { url = "https://files.pythonhosted.org/packages/af/c3/8a3b59c25070cc61dc517fbdfa5dc0904670c96f605cc69759dc09166b99/pyjwt-2.14.0.tar.gz", hash = "sha256:77283c83fb56ecf566a886c757a714bc83668e38156de2cce8263302f42e0b86", size = 113177, upload-time = "2026-09-11T13:11:54.638Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a3/5e/ecf12fdb62546d64385c158514e9b2b671f7832108ef2ecd2020ce0af2d1/pyjwt-2.13.0-py3-none-any.whl", hash = "sha256:66adcc2aff09b3f1bbd95fc1e1577df8ac8723c978552fd43304c8a290ac5728", size = 31274, upload-time = "2026-05-21T19:54:35.362Z" }, + { url = "https://files.pythonhosted.org/packages/9c/97/672cb32ce0dfea44b740cb7b4f97038463b9cf7c0ead1aacf595572851d6/pyjwt-2.14.0-py3-none-any.whl", hash = "sha256:ad0cef71c756a56e74863c2919cf0985f72decbcfcb550ee2f422e7c62b5eedc", size = 32896, upload-time = "2026-09-11T13:11:53.409Z" }, ] [[package]] @@ -1171,27 +1175,65 @@ wheels = [ [[package]] name = "types-boto3" -version = "1.43.89" +version = "1.43.96" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "botocore-stubs" }, { name = "types-s3transfer" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/df/8a/1d98ce50cd2f1a87b00005b442a3ed649855b7f32b709c1d83711a293901/types_boto3-1.43.89.tar.gz", hash = "sha256:325b0b546f89caa8ec7e0924bde93a7e2553e9e809c35fd56597dc9a2e49d627", size = 104896, upload-time = "2026-09-04T21:00:32.89Z" } +sdist = { url = "https://files.pythonhosted.org/packages/42/9f/579aedcd84882391564051c31f1a09ed2b127c3ce0fb41b10bfe7be87332/types_boto3-1.43.96.tar.gz", hash = "sha256:690f48c67d8e7b30540949bbff923aa2f85dc9f0e1432aea1421186438680188", size = 104887, upload-time = "2026-09-16T21:40:23.279Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/e6/da/e8e1dea970e869e78eac9c64ad09c73f7fa311ff4e8b839305cd03420f0f/types_boto3-1.43.96-py3-none-any.whl", hash = "sha256:7740b63afc07772cbc24393b8cffc2228c2406e844577735d8f6aa63b0df21d3", size = 71416, upload-time = "2026-09-16T21:40:18.96Z" }, +] + +[package.optional-dependencies] +apigatewaymanagementapi = [ + { name = "types-boto3-apigatewaymanagementapi" }, +] +dynamodb = [ + { name = "types-boto3-dynamodb" }, +] +kms = [ + { name = "types-boto3-kms" }, +] + +[[package]] +name = "types-boto3-apigatewaymanagementapi" +version = "1.43.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/f6/00/f46e0df6fe833fb1bec99f33e41462fe462046e8892fb7a03da3f542a486/types_boto3_apigatewaymanagementapi-1.43.0.tar.gz", hash = "sha256:b40cdf4eaedc72caaf3bf5d10827ab9e4c0e232e8356a8f2b37f848cbd5d5c5e", size = 15249, upload-time = "2026-04-29T22:58:41.771Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/ee/58/21c55451ea2b6a707137d8aa28e3ea6a65ff4f176a7a4e3930ef433a21bb/types_boto3_apigatewaymanagementapi-1.43.0-py3-none-any.whl", hash = "sha256:580a79bd6fd97e0679440270c5bfe6ecd573d341e13939851350153719d9bde4", size = 18431, upload-time = "2026-04-29T22:58:38.546Z" }, +] + +[[package]] +name = "types-boto3-dynamodb" +version = "1.43.64" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/02/65/90fa8cd19a4dccf25dbaa9d5f50181bbb94a3f1f2d1d7808c5f14a86a998/types_boto3_dynamodb-1.43.64.tar.gz", hash = "sha256:e5eb0fa9027bafea70a00ee59a562cf1ad1f39bb7548b0dc44ccc5e0607d5c7a", size = 49846, upload-time = "2026-08-04T21:37:40.036Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7c/cd/b385b389adf2ded6c8d49d195a0047b7af1184e7eb44e7ac4adbcc97507c/types_boto3_dynamodb-1.43.64-py3-none-any.whl", hash = "sha256:43e819c4f74a0aecea0f72a60d2fc2d0052c3ae9f8ebbd7ab62e35d59f7cbda0", size = 60080, upload-time = "2026-08-04T21:37:37.989Z" }, +] + +[[package]] +name = "types-boto3-kms" +version = "1.43.12" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/0d/46/7343b52e16eaa9dec7099cdd6a901317df583473b1658d04cb42885c8d03/types_boto3_kms-1.43.12.tar.gz", hash = "sha256:f9a06ca5a1cbf02f820208f1e84983a750daa1bce305bd11231961a9d770d9cd", size = 30696, upload-time = "2026-05-20T20:01:12.294Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b8/1b/04a78df1311ed51e1a54a7a88e2da9d1ab4114ff472d1c0fdc0be08140b9/types_boto3-1.43.89-py3-none-any.whl", hash = "sha256:0d04c2c246eb52042347ac8352c8b4753986b6ef10746d952c91051725358d73", size = 71419, upload-time = "2026-09-04T21:00:27.465Z" }, + { url = "https://files.pythonhosted.org/packages/37/e1/08af811394ca720a077a4a9fda7cce33c043819e56e0d722d21c977de444/types_boto3_kms-1.43.12-py3-none-any.whl", hash = "sha256:e3c2d0e510593920464aff052382fc31d7159c15cb2c439c5ad8988f6c8417e2", size = 38951, upload-time = "2026-05-20T20:01:08.731Z" }, ] [[package]] name = "types-requests" -version = "2.33.0.20260712" +version = "2.33.0.20260906" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "urllib3" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/db/51/703318f7b7be8bee126ec13bf615050f932d0179b8784420f3a0199cc769/types_requests-2.33.0.20260712.tar.gz", hash = "sha256:2141b67ab534a5c5cd2dac5034f2a35f42e699c5bf185eee608c5246a069d7fb", size = 25084, upload-time = "2026-07-12T05:14:20.455Z" } +sdist = { url = "https://files.pythonhosted.org/packages/c0/18/4c2c0290953f8b3b9612adfcb07b57f144ade3ad32a76764fca42b77c5f3/types_requests-2.33.0.20260906.tar.gz", hash = "sha256:76ab8a0fb736744a0c3deee7aa57b2927e301f078d9e61f5391b3e92002416b9", size = 25263, upload-time = "2026-09-06T06:35:47.707Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/62/e7/010c87f559e216d83f9dc51e939633fd0d0ead3377340181ab0e223cd3b5/types_requests-2.33.0.20260712-py3-none-any.whl", hash = "sha256:de027e28c171d3da529689cbfa023b0b4eab188c8dfa22fd834eebd2cee6e7bb", size = 21392, upload-time = "2026-07-12T05:14:19.616Z" }, + { url = "https://files.pythonhosted.org/packages/60/4c/51ec821d22a45b4162fa3f3e9e94ea4c0f8c49e82363b10955baa5438391/types_requests-2.33.0.20260906-py3-none-any.whl", hash = "sha256:9f53622652bd921ead7a54d665be1b8d518165c8b51cc9d68d944cd6b6bfa8fb", size = 21461, upload-time = "2026-09-06T06:35:45.856Z" }, ] [[package]]