From 3ae97d290399574bab0c5604e1875a5d92894005 Mon Sep 17 00:00:00 2001 From: Pranav Shridhar Date: Thu, 9 Jul 2026 17:03:51 +0530 Subject: [PATCH 1/2] Add HMAC verification for admin proxy requests Co-authored-by: Cursor --- app/core/config.py | 4 + app/core/hmac_auth.py | 161 +++++++++++++++++++++++++++++++++++ app/main.py | 3 + tests/core/__init__.py | 0 tests/core/test_hmac_auth.py | 54 ++++++++++++ 5 files changed, 222 insertions(+) create mode 100644 app/core/hmac_auth.py create mode 100644 tests/core/__init__.py create mode 100644 tests/core/test_hmac_auth.py diff --git a/app/core/config.py b/app/core/config.py index b9b446e..9035b2b 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -15,5 +15,9 @@ class Settings(BaseSettings): # Used as the default value for domain_setting.default_domain CLUSTER_ID: str = "IN" + # HMAC auth for Django → tars admin proxy requests + TARS_HMAC_SECRET: str = "dev-tars-hmac-secret" + TARS_HMAC_TIMESTAMP_TOLERANCE_SECONDS: int = 300 + settings = Settings() diff --git a/app/core/hmac_auth.py b/app/core/hmac_auth.py new file mode 100644 index 0000000..da8db13 --- /dev/null +++ b/app/core/hmac_auth.py @@ -0,0 +1,161 @@ +"""HMAC verification for Django → tars admin proxy requests. + +Matches django-rest-api `TarsClient` signing: + message = f"{timestamp}:{METHOD}:{path}:{body_str}" + signature = HMAC-SHA256(secret, message).hexdigest() + +Only admin paths are verified. Health checks and future SDK hot paths are skipped. +""" + +from __future__ import annotations + +import hashlib +import hmac +import logging +import time + +from starlette.datastructures import Headers +from starlette.responses import JSONResponse +from starlette.types import ASGIApp, Message, Receive, Scope, Send + +from app.core.config import settings + +logger = logging.getLogger(__name__) + +SIGNATURE_HEADER = "x-clickwrap-signature" +TIMESTAMP_HEADER = "x-clickwrap-timestamp" + +# Admin control-plane prefixes that require HMAC (Django proxy). +ADMIN_PATH_PREFIXES = ("/api/v2/",) + +# Explicitly excluded even if under a broader prefix later. +EXCLUDED_PATH_PREFIXES = ( + "/ht", + "/api/v3/public/", + "/docs", + "/openapi.json", + "/redoc", +) + + +def is_admin_path(path: str) -> bool: + if any(path == prefix or path.startswith(prefix) for prefix in EXCLUDED_PATH_PREFIXES): + return False + return any(path.startswith(prefix) for prefix in ADMIN_PATH_PREFIXES) + + +def build_signature_message(timestamp: str, method: str, path: str, body_str: str) -> bytes: + return f"{timestamp}:{method.upper()}:{path}:{body_str}".encode() + + +def compute_signature(secret: str, message: bytes) -> str: + return hmac.new( + secret.encode("utf-8"), + message, + digestmod=hashlib.sha256, + ).hexdigest() + + +def verify_hmac_headers( + *, + method: str, + path: str, + body_str: str, + signature: str | None, + timestamp: str | None, + secret: str, + tolerance_seconds: int, + now: int | None = None, +) -> str | None: + """Return an error detail string if verification fails, else None.""" + if not signature or not timestamp: + return "Missing HMAC signature headers" + + try: + ts = int(timestamp) + except ValueError: + return "Invalid HMAC timestamp" + + current = now if now is not None else int(time.time()) + if abs(current - ts) > tolerance_seconds: + return "HMAC timestamp outside allowed window" + + expected = compute_signature( + secret, + build_signature_message(timestamp, method, path, body_str), + ) + if not hmac.compare_digest(expected, signature): + return "Invalid HMAC signature" + + return None + + +class AdminHMACMiddleware: + """Pure ASGI middleware so the request body can be replayed to downstream apps.""" + + def __init__(self, app: ASGIApp) -> None: + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + path: str = scope["path"] + if not is_admin_path(path): + await self.app(scope, receive, send) + return + + body = await _read_body(receive) + headers = Headers(scope=scope) + method = scope["method"] + body_str = body.decode("utf-8") if body else "" + + error = verify_hmac_headers( + method=method, + path=path, + body_str=body_str, + signature=headers.get(SIGNATURE_HEADER), + timestamp=headers.get(TIMESTAMP_HEADER), + secret=settings.TARS_HMAC_SECRET, + tolerance_seconds=settings.TARS_HMAC_TIMESTAMP_TOLERANCE_SECONDS, + ) + if error is not None: + logger.warning( + "Admin HMAC verification failed", + extra={ + "path": path, + "method": method, + "reason": error, + }, + ) + response = JSONResponse(status_code=401, content={"detail": error}) + await response(scope, receive, send) + return + + await self.app(scope, _replay_receive(body), send) + + +async def _read_body(receive: Receive) -> bytes: + body = bytearray() + while True: + message = await receive() + if message["type"] != "http.request": + continue + body.extend(message.get("body", b"")) + if not message.get("more_body", False): + break + return bytes(body) + + +def _replay_receive(body: bytes) -> Receive: + sent = False + + async def receive() -> Message: + nonlocal sent + if sent: + return {"type": "http.disconnect"} + sent = True + return {"type": "http.request", "body": body, "more_body": False} + + return receive diff --git a/app/main.py b/app/main.py index 450cfb0..dc4460d 100644 --- a/app/main.py +++ b/app/main.py @@ -6,6 +6,7 @@ from fastapi.middleware.gzip import GZipMiddleware from app.core.config import settings +from app.core.hmac_auth import AdminHMACMiddleware from app.core.log_config import LoggingConfig dictConfig(LoggingConfig.to_dict()) @@ -25,7 +26,9 @@ async def lifespan(_app: FastAPI): lifespan=lifespan, ) +# Last added = outermost. HMAC must wrap GZip so it sees the raw request body. app.add_middleware(GZipMiddleware) +app.add_middleware(AdminHMACMiddleware) @app.get("/ht", tags=["ops"]) diff --git a/tests/core/__init__.py b/tests/core/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/core/test_hmac_auth.py b/tests/core/test_hmac_auth.py new file mode 100644 index 0000000..e140585 --- /dev/null +++ b/tests/core/test_hmac_auth.py @@ -0,0 +1,54 @@ +import pytest + +from app.core.hmac_auth import ( + build_signature_message, + compute_signature, + is_admin_path, + verify_hmac_headers, +) + + +@pytest.mark.unit +class TestHMACHelpers: + def test_is_admin_path(self): + assert is_admin_path("/api/v2/clickwraps") is True + assert is_admin_path("/api/v2/clickwraps/1") is True + assert is_admin_path("/ht") is False + assert is_admin_path("/api/v3/public/clickwrap/abc") is False + + def test_verify_valid_signature(self): + secret = "test-secret" + timestamp = "1710000000" + method = "POST" + path = "/api/v2/clickwraps" + body = '{"name":"A"}' + signature = compute_signature( + secret, build_signature_message(timestamp, method, path, body) + ) + + assert ( + verify_hmac_headers( + method=method, + path=path, + body_str=body, + signature=signature, + timestamp=timestamp, + secret=secret, + tolerance_seconds=300, + now=1710000000, + ) + is None + ) + + def test_verify_rejects_bad_signature(self): + error = verify_hmac_headers( + method="GET", + path="/api/v2/clickwraps", + body_str="", + signature="nope", + timestamp="1710000000", + secret="test-secret", + tolerance_seconds=300, + now=1710000000, + ) + assert error == "Invalid HMAC signature" From 45ea9d9f286e7c1445676bdca7a40859f7786fda Mon Sep 17 00:00:00 2001 From: Pranav Shridhar Date: Thu, 9 Jul 2026 17:04:01 +0530 Subject: [PATCH 2/2] Add clickwrap admin APIs for packet CRUD and mappings Co-authored-by: Cursor --- app/clickwrap/data/postgres/db_repo.py | 494 +++++++++++++++++- app/clickwrap/domain/domain_models.py | 198 ++++++- app/clickwrap/domain/use_cases/__init__.py | 15 + .../use_cases/create_packet_use_case.py | 61 +++ .../domain/use_cases/get_packet_use_case.py | 25 + .../domain/use_cases/list_packets_use_case.py | 29 + .../update_packet_mappings_use_case.py | 40 ++ .../use_cases/update_packet_use_case.py | 55 ++ app/clickwrap/exceptions.py | 32 ++ app/clickwrap/presentation/router.py | 129 +++++ app/clickwrap/presentation/schemas.py | 186 +++++++ app/core/deps.py | 57 ++ app/main.py | 2 + tests/clickwrap/data/test_packet_db_repo.py | 303 +++++++++-- tests/clickwrap/domain/__init__.py | 0 tests/clickwrap/domain/use_cases/__init__.py | 0 .../use_cases/test_create_packet_use_case.py | 44 ++ tests/clickwrap/presentation/__init__.py | 0 .../presentation/test_clickwrap_api.py | 247 +++++++++ tests/factories.py | 5 +- 20 files changed, 1857 insertions(+), 65 deletions(-) create mode 100644 app/clickwrap/domain/use_cases/create_packet_use_case.py create mode 100644 app/clickwrap/domain/use_cases/get_packet_use_case.py create mode 100644 app/clickwrap/domain/use_cases/list_packets_use_case.py create mode 100644 app/clickwrap/domain/use_cases/update_packet_mappings_use_case.py create mode 100644 app/clickwrap/domain/use_cases/update_packet_use_case.py create mode 100644 app/clickwrap/exceptions.py create mode 100644 app/clickwrap/presentation/router.py create mode 100644 app/clickwrap/presentation/schemas.py create mode 100644 app/core/deps.py create mode 100644 tests/clickwrap/domain/__init__.py create mode 100644 tests/clickwrap/domain/use_cases/__init__.py create mode 100644 tests/clickwrap/domain/use_cases/test_create_packet_use_case.py create mode 100644 tests/clickwrap/presentation/__init__.py create mode 100644 tests/clickwrap/presentation/test_clickwrap_api.py diff --git a/app/clickwrap/data/postgres/db_repo.py b/app/clickwrap/data/postgres/db_repo.py index 6bcf27a..1f9c86a 100644 --- a/app/clickwrap/data/postgres/db_repo.py +++ b/app/clickwrap/data/postgres/db_repo.py @@ -1,17 +1,55 @@ import logging +import math +import re +import uuid +from datetime import UTC, datetime +from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession from app.clickwrap.domain.domain_models import ( + DEFAULT_CLICKWRAP_TEXT_SINGLE_CHECKBOX, + AgreementSummary, + AgreementVersionSummary, + ClickwrapTextDomainModel, + PacketAgreementMappingDomainModel, + PacketAgreementMappingListDomainModel, + PacketCreateRequest, + PacketDetailDomainModel, PacketDomainModel, PacketFilterRequest, PacketListDomainModel, + PacketMappingsUpdateRequest, + PacketMinimalDomainModel, + PacketPaginatedListDomainModel, + PacketPaginatedRequest, + PacketSettingsDomainModel, + PacketUpdateRequest, +) +from app.clickwrap.exceptions import ( + AgreementNotFoundError, + PacketNotFoundError, + PacketPaginationError, +) +from app.db.enums import AgreementUiType, AgreementVersionStatus +from app.db.models import ( + Agreement, + AgreementVersion, + Packet, + PacketAgreementMapping, + PacketSettings, ) -from app.db.models import Packet logger = logging.getLogger(__name__) +def slugify(value: str) -> str: + value = value.lower().strip() + value = re.sub(r"[^\w\s-]", "", value) + value = re.sub(r"[-\s]+", "-", value) + return value.strip("-") + + class PacketDBRepository: def __init__(self, session: AsyncSession) -> None: self._session = session @@ -23,42 +61,452 @@ async def filter(self, request: PacketFilterRequest) -> PacketListDomainModel: Workspace isolation is always applied — request.workspace_id is a required field and is always included in the WHERE clause. - - Soft-delete: starts from Packet.objects() (excludes deleted) by default. - Pass include_deleted=True to use Packet.objects_including_deleted() instead. - - Args: - request.workspace_id: Required. Scopes results to this workspace. - request.packet_ids: Optional. Further filters to only these IDs. - request.include_deleted: When True, includes soft-deleted packets. """ - base = ( - Packet.objects_including_deleted() - if request.include_deleted - else Packet.objects() + query = self._base_filter_query(request) + result = await self._session.execute(query) + packets = result.scalars().unique().all() + + logger.info( + "PacketDBRepository.filter completed", + extra={ + "workspace_id": request.workspace_id, + "packet_count": len(packets), + "filtered_by_ids": request.packet_ids is not None, + "include_deleted": request.include_deleted, + }, ) - query = ( - base.where(Packet.workspace_id == request.workspace_id) - .order_by(Packet.id.desc()) + return PacketListDomainModel( + items=[self._to_packet_domain(p, include_settings=False) for p in packets] ) + async def get_count(self, request: PacketFilterRequest) -> int: + query = select(func.count()).select_from(Packet) + if not request.include_deleted: + query = query.where(Packet.is_deleted.is_(False)) + query = query.where(Packet.workspace_id == request.workspace_id) if request.packet_ids is not None: query = query.where(Packet.id.in_(request.packet_ids)) + if request.name_slugs is not None: + query = query.where(Packet.name_slug.in_(request.name_slugs)) + result = await self._session.execute(query) + return int(result.scalar_one()) + + async def get_by_id(self, packet_id: int, workspace_id: int) -> PacketDomainModel: + packet = await self._get_packet_orm( + packet_id=packet_id, workspace_id=workspace_id, load_settings=True + ) + settings = await self._get_settings_orm(packet.packet_settings_id, workspace_id) + return self._to_packet_domain(packet, settings=settings, include_settings=True) + + async def get_detail(self, packet_id: int, workspace_id: int) -> PacketDetailDomainModel: + packet = await self._get_packet_orm( + packet_id=packet_id, workspace_id=workspace_id, load_settings=True + ) + agreements = await self._get_mapped_agreement_summaries( + packet_id=packet_id, workspace_id=workspace_id + ) + settings = self._to_settings_domain( + await self._get_settings_orm(packet.packet_settings_id, workspace_id) + ) + return PacketDetailDomainModel( + id=packet.id, + name=packet.name, + name_slug=packet.name_slug, + description=packet.description, + public_id=packet.public_id, + workspace_id=packet.workspace_id, + settings=settings, + created_by_org_user_id=packet.created_by_org_user_id, + updated_by_org_user_id=packet.updated_by_org_user_id, + updated_by_org_user_at=packet.updated_by_org_user_at, + agreements=agreements, + ) + + async def get_paginated_list( + self, request: PacketPaginatedRequest + ) -> PacketPaginatedListDomainModel: + filter_request = PacketFilterRequest(workspace_id=request.workspace_id) + total_results = await self.get_count(filter_request) + total_pages = max(1, math.ceil(total_results / request.limit)) if total_results else 1 + if request.page > total_pages and total_results > 0: + raise PacketPaginationError() + + offset = (request.page - 1) * request.limit + query = ( + Packet.objects() + .where(Packet.workspace_id == request.workspace_id) + .order_by(Packet.id.desc()) + .offset(offset) + .limit(request.limit) + ) result = await self._session.execute(query) packets = result.scalars().all() + return PacketPaginatedListDomainModel( + page=request.page, + limit=request.limit, + total_results=total_results, + results=[ + PacketMinimalDomainModel( + id=p.id, + name=p.name, + description=p.description, + created_by_org_user_id=p.created_by_org_user_id, + updated_by_org_user_at=p.updated_by_org_user_at, + public_id=p.public_id, + ) + for p in packets + ], + ) + + async def create( + self, + request: PacketCreateRequest, + workspace_id: int, + org_user_id: int, + ) -> PacketDomainModel: + now = datetime.now(UTC).replace(tzinfo=None) + settings = PacketSettings( + workspace_id=workspace_id, + created_by_org_user_id=org_user_id, + updated_by_org_user_id=org_user_id, + agreement_ui_type=AgreementUiType.SINGLE_CHECKBOX, + clickwrap_texts=[{"text": DEFAULT_CLICKWRAP_TEXT_SINGLE_CHECKBOX}], + whitelisted_domains=[], + show_audit_click_status=False, + send_executed_audit_email=False, + allow_all_domains=False, + ) + self._session.add(settings) + await self._session.flush() + + packet = Packet( + workspace_id=workspace_id, + created_by_org_user_id=org_user_id, + updated_by_org_user_id=org_user_id, + updated_by_org_user_at=now, + name=request.name, + name_slug=slugify(request.name), + description=request.description, + public_id=uuid.uuid4(), + packet_settings_id=settings.id, + ) + self._session.add(packet) + await self._session.flush() + await self._session.refresh(packet) + await self._session.refresh(settings) + logger.info( - "PacketDBRepository.filter completed", + "PacketDBRepository.create completed", extra={ - "workspace_id": request.workspace_id, - "packet_count": len(packets), - "filtered_by_ids": request.packet_ids is not None, - "include_deleted": request.include_deleted, + "workspace_id": workspace_id, + "packet_id": packet.id, + "org_user_id": org_user_id, }, ) + return self._to_packet_domain(packet, settings=settings, include_settings=True) - return PacketListDomainModel( - items=[PacketDomainModel.model_validate(p) for p in packets] + async def update( + self, + packet_id: int, + workspace_id: int, + org_user_id: int, + request: PacketUpdateRequest, + ) -> PacketDomainModel: + packet = await self._get_packet_orm( + packet_id=packet_id, workspace_id=workspace_id, load_settings=True + ) + settings = await self._get_settings_orm(packet.packet_settings_id, workspace_id) + now = datetime.now(UTC).replace(tzinfo=None) + + if request.name is not None: + packet.name = request.name + packet.name_slug = slugify(request.name) + if request.description is not None: + packet.description = request.description + + if request.settings is not None: + settings_update = request.settings + if settings_update.type is not None: + settings.agreement_ui_type = settings_update.type + if settings_update.whitelisted_domains is not None: + settings.whitelisted_domains = settings_update.whitelisted_domains + if settings_update.clickwrap_texts is not None: + settings.clickwrap_texts = [ + {"text": item.text} for item in settings_update.clickwrap_texts + ] + if settings_update.send_executed_audit_email is not None: + settings.send_executed_audit_email = settings_update.send_executed_audit_email + if settings_update.show_audit_click_status is not None: + settings.show_audit_click_status = settings_update.show_audit_click_status + if settings_update.allow_all_domains is not None: + settings.allow_all_domains = settings_update.allow_all_domains + settings.updated_by_org_user_id = org_user_id + + packet.updated_by_org_user_id = org_user_id + packet.updated_by_org_user_at = now + + await self._session.flush() + await self._session.refresh(packet) + await self._session.refresh(settings) + + logger.info( + "PacketDBRepository.update completed", + extra={ + "workspace_id": workspace_id, + "packet_id": packet.id, + "org_user_id": org_user_id, + }, + ) + return self._to_packet_domain(packet, settings=settings, include_settings=True) + + async def update_mappings( + self, + packet_id: int, + workspace_id: int, + org_user_id: int, + request: PacketMappingsUpdateRequest, + ) -> PacketAgreementMappingListDomainModel: + await self._get_packet_orm( + packet_id=packet_id, workspace_id=workspace_id, load_settings=False + ) + + if request.remove_agreement_ids: + await self._soft_delete_mappings( + packet_id=packet_id, + workspace_id=workspace_id, + org_user_id=org_user_id, + agreement_ids=request.remove_agreement_ids, + ) + + if request.add_agreement_ids: + await self._add_mappings( + packet_id=packet_id, + workspace_id=workspace_id, + org_user_id=org_user_id, + agreement_ids=request.add_agreement_ids, + ) + + return await self.list_mappings(packet_id=packet_id, workspace_id=workspace_id) + + async def list_mappings( + self, packet_id: int, workspace_id: int + ) -> PacketAgreementMappingListDomainModel: + query = ( + PacketAgreementMapping.objects() + .where(PacketAgreementMapping.workspace_id == workspace_id) + .where(PacketAgreementMapping.packet_id == packet_id) + .order_by(PacketAgreementMapping.id.asc()) + ) + result = await self._session.execute(query) + mappings = result.scalars().all() + return PacketAgreementMappingListDomainModel( + items=[ + PacketAgreementMappingDomainModel.model_validate(mapping) for mapping in mappings + ] + ) + + def _base_filter_query(self, request: PacketFilterRequest): + base = Packet.objects_including_deleted() if request.include_deleted else Packet.objects() + query = base.where(Packet.workspace_id == request.workspace_id).order_by(Packet.id.desc()) + if request.packet_ids is not None: + query = query.where(Packet.id.in_(request.packet_ids)) + if request.name_slugs is not None: + query = query.where(Packet.name_slug.in_(request.name_slugs)) + return query + + async def _get_packet_orm( + self, packet_id: int, workspace_id: int, load_settings: bool + ) -> Packet: + query = ( + Packet.objects() + .where(Packet.id == packet_id) + .where(Packet.workspace_id == workspace_id) + ) + result = await self._session.execute(query) + packet = result.scalar_one_or_none() + if packet is None: + raise PacketNotFoundError(packet_id) + if load_settings: + # Ensure settings row is loaded in the same session + await self._get_settings_orm(packet.packet_settings_id, workspace_id) + return packet + + async def _get_settings_orm(self, settings_id: int, workspace_id: int) -> PacketSettings: + query = ( + PacketSettings.objects() + .where(PacketSettings.id == settings_id) + .where(PacketSettings.workspace_id == workspace_id) + ) + result = await self._session.execute(query) + settings = result.scalar_one_or_none() + if settings is None: + raise PacketNotFoundError() + return settings + + async def _get_mapped_agreement_summaries( + self, packet_id: int, workspace_id: int + ) -> list[AgreementSummary]: + mapping_query = ( + PacketAgreementMapping.objects() + .where(PacketAgreementMapping.packet_id == packet_id) + .where(PacketAgreementMapping.workspace_id == workspace_id) + .order_by(PacketAgreementMapping.id.asc()) + ) + mapping_result = await self._session.execute(mapping_query) + mappings = mapping_result.scalars().all() + if not mappings: + return [] + + agreement_ids = [m.agreement_id for m in mappings] + agreement_query = ( + Agreement.objects() + .where(Agreement.workspace_id == workspace_id) + .where(Agreement.id.in_(agreement_ids)) + ) + agreement_result = await self._session.execute(agreement_query) + agreements = {agreement.id: agreement for agreement in agreement_result.scalars().all()} + + version_query = ( + AgreementVersion.objects() + .where(AgreementVersion.workspace_id == workspace_id) + .where(AgreementVersion.agreement_id.in_(agreement_ids)) + .where(AgreementVersion.is_current.is_(True)) + .where(AgreementVersion.status == AgreementVersionStatus.PUBLISHED) + ) + version_result = await self._session.execute(version_query) + versions = {version.agreement_id: version for version in version_result.scalars().all()} + + summaries: list[AgreementSummary] = [] + for mapping in mappings: + if mapping.agreement_id not in agreements: + continue + version = versions.get(mapping.agreement_id) + summaries.append( + AgreementSummary( + id=mapping.agreement_id, + current_version=( + AgreementVersionSummary( + id=version.id, + name=version.name, + status=version.status, + modified_by_org_user_at=version.modified_by_org_user_at, + version_number=version.version_number, + sub_version_number=version.sub_version_number, + public_url=None, + ) + if version is not None + else None + ), + ) + ) + return summaries + + async def _soft_delete_mappings( + self, + packet_id: int, + workspace_id: int, + org_user_id: int, + agreement_ids: list[int], + ) -> None: + now = datetime.now(UTC).replace(tzinfo=None) + query = ( + PacketAgreementMapping.objects() + .where(PacketAgreementMapping.packet_id == packet_id) + .where(PacketAgreementMapping.workspace_id == workspace_id) + .where(PacketAgreementMapping.agreement_id.in_(agreement_ids)) + ) + result = await self._session.execute(query) + mappings = result.scalars().all() + for mapping in mappings: + mapping.is_deleted = True + mapping.deleted_at = now + mapping.deleted_by_org_user_id = org_user_id + mapping.updated_by_org_user_id = org_user_id + await self._session.flush() + + async def _add_mappings( + self, + packet_id: int, + workspace_id: int, + org_user_id: int, + agreement_ids: list[int], + ) -> None: + unique_ids = list(dict.fromkeys(agreement_ids)) + agreement_query = ( + Agreement.objects() + .where(Agreement.workspace_id == workspace_id) + .where(Agreement.id.in_(unique_ids)) + ) + agreement_result = await self._session.execute(agreement_query) + found_ids = {agreement.id for agreement in agreement_result.scalars().all()} + missing = [agreement_id for agreement_id in unique_ids if agreement_id not in found_ids] + if missing: + raise AgreementNotFoundError(missing[0]) + + existing_query = ( + PacketAgreementMapping.objects() + .where(PacketAgreementMapping.packet_id == packet_id) + .where(PacketAgreementMapping.workspace_id == workspace_id) + .where(PacketAgreementMapping.agreement_id.in_(unique_ids)) + ) + existing_result = await self._session.execute(existing_query) + already_mapped = {mapping.agreement_id for mapping in existing_result.scalars().all()} + + for agreement_id in unique_ids: + if agreement_id in already_mapped: + continue + self._session.add( + PacketAgreementMapping( + workspace_id=workspace_id, + created_by_org_user_id=org_user_id, + updated_by_org_user_id=org_user_id, + packet_id=packet_id, + agreement_id=agreement_id, + ) + ) + await self._session.flush() + + def _to_settings_domain(self, settings: PacketSettings) -> PacketSettingsDomainModel: + raw_texts = settings.clickwrap_texts or [] + clickwrap_texts = [ + ClickwrapTextDomainModel(text=item["text"] if isinstance(item, dict) else str(item)) + for item in raw_texts + ] + return PacketSettingsDomainModel( + id=settings.id, + type=AgreementUiType(settings.agreement_ui_type), + clickwrap_texts=clickwrap_texts, + whitelisted_domains=list(settings.whitelisted_domains or []), + allow_all_domains=settings.allow_all_domains, + send_executed_audit_email=settings.send_executed_audit_email, + show_audit_click_status=settings.show_audit_click_status, + ) + + def _to_packet_domain( + self, + packet: Packet, + *, + settings: PacketSettings | None = None, + include_settings: bool, + ) -> PacketDomainModel: + return PacketDomainModel( + id=packet.id, + workspace_id=packet.workspace_id, + name=packet.name, + name_slug=packet.name_slug, + description=packet.description, + public_id=packet.public_id, + packet_settings_id=packet.packet_settings_id, + created_by_org_user_id=packet.created_by_org_user_id, + updated_by_org_user_id=packet.updated_by_org_user_id, + updated_by_org_user_at=packet.updated_by_org_user_at, + created_at=packet.created_at, + updated_at=packet.updated_at, + is_deleted=packet.is_deleted, + deleted_at=packet.deleted_at, + deleted_by_org_user_id=packet.deleted_by_org_user_id, + settings=self._to_settings_domain(settings) if include_settings and settings else None, ) diff --git a/app/clickwrap/domain/domain_models.py b/app/clickwrap/domain/domain_models.py index 7ad6196..73c44dc 100644 --- a/app/clickwrap/domain/domain_models.py +++ b/app/clickwrap/domain/domain_models.py @@ -1,7 +1,46 @@ +import re import uuid from datetime import datetime -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator + +from app.db.enums import AgreementUiType + +DEFAULT_CLICKWRAP_TEXT_SINGLE_CHECKBOX = "I agree to {#agreement_list#} as per the laws." +AGREEMENT_ID_REGEX = r"\{#agreement_([1-9]\d*)#\}" +CLICKWRAP_DOMAIN_VALIDATION_REGEX = ( + r"^(?!.*\.\.)(?!.*-$)(?!.*_$)(?!.*\.$)(?!-)(?!_)" + r"[a-zA-Z0-9_](?:[a-zA-Z0-9_-]*[a-zA-Z0-9_]\.)+[a-zA-Z]{1,63}$" +) + + +class ClickwrapTextDomainModel(BaseModel): + text: str = Field(min_length=1) + + +class PacketSettingsDomainModel(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + type: AgreementUiType + clickwrap_texts: list[ClickwrapTextDomainModel] + whitelisted_domains: list[str] = Field(default_factory=list) + allow_all_domains: bool = False + send_executed_audit_email: bool = False + show_audit_click_status: bool = False + + @field_validator("whitelisted_domains") + @classmethod + def validate_whitelisted_domains(cls, domains: list[str]) -> list[str]: + for domain in domains: + if not re.match(CLICKWRAP_DOMAIN_VALIDATION_REGEX, domain): + raise ValueError(f"Invalid domain in settings: {domain}.") + return domains + + @model_validator(mode="after") + def validate_clickwrap_texts_for_type(self) -> "PacketSettingsDomainModel": + _validate_clickwrap_texts_for_type(self.type, self.clickwrap_texts) + return self class PacketDomainModel(BaseModel): @@ -27,6 +66,8 @@ class PacketDomainModel(BaseModel): deleted_at: datetime | None deleted_by_org_user_id: int | None + settings: PacketSettingsDomainModel | None = None + class PacketListDomainModel(BaseModel): items: list[PacketDomainModel] @@ -35,4 +76,159 @@ class PacketListDomainModel(BaseModel): class PacketFilterRequest(BaseModel): workspace_id: int packet_ids: list[int] | None = None + name_slugs: list[str] | None = None include_deleted: bool = False + + +class PacketCreateRequest(BaseModel): + name: str = Field(min_length=1, max_length=100) + description: str = Field(min_length=1, max_length=500) + + +class PacketSettingsUpdateRequest(BaseModel): + type: AgreementUiType | None = None + whitelisted_domains: list[str] | None = None + clickwrap_texts: list[ClickwrapTextDomainModel] | None = None + send_executed_audit_email: bool | None = None + show_audit_click_status: bool | None = None + allow_all_domains: bool | None = None + + @field_validator("whitelisted_domains") + @classmethod + def validate_whitelisted_domains(cls, domains: list[str] | None) -> list[str] | None: + if domains is None: + return None + for domain in domains: + if not re.match(CLICKWRAP_DOMAIN_VALIDATION_REGEX, domain): + raise ValueError(f"Invalid domain in settings: {domain}.") + return domains + + @model_validator(mode="after") + def validate_optional_clickwrap_texts(self) -> "PacketSettingsUpdateRequest": + if self.clickwrap_texts is None: + return self + ui_type = self.type or AgreementUiType.SINGLE_CHECKBOX + _validate_clickwrap_texts_for_type(ui_type, self.clickwrap_texts) + return self + + +class PacketUpdateRequest(BaseModel): + name: str | None = Field(default=None, min_length=1, max_length=100) + description: str | None = Field(default=None, max_length=500) + settings: PacketSettingsUpdateRequest | None = None + + +class PacketPaginatedRequest(BaseModel): + workspace_id: int + page: int = Field(default=1, ge=1) + limit: int = Field(default=10, ge=1, le=100) + + +class PacketMinimalDomainModel(BaseModel): + id: int + name: str + description: str | None + created_by_org_user_id: int + updated_by_org_user_at: datetime | None + public_id: uuid.UUID + + +class PacketPaginatedListDomainModel(BaseModel): + page: int + limit: int + total_results: int + results: list[PacketMinimalDomainModel] + + +class AgreementVersionSummary(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + name: str + status: str | None + modified_by_org_user_at: datetime | None = None + version_number: int + sub_version_number: int + public_url: str | None = None + + @property + def full_version_number(self) -> str: + return f"{self.version_number}.{self.sub_version_number}" + + +class AgreementSummary(BaseModel): + id: int + current_version: AgreementVersionSummary | None = None + + +class PacketDetailDomainModel(BaseModel): + id: int + name: str + name_slug: str + description: str | None + public_id: uuid.UUID + workspace_id: int + settings: PacketSettingsDomainModel + created_by_org_user_id: int + updated_by_org_user_id: int | None + updated_by_org_user_at: datetime | None + agreements: list[AgreementSummary] = Field(default_factory=list) + + +class PacketAgreementMappingDomainModel(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + packet_id: int + agreement_id: int + workspace_id: int + created_by_org_user_id: int + + +class PacketAgreementMappingListDomainModel(BaseModel): + items: list[PacketAgreementMappingDomainModel] + + +class PacketMappingsUpdateRequest(BaseModel): + add_agreement_ids: list[int] | None = None + remove_agreement_ids: list[int] | None = None + + @model_validator(mode="after") + def require_at_least_one_change(self) -> "PacketMappingsUpdateRequest": + if not self.add_agreement_ids and not self.remove_agreement_ids: + raise ValueError( + "At least one of add_agreement_ids or remove_agreement_ids must be present" + ) + return self + + +def _validate_clickwrap_texts_for_type( + ui_type: AgreementUiType, clickwrap_texts: list[ClickwrapTextDomainModel] +) -> None: + if ui_type in (AgreementUiType.SINGLE_CHECKBOX, AgreementUiType.INLINE): + if len(clickwrap_texts) > 1: + raise ValueError( + "clickwrap_texts cannot have more than one item when type is single checkbox/inline" + ) + if "{#agreement_list#}" not in clickwrap_texts[0].text: + raise ValueError( + "Compulsory to have agreement_list in text field when type is single checkbox/inline" + ) + if re.search(AGREEMENT_ID_REGEX, clickwrap_texts[0].text): + raise ValueError( + "Not allowed to have agreement_id in text field when type is single checkbox/inline" + ) + return + + if ui_type == AgreementUiType.MULTIPLE_CHECKBOX: + if len(clickwrap_texts) == 1: + raise ValueError("clickwrap_texts cannot have one item when type is multiple checkbox") + for clickwrap_text in clickwrap_texts: + if "{#agreement_list#}" in clickwrap_text.text: + raise ValueError( + "Not allowed to have agreement_list in text field when type is multiple checkbox" + ) + if not re.search(AGREEMENT_ID_REGEX, clickwrap_text.text): + raise ValueError( + "Compulsory to have agreement_id in text field when type is multiple checkbox" + ) diff --git a/app/clickwrap/domain/use_cases/__init__.py b/app/clickwrap/domain/use_cases/__init__.py index e69de29..0603fd3 100644 --- a/app/clickwrap/domain/use_cases/__init__.py +++ b/app/clickwrap/domain/use_cases/__init__.py @@ -0,0 +1,15 @@ +from app.clickwrap.domain.use_cases.create_packet_use_case import CreatePacketUseCase +from app.clickwrap.domain.use_cases.get_packet_use_case import GetPacketUseCase +from app.clickwrap.domain.use_cases.list_packets_use_case import ListPacketsUseCase +from app.clickwrap.domain.use_cases.update_packet_mappings_use_case import ( + UpdatePacketMappingsUseCase, +) +from app.clickwrap.domain.use_cases.update_packet_use_case import UpdatePacketUseCase + +__all__ = [ + "CreatePacketUseCase", + "GetPacketUseCase", + "ListPacketsUseCase", + "UpdatePacketMappingsUseCase", + "UpdatePacketUseCase", +] diff --git a/app/clickwrap/domain/use_cases/create_packet_use_case.py b/app/clickwrap/domain/use_cases/create_packet_use_case.py new file mode 100644 index 0000000..45ea01c --- /dev/null +++ b/app/clickwrap/domain/use_cases/create_packet_use_case.py @@ -0,0 +1,61 @@ +import logging + +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from app.clickwrap.data.postgres.db_repo import PacketDBRepository, slugify +from app.clickwrap.domain.domain_models import ( + PacketCreateRequest, + PacketDomainModel, + PacketFilterRequest, +) +from app.clickwrap.exceptions import PacketInvalidNameError + +logger = logging.getLogger(__name__) + + +class CreatePacketUseCase: + def __init__(self, session: AsyncSession, repo: PacketDBRepository | None = None) -> None: + self._session = session + self._repo = repo or PacketDBRepository(session) + + async def execute( + self, + request: PacketCreateRequest, + workspace_id: int, + org_user_id: int, + ) -> PacketDomainModel: + name_slug = slugify(request.name) + existing_count = await self._repo.get_count( + PacketFilterRequest(workspace_id=workspace_id, name_slugs=[name_slug]) + ) + if existing_count > 0: + logger.info( + "CreatePacketUseCase rejected duplicate name", + extra={ + "workspace_id": workspace_id, + "org_user_id": org_user_id, + "packet_name": request.name, + "name_slug": name_slug, + }, + ) + raise PacketInvalidNameError() + + try: + packet = await self._repo.create( + request=request, + workspace_id=workspace_id, + org_user_id=org_user_id, + ) + except IntegrityError as exc: + raise PacketInvalidNameError() from exc + + logger.info( + "CreatePacketUseCase completed", + extra={ + "workspace_id": workspace_id, + "org_user_id": org_user_id, + "packet_id": packet.id, + }, + ) + return packet diff --git a/app/clickwrap/domain/use_cases/get_packet_use_case.py b/app/clickwrap/domain/use_cases/get_packet_use_case.py new file mode 100644 index 0000000..af39211 --- /dev/null +++ b/app/clickwrap/domain/use_cases/get_packet_use_case.py @@ -0,0 +1,25 @@ +import logging + +from sqlalchemy.ext.asyncio import AsyncSession + +from app.clickwrap.data.postgres.db_repo import PacketDBRepository +from app.clickwrap.domain.domain_models import PacketDetailDomainModel + +logger = logging.getLogger(__name__) + + +class GetPacketUseCase: + def __init__(self, session: AsyncSession, repo: PacketDBRepository | None = None) -> None: + self._repo = repo or PacketDBRepository(session) + + async def execute(self, packet_id: int, workspace_id: int) -> PacketDetailDomainModel: + detail = await self._repo.get_detail(packet_id=packet_id, workspace_id=workspace_id) + logger.info( + "GetPacketUseCase completed", + extra={ + "workspace_id": workspace_id, + "packet_id": packet_id, + "agreement_count": len(detail.agreements), + }, + ) + return detail diff --git a/app/clickwrap/domain/use_cases/list_packets_use_case.py b/app/clickwrap/domain/use_cases/list_packets_use_case.py new file mode 100644 index 0000000..cb2d749 --- /dev/null +++ b/app/clickwrap/domain/use_cases/list_packets_use_case.py @@ -0,0 +1,29 @@ +import logging + +from sqlalchemy.ext.asyncio import AsyncSession + +from app.clickwrap.data.postgres.db_repo import PacketDBRepository +from app.clickwrap.domain.domain_models import ( + PacketPaginatedListDomainModel, + PacketPaginatedRequest, +) + +logger = logging.getLogger(__name__) + + +class ListPacketsUseCase: + def __init__(self, session: AsyncSession, repo: PacketDBRepository | None = None) -> None: + self._repo = repo or PacketDBRepository(session) + + async def execute(self, request: PacketPaginatedRequest) -> PacketPaginatedListDomainModel: + result = await self._repo.get_paginated_list(request) + logger.info( + "ListPacketsUseCase completed", + extra={ + "workspace_id": request.workspace_id, + "page": request.page, + "limit": request.limit, + "total_results": result.total_results, + }, + ) + return result diff --git a/app/clickwrap/domain/use_cases/update_packet_mappings_use_case.py b/app/clickwrap/domain/use_cases/update_packet_mappings_use_case.py new file mode 100644 index 0000000..d9068e7 --- /dev/null +++ b/app/clickwrap/domain/use_cases/update_packet_mappings_use_case.py @@ -0,0 +1,40 @@ +import logging + +from sqlalchemy.ext.asyncio import AsyncSession + +from app.clickwrap.data.postgres.db_repo import PacketDBRepository +from app.clickwrap.domain.domain_models import ( + PacketAgreementMappingListDomainModel, + PacketMappingsUpdateRequest, +) + +logger = logging.getLogger(__name__) + + +class UpdatePacketMappingsUseCase: + def __init__(self, session: AsyncSession, repo: PacketDBRepository | None = None) -> None: + self._repo = repo or PacketDBRepository(session) + + async def execute( + self, + packet_id: int, + workspace_id: int, + org_user_id: int, + request: PacketMappingsUpdateRequest, + ) -> PacketAgreementMappingListDomainModel: + result = await self._repo.update_mappings( + packet_id=packet_id, + workspace_id=workspace_id, + org_user_id=org_user_id, + request=request, + ) + logger.info( + "UpdatePacketMappingsUseCase completed", + extra={ + "workspace_id": workspace_id, + "packet_id": packet_id, + "org_user_id": org_user_id, + "mapping_count": len(result.items), + }, + ) + return result diff --git a/app/clickwrap/domain/use_cases/update_packet_use_case.py b/app/clickwrap/domain/use_cases/update_packet_use_case.py new file mode 100644 index 0000000..986e04a --- /dev/null +++ b/app/clickwrap/domain/use_cases/update_packet_use_case.py @@ -0,0 +1,55 @@ +import logging + +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from app.clickwrap.data.postgres.db_repo import PacketDBRepository, slugify +from app.clickwrap.domain.domain_models import ( + PacketDetailDomainModel, + PacketFilterRequest, + PacketUpdateRequest, +) +from app.clickwrap.exceptions import PacketInvalidNameError + +logger = logging.getLogger(__name__) + + +class UpdatePacketUseCase: + def __init__(self, session: AsyncSession, repo: PacketDBRepository | None = None) -> None: + self._repo = repo or PacketDBRepository(session) + + async def execute( + self, + packet_id: int, + workspace_id: int, + org_user_id: int, + request: PacketUpdateRequest, + ) -> PacketDetailDomainModel: + if request.name is not None: + name_slug = slugify(request.name) + existing = await self._repo.filter( + PacketFilterRequest(workspace_id=workspace_id, name_slugs=[name_slug]) + ) + if any(item.id != packet_id for item in existing.items): + raise PacketInvalidNameError() + + try: + await self._repo.update( + packet_id=packet_id, + workspace_id=workspace_id, + org_user_id=org_user_id, + request=request, + ) + except IntegrityError as exc: + raise PacketInvalidNameError() from exc + + detail = await self._repo.get_detail(packet_id=packet_id, workspace_id=workspace_id) + logger.info( + "UpdatePacketUseCase completed", + extra={ + "workspace_id": workspace_id, + "packet_id": packet_id, + "org_user_id": org_user_id, + }, + ) + return detail diff --git a/app/clickwrap/exceptions.py b/app/clickwrap/exceptions.py new file mode 100644 index 0000000..a91c721 --- /dev/null +++ b/app/clickwrap/exceptions.py @@ -0,0 +1,32 @@ +class ClickwrapError(Exception): + """Base error for clickwrap/packet domain failures.""" + + +class PacketNotFoundError(ClickwrapError): + def __init__(self, packet_id: int | None = None) -> None: + self.packet_id = packet_id + message = f"Packet {packet_id} not found" if packet_id is not None else "Packet not found" + super().__init__(message) + + +class PacketInvalidNameError(ClickwrapError): + def __init__(self) -> None: + super().__init__( + "Clickthrough with the same name already exists. Please rename your clickthrough." + ) + + +class PacketPaginationError(ClickwrapError): + def __init__(self) -> None: + super().__init__("List Pagination out of range.") + + +class AgreementNotFoundError(ClickwrapError): + def __init__(self, agreement_id: int | None = None) -> None: + self.agreement_id = agreement_id + message = ( + f"Agreement {agreement_id} not found" + if agreement_id is not None + else "Agreement not found" + ) + super().__init__(message) diff --git a/app/clickwrap/presentation/router.py b/app/clickwrap/presentation/router.py new file mode 100644 index 0000000..a5214f5 --- /dev/null +++ b/app/clickwrap/presentation/router.py @@ -0,0 +1,129 @@ +import logging + +from fastapi import APIRouter, Depends, HTTPException, Query, status +from sqlalchemy.ext.asyncio import AsyncSession + +from app.clickwrap.domain.domain_models import ( + PacketCreateRequest, + PacketMappingsUpdateRequest, + PacketPaginatedRequest, + PacketUpdateRequest, +) +from app.clickwrap.domain.use_cases.create_packet_use_case import CreatePacketUseCase +from app.clickwrap.domain.use_cases.get_packet_use_case import GetPacketUseCase +from app.clickwrap.domain.use_cases.list_packets_use_case import ListPacketsUseCase +from app.clickwrap.domain.use_cases.update_packet_mappings_use_case import ( + UpdatePacketMappingsUseCase, +) +from app.clickwrap.domain.use_cases.update_packet_use_case import UpdatePacketUseCase +from app.clickwrap.exceptions import ( + AgreementNotFoundError, + PacketInvalidNameError, + PacketNotFoundError, + PacketPaginationError, +) +from app.clickwrap.presentation.schemas import ( + PacketDetailResponse, + PacketMappingResponse, + PacketMinimalResponse, + PacketPaginatedResponse, +) +from app.core.deps import RequestContext, get_db_session, get_request_context + +logger = logging.getLogger(__name__) + +router = APIRouter(prefix="/api/v2/clickwraps", tags=["clickwraps"]) + + +@router.get("") +@router.get("/") +async def list_packets( + page: int = Query(default=1, ge=1), + limit: int = Query(default=10, ge=1, le=100), + ctx: RequestContext = Depends(get_request_context), + session: AsyncSession = Depends(get_db_session), +) -> PacketPaginatedResponse: + try: + result = await ListPacketsUseCase(session).execute( + PacketPaginatedRequest(workspace_id=ctx.workspace_id, page=page, limit=limit) + ) + except PacketPaginationError as exc: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc + return PacketPaginatedResponse.from_domain(result) + + +@router.post("", status_code=status.HTTP_201_CREATED) +@router.post("/", status_code=status.HTTP_201_CREATED) +async def create_packet( + body: PacketCreateRequest, + ctx: RequestContext = Depends(get_request_context), + session: AsyncSession = Depends(get_db_session), +) -> PacketMinimalResponse: + try: + packet = await CreatePacketUseCase(session).execute( + request=body, + workspace_id=ctx.workspace_id, + org_user_id=ctx.org_user_id, + ) + except PacketInvalidNameError as exc: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc + return PacketMinimalResponse.from_packet_domain(packet) + + +@router.get("/{clickwrap_id}") +@router.get("/{clickwrap_id}/") +async def get_packet( + clickwrap_id: int, + ctx: RequestContext = Depends(get_request_context), + session: AsyncSession = Depends(get_db_session), +) -> PacketDetailResponse: + try: + detail = await GetPacketUseCase(session).execute( + packet_id=clickwrap_id, workspace_id=ctx.workspace_id + ) + except PacketNotFoundError as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc + return PacketDetailResponse.from_domain(detail) + + +@router.patch("/{clickwrap_id}") +@router.patch("/{clickwrap_id}/") +async def update_packet( + clickwrap_id: int, + body: PacketUpdateRequest, + ctx: RequestContext = Depends(get_request_context), + session: AsyncSession = Depends(get_db_session), +) -> PacketDetailResponse: + try: + detail = await UpdatePacketUseCase(session).execute( + packet_id=clickwrap_id, + workspace_id=ctx.workspace_id, + org_user_id=ctx.org_user_id, + request=body, + ) + except PacketNotFoundError as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc + except PacketInvalidNameError as exc: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc + return PacketDetailResponse.from_domain(detail) + + +@router.patch("/{clickwrap_id}/clickwrap-agreement-mappings") +async def update_packet_mappings( + clickwrap_id: int, + body: PacketMappingsUpdateRequest, + ctx: RequestContext = Depends(get_request_context), + session: AsyncSession = Depends(get_db_session), +) -> list[PacketMappingResponse]: + try: + result = await UpdatePacketMappingsUseCase(session).execute( + packet_id=clickwrap_id, + workspace_id=ctx.workspace_id, + org_user_id=ctx.org_user_id, + request=body, + ) + except PacketNotFoundError as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc + except AgreementNotFoundError as exc: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc + return [PacketMappingResponse.from_domain(item) for item in result.items] diff --git a/app/clickwrap/presentation/schemas.py b/app/clickwrap/presentation/schemas.py new file mode 100644 index 0000000..fb4321b --- /dev/null +++ b/app/clickwrap/presentation/schemas.py @@ -0,0 +1,186 @@ +from datetime import datetime +from uuid import UUID + +from pydantic import BaseModel, ConfigDict, Field, computed_field + +from app.clickwrap.domain.domain_models import ( + AgreementSummary, + AgreementVersionSummary, + PacketAgreementMappingDomainModel, + PacketDetailDomainModel, + PacketDomainModel, + PacketMinimalDomainModel, + PacketPaginatedListDomainModel, + PacketSettingsDomainModel, +) +from app.db.enums import AgreementUiType + + +class ClickwrapTextResponse(BaseModel): + text: str + + +class PacketSettingsResponse(BaseModel): + id: int + type: AgreementUiType + clickwrap_texts: list[ClickwrapTextResponse] + whitelisted_domains: list[str] = Field(default_factory=list) + allow_all_domains: bool = False + send_executed_audit_email: bool = False + show_audit_click_status: bool = False + + @classmethod + def from_domain(cls, settings: PacketSettingsDomainModel) -> "PacketSettingsResponse": + return cls( + id=settings.id, + type=settings.type, + clickwrap_texts=[ + ClickwrapTextResponse(text=item.text) for item in settings.clickwrap_texts + ], + whitelisted_domains=settings.whitelisted_domains, + allow_all_domains=settings.allow_all_domains, + send_executed_audit_email=settings.send_executed_audit_email, + show_audit_click_status=settings.show_audit_click_status, + ) + + +class AgreementVersionResponse(BaseModel): + id: int + name: str + status: str | None + modified_by_org_user_at: datetime | None = None + version_number: int + sub_version_number: int + public_url: str | None = None + + @computed_field # type: ignore[prop-decorator] + @property + def full_version_number(self) -> str: + return f"{self.version_number}.{self.sub_version_number}" + + @classmethod + def from_domain(cls, version: AgreementVersionSummary) -> "AgreementVersionResponse": + return cls( + id=version.id, + name=version.name, + status=version.status, + modified_by_org_user_at=version.modified_by_org_user_at, + version_number=version.version_number, + sub_version_number=version.sub_version_number, + public_url=version.public_url, + ) + + +class AgreementResponse(BaseModel): + id: int + current_version: AgreementVersionResponse | None = None + + @classmethod + def from_domain(cls, agreement: AgreementSummary) -> "AgreementResponse": + return cls( + id=agreement.id, + current_version=( + AgreementVersionResponse.from_domain(agreement.current_version) + if agreement.current_version is not None + else None + ), + ) + + +class PacketMinimalResponse(BaseModel): + id: int + name: str + description: str | None + created_by_org_user_id: int + updated_by_org_user_at: datetime | None + public_id: UUID + + @classmethod + def from_domain(cls, packet: PacketMinimalDomainModel) -> "PacketMinimalResponse": + return cls( + id=packet.id, + name=packet.name, + description=packet.description, + created_by_org_user_id=packet.created_by_org_user_id, + updated_by_org_user_at=packet.updated_by_org_user_at, + public_id=packet.public_id, + ) + + @classmethod + def from_packet_domain(cls, packet: PacketDomainModel) -> "PacketMinimalResponse": + return cls( + id=packet.id, + name=packet.name, + description=packet.description, + created_by_org_user_id=packet.created_by_org_user_id, + updated_by_org_user_at=packet.updated_by_org_user_at, + public_id=packet.public_id, + ) + + +class PacketPaginatedResponse(BaseModel): + page: int + limit: int + total_results: int + results: list[PacketMinimalResponse] + + @classmethod + def from_domain(cls, page: PacketPaginatedListDomainModel) -> "PacketPaginatedResponse": + return cls( + page=page.page, + limit=page.limit, + total_results=page.total_results, + results=[PacketMinimalResponse.from_domain(item) for item in page.results], + ) + + +class PacketDetailResponse(BaseModel): + id: int + name: str + name_slug: str + description: str | None + public_id: UUID + workspace_id: int + settings: PacketSettingsResponse + created_by_org_user_id: int + updated_by_org_user_id: int | None + updated_by_org_user_at: datetime | None + agreements: list[AgreementResponse] = Field(default_factory=list) + + @classmethod + def from_domain(cls, detail: PacketDetailDomainModel) -> "PacketDetailResponse": + return cls( + id=detail.id, + name=detail.name, + name_slug=detail.name_slug, + description=detail.description, + public_id=detail.public_id, + workspace_id=detail.workspace_id, + settings=PacketSettingsResponse.from_domain(detail.settings), + created_by_org_user_id=detail.created_by_org_user_id, + updated_by_org_user_id=detail.updated_by_org_user_id, + updated_by_org_user_at=detail.updated_by_org_user_at, + agreements=[ + AgreementResponse.from_domain(agreement) for agreement in detail.agreements + ], + ) + + +class PacketMappingResponse(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: int + packet_id: int + agreement_id: int + workspace_id: int + created_by_org_user_id: int + + @classmethod + def from_domain(cls, mapping: PacketAgreementMappingDomainModel) -> "PacketMappingResponse": + return cls( + id=mapping.id, + packet_id=mapping.packet_id, + agreement_id=mapping.agreement_id, + workspace_id=mapping.workspace_id, + created_by_org_user_id=mapping.created_by_org_user_id, + ) diff --git a/app/core/deps.py b/app/core/deps.py new file mode 100644 index 0000000..a282c36 --- /dev/null +++ b/app/core/deps.py @@ -0,0 +1,57 @@ +from collections.abc import AsyncGenerator +from dataclasses import dataclass + +from fastapi import Header, HTTPException, status +from sqlalchemy.ext.asyncio import AsyncSession + +from app.db.postgres import AsyncSessionLocal + + +@dataclass(frozen=True, slots=True) +class RequestContext: + workspace_id: int + org_user_id: int + user_id: int | None = None + request_id: str | None = None + client_ip: str | None = None + + +async def get_db_session() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSessionLocal() as session: + try: + yield session + await session.commit() + except Exception: + await session.rollback() + raise + + +def get_request_context( + x_workspace_id: str | None = Header(default=None, alias="X-Workspace-ID"), + x_org_user_id: str | None = Header(default=None, alias="X-Org-User-ID"), + x_user_id: str | None = Header(default=None, alias="X-User-ID"), + x_request_id: str | None = Header(default=None, alias="X-Request-ID"), + x_client_ip: str | None = Header(default=None, alias="X-Client-IP"), +) -> RequestContext: + if not x_workspace_id or not x_org_user_id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="X-Workspace-ID and X-Org-User-ID headers are required", + ) + try: + workspace_id = int(x_workspace_id) + org_user_id = int(x_org_user_id) + user_id = int(x_user_id) if x_user_id else None + except ValueError as exc: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Workspace and org user header values must be integers", + ) from exc + + return RequestContext( + workspace_id=workspace_id, + org_user_id=org_user_id, + user_id=user_id, + request_id=x_request_id, + client_ip=x_client_ip, + ) diff --git a/app/main.py b/app/main.py index dc4460d..5d78a11 100644 --- a/app/main.py +++ b/app/main.py @@ -5,6 +5,7 @@ from fastapi import FastAPI from fastapi.middleware.gzip import GZipMiddleware +from app.clickwrap.presentation.router import router as clickwrap_router from app.core.config import settings from app.core.hmac_auth import AdminHMACMiddleware from app.core.log_config import LoggingConfig @@ -29,6 +30,7 @@ async def lifespan(_app: FastAPI): # Last added = outermost. HMAC must wrap GZip so it sees the raw request body. app.add_middleware(GZipMiddleware) app.add_middleware(AdminHMACMiddleware) +app.include_router(clickwrap_router) @app.get("/ht", tags=["ops"]) diff --git a/tests/clickwrap/data/test_packet_db_repo.py b/tests/clickwrap/data/test_packet_db_repo.py index 477046c..91de874 100644 --- a/tests/clickwrap/data/test_packet_db_repo.py +++ b/tests/clickwrap/data/test_packet_db_repo.py @@ -1,48 +1,48 @@ +import uuid +from datetime import UTC, datetime + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + from app.clickwrap.data.postgres.db_repo import PacketDBRepository from app.clickwrap.domain.domain_models import ( - PacketDomainModel, + ClickwrapTextDomainModel, + PacketCreateRequest, PacketFilterRequest, - PacketListDomainModel, + PacketMappingsUpdateRequest, + PacketPaginatedRequest, + PacketSettingsUpdateRequest, + PacketUpdateRequest, ) +from app.clickwrap.exceptions import ( + AgreementNotFoundError, + PacketNotFoundError, + PacketPaginationError, +) +from app.db.enums import AgreementUiType, AgreementVersionStatus +from app.db.models import Agreement, AgreementVersion, PacketAgreementMapping from tests.base_db_repo_test import BaseDBRepoTestCase from tests.factories import PacketFactory -class TestPacketDBRepositoryFilter(BaseDBRepoTestCase): - async def test_filter_workspace_isolation(self, db_session): - """ - Each workspace sees only its own packets, ordered by id DESC. - - Creates packets in two workspaces and asserts that a filter scoped to - workspace A returns only workspace A's packets (newest first), and - workspace B's packets are invisible to it. - """ - # setup +class TestPacketDBRepository(BaseDBRepoTestCase): + async def test_filter_workspace_isolation(self, db_session: AsyncSession): ws_a, ws_b = 100, 101 packet_a1 = await PacketFactory.create(db_session, workspace_id=ws_a) packet_a2 = await PacketFactory.create(db_session, workspace_id=ws_a) - await PacketFactory.create(db_session, workspace_id=ws_b) # must not appear + await PacketFactory.create(db_session, workspace_id=ws_b) repo = PacketDBRepository(db_session) - - # make the call result = await repo.filter(PacketFilterRequest(workspace_id=ws_a)) - # assert — workspace isolation + id DESC ordering - assert result == PacketListDomainModel( - items=[ - PacketDomainModel.model_validate(packet_a2), - PacketDomainModel.model_validate(packet_a1), - ] - ) + assert [p.id for p in result.items] == [packet_a2.id, packet_a1.id] - async def test_filter_by_packet_ids(self, db_session): + async def test_filter_by_packet_ids(self, db_session: AsyncSession): """ When packet_ids is provided, only those IDs are returned. Cross-workspace IDs are silently excluded (workspace isolation still applies even when an explicit ID list is passed). """ - # setup ws_a, ws_b = 200, 201 packet_a1 = await PacketFactory.create(db_session, workspace_id=ws_a) packet_a2 = await PacketFactory.create(db_session, workspace_id=ws_a) @@ -50,44 +50,267 @@ async def test_filter_by_packet_ids(self, db_session): repo = PacketDBRepository(db_session) - # only packet_a1 requested result = await repo.filter( PacketFilterRequest(workspace_id=ws_a, packet_ids=[packet_a1.id]) ) assert len(result.items) == 1 assert result.items[0].id == packet_a1.id - # packet_a2 + cross-workspace packet_b — packet_b must be excluded result = await repo.filter( PacketFilterRequest(workspace_id=ws_a, packet_ids=[packet_a2.id, packet_b.id]) ) assert len(result.items) == 1 assert result.items[0].id == packet_a2.id - async def test_filter_soft_delete_behaviour(self, db_session): + async def test_filter_soft_delete_behaviour(self, db_session: AsyncSession): """ Soft-deleted packets are excluded by default (Packet.objects()). Passing include_deleted=True switches to Packet.objects_including_deleted() and returns all packets regardless of deletion status. """ - # setup workspace_id = 300 active = await PacketFactory.create(db_session, workspace_id=workspace_id) - deleted = await PacketFactory.create( - db_session, workspace_id=workspace_id, is_deleted=True - ) + deleted = await PacketFactory.create(db_session, workspace_id=workspace_id, is_deleted=True) repo = PacketDBRepository(db_session) - # default — deleted packet must not appear - result = await repo.filter(PacketFilterRequest(workspace_id=workspace_id)) - assert len(result.items) == 1 - assert result.items[0].id == active.id + default_result = await repo.filter(PacketFilterRequest(workspace_id=workspace_id)) + assert [p.id for p in default_result.items] == [active.id] - # opt-in — both packets returned - result = await repo.filter( + including_deleted = await repo.filter( PacketFilterRequest(workspace_id=workspace_id, include_deleted=True) ) - result_ids = {item.id for item in result.items} - assert active.id in result_ids - assert deleted.id in result_ids + assert {p.id for p in including_deleted.items} == {active.id, deleted.id} + + async def test_create_packet_with_default_settings(self, db_session: AsyncSession): + repo = PacketDBRepository(db_session) + created = await repo.create( + request=PacketCreateRequest( + name="Vendor Onboarding", + description="Onboarding pack", + ), + workspace_id=42, + org_user_id=7, + ) + + assert created.id > 0 + assert created.workspace_id == 42 + assert created.name == "Vendor Onboarding" + assert created.name_slug == "vendor-onboarding" + assert created.description == "Onboarding pack" + assert created.created_by_org_user_id == 7 + assert created.updated_by_org_user_id == 7 + assert created.settings is not None + assert created.settings.type == AgreementUiType.SINGLE_CHECKBOX + assert created.settings.clickwrap_texts[0].text == ( + "I agree to {#agreement_list#} as per the laws." + ) + + fetched = await repo.get_by_id(packet_id=created.id, workspace_id=42) + assert fetched.id == created.id + assert fetched.settings is not None + assert fetched.settings.id == created.settings.id + + async def test_get_by_id_raises_for_other_workspace(self, db_session: AsyncSession): + packet = await PacketFactory.create(db_session, workspace_id=10) + repo = PacketDBRepository(db_session) + + with pytest.raises(PacketNotFoundError): + await repo.get_by_id(packet_id=packet.id, workspace_id=99) + + async def test_get_by_id_excludes_soft_deleted(self, db_session: AsyncSession): + packet = await PacketFactory.create(db_session, workspace_id=11, is_deleted=True) + repo = PacketDBRepository(db_session) + + with pytest.raises(PacketNotFoundError): + await repo.get_by_id(packet_id=packet.id, workspace_id=11) + + async def test_get_count_by_name_slug(self, db_session: AsyncSession): + await PacketFactory.create(db_session, workspace_id=12, name_slug="vendor-onboarding") + repo = PacketDBRepository(db_session) + + count = await repo.get_count( + PacketFilterRequest(workspace_id=12, name_slugs=["vendor-onboarding"]) + ) + assert count == 1 + + other = await repo.get_count(PacketFilterRequest(workspace_id=12, name_slugs=["missing"])) + assert other == 0 + + async def test_get_paginated_list(self, db_session: AsyncSession): + workspace_id = 13 + p1 = await PacketFactory.create(db_session, workspace_id=workspace_id) + p2 = await PacketFactory.create(db_session, workspace_id=workspace_id) + p3 = await PacketFactory.create(db_session, workspace_id=workspace_id) + repo = PacketDBRepository(db_session) + + page_1 = await repo.get_paginated_list( + PacketPaginatedRequest(workspace_id=workspace_id, page=1, limit=2) + ) + assert page_1.total_results == 3 + assert [r.id for r in page_1.results] == [p3.id, p2.id] + + page_2 = await repo.get_paginated_list( + PacketPaginatedRequest(workspace_id=workspace_id, page=2, limit=2) + ) + assert [r.id for r in page_2.results] == [p1.id] + + with pytest.raises(PacketPaginationError): + await repo.get_paginated_list( + PacketPaginatedRequest(workspace_id=workspace_id, page=5, limit=2) + ) + + async def test_update_packet_name_and_settings(self, db_session: AsyncSession): + packet = await PacketFactory.create( + db_session, + workspace_id=14, + name="Old Name", + name_slug="old-name", + description="Old desc", + ) + repo = PacketDBRepository(db_session) + + updated = await repo.update( + packet_id=packet.id, + workspace_id=14, + org_user_id=99, + request=PacketUpdateRequest( + name="New Name", + description="New desc", + settings=PacketSettingsUpdateRequest( + type=AgreementUiType.SINGLE_CHECKBOX, + allow_all_domains=True, + whitelisted_domains=["app.acme.com"], + clickwrap_texts=[ + ClickwrapTextDomainModel( + text="I agree to {#agreement_list#} as per the laws." + ) + ], + send_executed_audit_email=True, + show_audit_click_status=True, + ), + ), + ) + + assert updated.name == "New Name" + assert updated.name_slug == "new-name" + assert updated.description == "New desc" + assert updated.updated_by_org_user_id == 99 + assert updated.settings is not None + assert updated.settings.allow_all_domains is True + assert updated.settings.whitelisted_domains == ["app.acme.com"] + assert updated.settings.send_executed_audit_email is True + assert updated.settings.show_audit_click_status is True + + async def test_update_mappings_add_and_remove(self, db_session: AsyncSession): + workspace_id = 15 + org_user_id = 3 + packet = await PacketFactory.create(db_session, workspace_id=workspace_id) + agreement_1 = await _create_agreement(db_session, workspace_id, org_user_id) + agreement_2 = await _create_agreement(db_session, workspace_id, org_user_id) + agreement_3 = await _create_agreement(db_session, workspace_id, org_user_id) + + existing = PacketAgreementMapping( + workspace_id=workspace_id, + created_by_org_user_id=org_user_id, + packet_id=packet.id, + agreement_id=agreement_1.id, + ) + db_session.add(existing) + await db_session.flush() + + repo = PacketDBRepository(db_session) + result = await repo.update_mappings( + packet_id=packet.id, + workspace_id=workspace_id, + org_user_id=org_user_id, + request=PacketMappingsUpdateRequest( + add_agreement_ids=[agreement_2.id, agreement_3.id], + remove_agreement_ids=[agreement_1.id], + ), + ) + + assert sorted(m.agreement_id for m in result.items) == [ + agreement_2.id, + agreement_3.id, + ] + + async def test_update_mappings_unknown_agreement_raises(self, db_session: AsyncSession): + packet = await PacketFactory.create(db_session, workspace_id=16) + repo = PacketDBRepository(db_session) + + with pytest.raises(AgreementNotFoundError): + await repo.update_mappings( + packet_id=packet.id, + workspace_id=16, + org_user_id=1, + request=PacketMappingsUpdateRequest(add_agreement_ids=[999999]), + ) + + async def test_get_detail_includes_mapped_agreements(self, db_session: AsyncSession): + workspace_id = 17 + org_user_id = 4 + packet = await PacketFactory.create(db_session, workspace_id=workspace_id) + agreement = await _create_agreement(db_session, workspace_id, org_user_id) + version = await _create_published_version( + db_session, workspace_id, org_user_id, agreement.id + ) + db_session.add( + PacketAgreementMapping( + workspace_id=workspace_id, + created_by_org_user_id=org_user_id, + packet_id=packet.id, + agreement_id=agreement.id, + ) + ) + await db_session.flush() + + repo = PacketDBRepository(db_session) + detail = await repo.get_detail(packet_id=packet.id, workspace_id=workspace_id) + + assert detail.id == packet.id + assert detail.settings is not None + assert len(detail.agreements) == 1 + assert detail.agreements[0].id == agreement.id + assert detail.agreements[0].current_version is not None + assert detail.agreements[0].current_version.id == version.id + + +async def _create_agreement( + session: AsyncSession, workspace_id: int, org_user_id: int +) -> Agreement: + agreement = Agreement( + workspace_id=workspace_id, + created_by_org_user_id=org_user_id, + url_slug=f"agr-{uuid.uuid4().hex[:12]}", + ) + session.add(agreement) + await session.flush() + await session.refresh(agreement) + return agreement + + +async def _create_published_version( + session: AsyncSession, + workspace_id: int, + org_user_id: int, + agreement_id: int, +) -> AgreementVersion: + version = AgreementVersion( + workspace_id=workspace_id, + created_by_org_user_id=org_user_id, + agreement_id=agreement_id, + name="MSA", + name_slug="msa", + status=AgreementVersionStatus.PUBLISHED, + version_number=1, + sub_version_number=0, + is_current=True, + modified_by_org_user_at=datetime.now(UTC).replace(tzinfo=None), + published_at=datetime.now(UTC).replace(tzinfo=None), + published_by_org_user_id=org_user_id, + ) + session.add(version) + await session.flush() + await session.refresh(version) + return version diff --git a/tests/clickwrap/domain/__init__.py b/tests/clickwrap/domain/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/clickwrap/domain/use_cases/__init__.py b/tests/clickwrap/domain/use_cases/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/clickwrap/domain/use_cases/test_create_packet_use_case.py b/tests/clickwrap/domain/use_cases/test_create_packet_use_case.py new file mode 100644 index 0000000..64c4175 --- /dev/null +++ b/tests/clickwrap/domain/use_cases/test_create_packet_use_case.py @@ -0,0 +1,44 @@ +from unittest.mock import AsyncMock + +import pytest + +from app.clickwrap.domain.domain_models import ( + PacketCreateRequest, + PacketDomainModel, + PacketFilterRequest, +) +from app.clickwrap.domain.use_cases.create_packet_use_case import CreatePacketUseCase +from app.clickwrap.exceptions import PacketInvalidNameError + + +@pytest.mark.unit +class TestCreatePacketUseCase: + async def test_execute_creates_when_name_available(self): + repo = AsyncMock() + repo.get_count.return_value = 0 + created = AsyncMock(spec=PacketDomainModel) + created.id = 10 + repo.create.return_value = created + + use_case = CreatePacketUseCase(session=AsyncMock(), repo=repo) + request = PacketCreateRequest(name="Packet A", description="desc") + result = await use_case.execute(request=request, workspace_id=1, org_user_id=2) + + assert result is created + repo.get_count.assert_awaited_once_with( + PacketFilterRequest(workspace_id=1, name_slugs=["packet-a"]) + ) + repo.create.assert_awaited_once_with(request=request, workspace_id=1, org_user_id=2) + + async def test_execute_raises_on_duplicate_name(self): + repo = AsyncMock() + repo.get_count.return_value = 1 + use_case = CreatePacketUseCase(session=AsyncMock(), repo=repo) + + with pytest.raises(PacketInvalidNameError): + await use_case.execute( + request=PacketCreateRequest(name="Packet A", description="desc"), + workspace_id=1, + org_user_id=2, + ) + repo.create.assert_not_awaited() diff --git a/tests/clickwrap/presentation/__init__.py b/tests/clickwrap/presentation/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/clickwrap/presentation/test_clickwrap_api.py b/tests/clickwrap/presentation/test_clickwrap_api.py new file mode 100644 index 0000000..d26bdee --- /dev/null +++ b/tests/clickwrap/presentation/test_clickwrap_api.py @@ -0,0 +1,247 @@ +import json +import time + +import pytest +from httpx import ASGITransport, AsyncClient +from sqlalchemy.ext.asyncio import AsyncSession + +from app.core.config import settings +from app.core.deps import get_db_session +from app.core.hmac_auth import ( + build_signature_message, + compute_signature, +) +from app.db.models import Agreement, PacketAgreementMapping +from app.main import app +from tests.factories import PacketFactory + + +def _sign_headers( + method: str, + path: str, + body: dict | None = None, + *, + workspace_id: int = 42, + org_user_id: int = 7, + timestamp: int | None = None, + signature: str | None = None, +) -> dict[str, str]: + body_str = json.dumps(body, sort_keys=True) if body is not None else "" + ts = str(timestamp if timestamp is not None else int(time.time())) + message = build_signature_message(ts, method, path, body_str) + sig = ( + signature + if signature is not None + else compute_signature(settings.TARS_HMAC_SECRET, message) + ) + return { + "X-Workspace-ID": str(workspace_id), + "X-Org-User-ID": str(org_user_id), + "X-User-ID": "100", + "Content-Type": "application/json", + "X-Clickwrap-Timestamp": ts, + "X-Clickwrap-Signature": sig, + } + + +@pytest.fixture +async def api_client(db_session: AsyncSession): + async def _override_db_session(): + yield db_session + + app.dependency_overrides[get_db_session] = _override_db_session + transport = ASGITransport(app=app) + async with AsyncClient(transport=transport, base_url="http://test") as client: + yield client + app.dependency_overrides.clear() + + +@pytest.mark.integration +class TestClickwrapAPI: + async def test_create_list_get_update_flow(self, api_client): + create_body = {"name": "Vendor Onboarding", "description": "Pack desc"} + create_resp = await api_client.post( + "/api/v2/clickwraps", + headers=_sign_headers("POST", "/api/v2/clickwraps", create_body), + content=json.dumps(create_body, sort_keys=True), + ) + assert create_resp.status_code == 201 + created = create_resp.json() + assert created["name"] == "Vendor Onboarding" + assert created["description"] == "Pack desc" + assert "contract_type_id" not in created + packet_id = created["id"] + + list_resp = await api_client.get( + "/api/v2/clickwraps", + headers=_sign_headers("GET", "/api/v2/clickwraps"), + params={"page": 1, "limit": 10}, + ) + assert list_resp.status_code == 200 + listed = list_resp.json() + assert listed["total_results"] == 1 + assert listed["results"][0]["id"] == packet_id + + get_path = f"/api/v2/clickwraps/{packet_id}" + get_resp = await api_client.get( + get_path, + headers=_sign_headers("GET", get_path), + ) + assert get_resp.status_code == 200 + detail = get_resp.json() + assert detail["name"] == "Vendor Onboarding" + assert detail["settings"]["type"] == "SINGLE_CHECKBOX" + assert detail["agreements"] == [] + assert "contract_type_id" not in detail + + patch_body = { + "name": "Vendor Onboarding v2", + "settings": { + "allow_all_domains": True, + "whitelisted_domains": ["app.acme.com"], + }, + } + patch_path = f"/api/v2/clickwraps/{packet_id}" + patch_resp = await api_client.patch( + patch_path, + headers=_sign_headers("PATCH", patch_path, patch_body), + content=json.dumps(patch_body, sort_keys=True), + ) + assert patch_resp.status_code == 200 + updated = patch_resp.json() + assert updated["name"] == "Vendor Onboarding v2" + assert updated["name_slug"] == "vendor-onboarding-v2" + assert updated["settings"]["allow_all_domains"] is True + assert updated["settings"]["whitelisted_domains"] == ["app.acme.com"] + + async def test_create_duplicate_name_returns_400(self, api_client, db_session): + await PacketFactory.create( + db_session, + workspace_id=42, + name="Dup Name", + name_slug="dup-name", + ) + body = {"name": "Dup Name", "description": "desc"} + resp = await api_client.post( + "/api/v2/clickwraps", + headers=_sign_headers("POST", "/api/v2/clickwraps", body), + content=json.dumps(body, sort_keys=True), + ) + assert resp.status_code == 400 + assert "same name already exists" in resp.json()["detail"] + + async def test_get_missing_packet_returns_404(self, api_client): + path = "/api/v2/clickwraps/999999" + resp = await api_client.get( + path, + headers=_sign_headers("GET", path), + ) + assert resp.status_code == 404 + + async def test_missing_workspace_header_returns_400(self, api_client): + headers = _sign_headers("GET", "/api/v2/clickwraps") + del headers["X-Workspace-ID"] + resp = await api_client.get( + "/api/v2/clickwraps", + headers=headers, + ) + assert resp.status_code == 400 + + async def test_update_mappings(self, api_client, db_session): + workspace_id = 42 + org_user_id = 7 + packet = await PacketFactory.create(db_session, workspace_id=workspace_id) + agreement = Agreement( + workspace_id=workspace_id, + created_by_org_user_id=org_user_id, + url_slug="msa-api-test", + ) + db_session.add(agreement) + await db_session.flush() + + path = f"/api/v2/clickwraps/{packet.id}/clickwrap-agreement-mappings" + add_body = {"add_agreement_ids": [agreement.id]} + add_resp = await api_client.patch( + path, + headers=_sign_headers( + "PATCH", + path, + add_body, + workspace_id=workspace_id, + org_user_id=org_user_id, + ), + content=json.dumps(add_body, sort_keys=True), + ) + assert add_resp.status_code == 200 + mappings = add_resp.json() + assert len(mappings) == 1 + assert mappings[0]["agreement_id"] == agreement.id + assert mappings[0]["packet_id"] == packet.id + + remove_body = {"remove_agreement_ids": [agreement.id]} + remove_resp = await api_client.patch( + path, + headers=_sign_headers( + "PATCH", + path, + remove_body, + workspace_id=workspace_id, + org_user_id=org_user_id, + ), + content=json.dumps(remove_body, sort_keys=True), + ) + assert remove_resp.status_code == 200 + assert remove_resp.json() == [] + + remaining = ( + ( + await db_session.execute( + PacketAgreementMapping.objects().where( + PacketAgreementMapping.packet_id == packet.id + ) + ) + ) + .scalars() + .all() + ) + assert remaining == [] + + +@pytest.mark.integration +class TestAdminHMACMiddleware: + async def test_health_check_skips_hmac(self, api_client): + resp = await api_client.get("/ht") + assert resp.status_code == 200 + assert resp.json() == {"status": "ok"} + + async def test_missing_signature_returns_401(self, api_client): + resp = await api_client.get( + "/api/v2/clickwraps", + headers={ + "X-Workspace-ID": "42", + "X-Org-User-ID": "7", + }, + ) + assert resp.status_code == 401 + assert resp.json()["detail"] == "Missing HMAC signature headers" + + async def test_invalid_signature_returns_401(self, api_client): + headers = _sign_headers( + "GET", + "/api/v2/clickwraps", + signature="deadbeef", + ) + resp = await api_client.get("/api/v2/clickwraps", headers=headers) + assert resp.status_code == 401 + assert resp.json()["detail"] == "Invalid HMAC signature" + + async def test_expired_timestamp_returns_401(self, api_client): + expired = int(time.time()) - settings.TARS_HMAC_TIMESTAMP_TOLERANCE_SECONDS - 10 + headers = _sign_headers( + "GET", + "/api/v2/clickwraps", + timestamp=expired, + ) + resp = await api_client.get("/api/v2/clickwraps", headers=headers) + assert resp.status_code == 401 + assert resp.json()["detail"] == "HMAC timestamp outside allowed window" diff --git a/tests/factories.py b/tests/factories.py index fbac392..39448b0 100644 --- a/tests/factories.py +++ b/tests/factories.py @@ -21,7 +21,10 @@ async def create(cls, session: AsyncSession, **kwargs: object) -> PacketSettings workspace_id=kwargs.get("workspace_id", n), created_by_org_user_id=kwargs.get("created_by_org_user_id", n), agreement_ui_type=kwargs.get("agreement_ui_type", AgreementUiType.SINGLE_CHECKBOX), - clickwrap_texts=kwargs.get("clickwrap_texts", [{"text": "I agree to the terms."}]), + clickwrap_texts=kwargs.get( + "clickwrap_texts", + [{"text": "I agree to {#agreement_list#} as per the laws."}], + ), whitelisted_domains=kwargs.get("whitelisted_domains", []), show_audit_click_status=kwargs.get("show_audit_click_status", False), send_executed_audit_email=kwargs.get("send_executed_audit_email", False),