diff --git a/apps/api/app/api/v1/system.py b/apps/api/app/api/v1/system.py index edf18b8..cdf820b 100644 --- a/apps/api/app/api/v1/system.py +++ b/apps/api/app/api/v1/system.py @@ -7,11 +7,12 @@ ready to serve traffic". from __future__ import annotations +import uuid from typing import Literal import httpx import redis.asyncio as redis_asyncio -from fastapi import APIRouter, Depends, Request +from fastapi import APIRouter, Depends, Request, status from pydantic import BaseModel from sqlalchemy import text from sqlalchemy.ext.asyncio import AsyncSession @@ -23,8 +24,9 @@ from app.core.security import get_client_ip, is_localhost from app.db.session import get_db from app.models.enums import SystemSecretKey from app.models.user import User +from app.schemas.provisioning import PendingProvisioningCreate, PendingProvisioningResponse from app.schemas.system_secret import SetSystemSecretRequest, SystemSecretStatus -from app.services import system_secret_service +from app.services import provisioning_service, system_secret_service from app.services.enrichment_service import estimate_max_credits_per_company logger = get_logger(__name__) @@ -194,3 +196,39 @@ async def system_logs( """Most-recent-first view into the application's live log stream (capped at the last 500 entries app-wide, see `core/logging.py`).""" return [LogEntryResponse(**entry) for entry in await get_recent_logs(settings, limit=100)] + + +@router.get("/system/pending-provisioning", response_model=list[PendingProvisioningResponse]) +async def list_pending_provisioning( + db: AsyncSession = Depends(get_db), _admin: User = Depends(require_admin) +) -> list[PendingProvisioningResponse]: + """Queued not-yet-existing accounts - see app.services.provisioning_service.""" + return [ + PendingProvisioningResponse.model_validate(p, from_attributes=True) + for p in await provisioning_service.list_pending(db) + ] + + +@router.post( + "/system/pending-provisioning", + response_model=PendingProvisioningResponse, + status_code=status.HTTP_201_CREATED, +) +async def create_pending_provisioning( + payload: PendingProvisioningCreate, + db: AsyncSession = Depends(get_db), + _admin: User = Depends(require_admin), +) -> PendingProvisioningResponse: + record = await provisioning_service.create_pending( + db, payload.email, payload.source_user_id, payload.make_admin + ) + return PendingProvisioningResponse.model_validate(record, from_attributes=True) + + +@router.delete("/system/pending-provisioning/{pending_id}", status_code=status.HTTP_204_NO_CONTENT) +async def delete_pending_provisioning( + pending_id: uuid.UUID, + db: AsyncSession = Depends(get_db), + _admin: User = Depends(require_admin), +) -> None: + await provisioning_service.delete_pending(db, pending_id) diff --git a/apps/api/app/models/__init__.py b/apps/api/app/models/__init__.py index 76ba0dd..8aeee15 100644 --- a/apps/api/app/models/__init__.py +++ b/apps/api/app/models/__init__.py @@ -22,6 +22,7 @@ from app.models.notification_destination import ( # noqa: F401 NotificationDestinationCompany, ) from app.models.password_history import PasswordHistoryEntry # noqa: F401 +from app.models.pending_provisioning import PendingProvisioning # noqa: F401 from app.models.refresh_token import RefreshToken # noqa: F401 from app.models.report import Report # noqa: F401 from app.models.snapshot import Snapshot # noqa: F401 diff --git a/apps/api/app/models/pending_provisioning.py b/apps/api/app/models/pending_provisioning.py new file mode 100644 index 0000000..857799b --- /dev/null +++ b/apps/api/app/models/pending_provisioning.py @@ -0,0 +1,31 @@ +"""Lets an admin pre-configure what a not-yet-existing account should get +the moment someone verifies their email with a matching address - copies +of another user's per-user API keys and companies (with their full data: +sources, enrichment, reports, snapshots, etc.), an email notification +destination for the new account's own address, and optionally an admin +promotion. Built for demoing the app to people who don't have accounts +yet without hand-entering everything for them after the fact. + +Applied at email-verification time, not raw registration, since +registering an email address doesn't prove you own it - see +auth_service.verify_email / provisioning_service.apply_if_pending. +""" + +from __future__ import annotations + +import uuid + +from sqlalchemy import Boolean, ForeignKey, String +from sqlalchemy.orm import Mapped, mapped_column + +from app.db.base import Base, TimestampMixin, UUIDPrimaryKeyMixin + + +class PendingProvisioning(Base, UUIDPrimaryKeyMixin, TimestampMixin): + __tablename__ = "pending_provisionings" + + email: Mapped[str] = mapped_column(String(320), unique=True, index=True) + source_user_id: Mapped[uuid.UUID] = mapped_column( + ForeignKey("users.id", ondelete="CASCADE"), index=True + ) + make_admin: Mapped[bool] = mapped_column(Boolean, default=False) diff --git a/apps/api/app/repositories/pending_provisioning_repository.py b/apps/api/app/repositories/pending_provisioning_repository.py new file mode 100644 index 0000000..a37d643 --- /dev/null +++ b/apps/api/app/repositories/pending_provisioning_repository.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +import uuid + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.models.pending_provisioning import PendingProvisioning + + +class PendingProvisioningRepository: + def __init__(self, db: AsyncSession) -> None: + self.db = db + + async def list_all(self) -> list[PendingProvisioning]: + result = await self.db.execute(select(PendingProvisioning)) + return list(result.scalars().all()) + + async def get_by_email(self, email: str) -> PendingProvisioning | None: + result = await self.db.execute( + select(PendingProvisioning).where(PendingProvisioning.email == email.lower()) + ) + return result.scalar_one_or_none() + + async def create( + self, email: str, source_user_id: uuid.UUID, make_admin: bool + ) -> PendingProvisioning: + record = PendingProvisioning( + email=email.lower(), source_user_id=source_user_id, make_admin=make_admin + ) + self.db.add(record) + await self.db.flush() + return record + + async def delete(self, record: PendingProvisioning) -> None: + await self.db.delete(record) + await self.db.flush() diff --git a/apps/api/app/schemas/provisioning.py b/apps/api/app/schemas/provisioning.py new file mode 100644 index 0000000..f0604a2 --- /dev/null +++ b/apps/api/app/schemas/provisioning.py @@ -0,0 +1,20 @@ +from __future__ import annotations + +import uuid +from datetime import datetime + +from pydantic import BaseModel, EmailStr + + +class PendingProvisioningCreate(BaseModel): + email: EmailStr + source_user_id: uuid.UUID + make_admin: bool = False + + +class PendingProvisioningResponse(BaseModel): + id: uuid.UUID + email: str + source_user_id: uuid.UUID + make_admin: bool + created_at: datetime diff --git a/apps/api/app/services/auth_service.py b/apps/api/app/services/auth_service.py index 2973800..a4ae85f 100644 --- a/apps/api/app/services/auth_service.py +++ b/apps/api/app/services/auth_service.py @@ -52,7 +52,7 @@ from app.schemas.auth import ( TokenResponse, VerifyEmailRequest, ) -from app.services import ip_throttle_service, security_email_service +from app.services import ip_throttle_service, provisioning_service, security_email_service EMAIL_CODE_VALID_HOURS = 36 @@ -202,6 +202,11 @@ async def verify_email(db: AsyncSession, client_ip: str, payload: VerifyEmailReq await code_repo.mark_used(record) user.email_verified = True + # Verifying an email is the point at which the app can actually trust + # someone owns this address - if an admin queued this address up for + # pre-provisioning (see provisioning_service), this is where it fires, + # never at raw registration (which proves nothing about ownership). + await provisioning_service.apply_if_pending(db, user) await ip_throttle_service.reset_on_success(db, client_ip, ThrottleAction.VERIFY_EMAIL_CODE) await UserSecurityEventRepository(db).create( user_id=user.id, event_type=SecurityEventType.EMAIL_VERIFIED, ip_address=client_ip diff --git a/apps/api/app/services/provisioning_service.py b/apps/api/app/services/provisioning_service.py new file mode 100644 index 0000000..7ddc388 --- /dev/null +++ b/apps/api/app/services/provisioning_service.py @@ -0,0 +1,295 @@ +"""Pre-provisions a not-yet-existing account: the moment someone verifies +an email address an admin has queued up (see PendingProvisioning), they +get a deep copy of another user's per-user API keys and companies - +including every tab's worth of data (sources, source documents, monitoring +runs, reports, snapshots, detected changes, enrichment) - plus an email +notification destination for their own address linked to the copied +companies, and optionally an admin promotion. + +Applied at email-verification time (see auth_service.verify_email), not +raw registration - registering an email doesn't prove ownership of it, +verifying does. +""" + +from __future__ import annotations + +import uuid +from typing import Any + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.core.errors import ConflictError, NotFoundError +from app.models.company import Company, CompanyAlias, Competitor +from app.models.company_enrichment import CompanyEnrichment +from app.models.detected_change import DetectedChange +from app.models.enums import NotificationType, SeverityLevel +from app.models.monitor_configuration import MonitorConfiguration +from app.models.monitoring_run import MonitoringRun +from app.models.notification_destination import ( + NotificationDestination, + NotificationDestinationCompany, +) +from app.models.pending_provisioning import PendingProvisioning +from app.models.report import Report +from app.models.snapshot import Snapshot +from app.models.source import Source +from app.models.source_document import SourceDocument +from app.models.user import User +from app.models.user_api_key import UserApiKey +from app.repositories.pending_provisioning_repository import PendingProvisioningRepository +from app.repositories.user_repository import UserRepository + + +def _copy_row[ModelT]( + model_cls: type[ModelT], source_row: Any, overrides: dict[str, Any] +) -> ModelT: + """A new ORM instance of the same class, with every column value copied + from `source_row` except whatever `overrides` replaces (always at least + `id`, plus any foreign keys that need remapping to the new owner's + copies of their parents). Generic over every table involved so adding a + new company-related table later doesn't require touching this file.""" + data = { + col.name: getattr(source_row, col.name) for col in model_cls.__table__.columns # type: ignore[attr-defined] + } + data.update(overrides) + return model_cls(**data) + + +async def _clone_company( + db: AsyncSession, source_company_id: uuid.UUID, target_user_id: uuid.UUID +) -> Company: + id_map: dict[uuid.UUID, uuid.UUID] = {} + + def new_id(old_id: uuid.UUID) -> uuid.UUID: + if old_id not in id_map: + id_map[old_id] = uuid.uuid4() + return id_map[old_id] + + company = await db.get(Company, source_company_id) + assert company is not None + + new_company = _copy_row(Company, company, {"id": new_id(company.id), "user_id": target_user_id}) + db.add(new_company) + await db.flush() + + for alias in ( + await db.execute(select(CompanyAlias).where(CompanyAlias.company_id == source_company_id)) + ).scalars(): + db.add(_copy_row(CompanyAlias, alias, {"id": uuid.uuid4(), "company_id": new_company.id})) + + for competitor in ( + await db.execute(select(Competitor).where(Competitor.company_id == source_company_id)) + ).scalars(): + db.add( + _copy_row(Competitor, competitor, {"id": uuid.uuid4(), "company_id": new_company.id}) + ) + + config = ( + await db.execute( + select(MonitorConfiguration).where(MonitorConfiguration.company_id == source_company_id) + ) + ).scalar_one_or_none() + if config is not None: + db.add( + _copy_row( + MonitorConfiguration, config, {"id": uuid.uuid4(), "company_id": new_company.id} + ) + ) + + enrichment = ( + await db.execute( + select(CompanyEnrichment).where(CompanyEnrichment.company_id == source_company_id) + ) + ).scalar_one_or_none() + if enrichment is not None: + db.add( + _copy_row( + CompanyEnrichment, enrichment, {"id": uuid.uuid4(), "company_id": new_company.id} + ) + ) + + sources = ( + (await db.execute(select(Source).where(Source.company_id == source_company_id))) + .scalars() + .all() + ) + for source in sources: + db.add(_copy_row(Source, source, {"id": new_id(source.id), "company_id": new_company.id})) + await db.flush() + + for doc in ( + await db.execute( + select(SourceDocument).where(SourceDocument.company_id == source_company_id) + ) + ).scalars(): + db.add( + _copy_row( + SourceDocument, + doc, + { + "id": uuid.uuid4(), + "company_id": new_company.id, + "source_id": new_id(doc.source_id), + }, + ) + ) + + runs = ( + ( + await db.execute( + select(MonitoringRun).where(MonitoringRun.company_id == source_company_id) + ) + ) + .scalars() + .all() + ) + for run in runs: + db.add(_copy_row(MonitoringRun, run, {"id": new_id(run.id), "company_id": new_company.id})) + await db.flush() + + for report in ( + await db.execute(select(Report).where(Report.company_id == source_company_id)) + ).scalars(): + db.add( + _copy_row( + Report, + report, + { + "id": uuid.uuid4(), + "company_id": new_company.id, + "monitoring_run_id": ( + new_id(report.monitoring_run_id) if report.monitoring_run_id else None + ), + }, + ) + ) + + snapshots = ( + (await db.execute(select(Snapshot).where(Snapshot.company_id == source_company_id))) + .scalars() + .all() + ) + for snapshot in snapshots: + db.add( + _copy_row( + Snapshot, + snapshot, + { + "id": new_id(snapshot.id), + "company_id": new_company.id, + "source_id": new_id(snapshot.source_id) if snapshot.source_id else None, + "monitoring_run_id": ( + new_id(snapshot.monitoring_run_id) if snapshot.monitoring_run_id else None + ), + }, + ) + ) + await db.flush() + + for change in ( + await db.execute( + select(DetectedChange).where(DetectedChange.company_id == source_company_id) + ) + ).scalars(): + db.add( + _copy_row( + DetectedChange, + change, + { + "id": uuid.uuid4(), + "company_id": new_company.id, + "source_id": new_id(change.source_id) if change.source_id else None, + "monitoring_run_id": ( + new_id(change.monitoring_run_id) if change.monitoring_run_id else None + ), + "previous_snapshot_id": ( + new_id(change.previous_snapshot_id) if change.previous_snapshot_id else None + ), + "current_snapshot_id": ( + new_id(change.current_snapshot_id) if change.current_snapshot_id else None + ), + }, + ) + ) + + return new_company + + +async def _copy_api_keys( + db: AsyncSession, source_user_id: uuid.UUID, target_user_id: uuid.UUID +) -> None: + keys = ( + (await db.execute(select(UserApiKey).where(UserApiKey.user_id == source_user_id))) + .scalars() + .all() + ) + for key in keys: + db.add(_copy_row(UserApiKey, key, {"id": uuid.uuid4(), "user_id": target_user_id})) + + +async def create_pending( + db: AsyncSession, email: str, source_user_id: uuid.UUID, make_admin: bool +) -> PendingProvisioning: + if await UserRepository(db).get_by_email(email) is not None: + raise ConflictError("An account with this email already exists") + if await PendingProvisioningRepository(db).get_by_email(email) is not None: + raise ConflictError("This email is already queued for provisioning") + record = await PendingProvisioningRepository(db).create(email, source_user_id, make_admin) + await db.commit() + return record + + +async def list_pending(db: AsyncSession) -> list[PendingProvisioning]: + return await PendingProvisioningRepository(db).list_all() + + +async def delete_pending(db: AsyncSession, pending_id: uuid.UUID) -> None: + repo = PendingProvisioningRepository(db) + match = next((p for p in await repo.list_all() if p.id == pending_id), None) + if match is None: + raise NotFoundError("Pending provisioning entry not found") + await repo.delete(match) + await db.commit() + + +async def apply_if_pending(db: AsyncSession, new_user: User) -> bool: + """Called right after a fresh account's email gets verified. Returns + True if a pending provisioning entry matched and was applied (and + consumed - a pending entry only ever fires once).""" + pending = await PendingProvisioningRepository(db).get_by_email(new_user.email) + if pending is None: + return False + + await _copy_api_keys(db, pending.source_user_id, new_user.id) + + source_companies = ( + (await db.execute(select(Company).where(Company.user_id == pending.source_user_id))) + .scalars() + .all() + ) + new_company_ids = [ + (await _clone_company(db, company.id, new_user.id)).id for company in source_companies + ] + + if new_company_ids: + destination = NotificationDestination( + user_id=new_user.id, + type=NotificationType.EMAIL, + destination_value=new_user.email, + verified=False, + enabled=True, + minimum_severity=SeverityLevel.MEDIUM, + ) + db.add(destination) + await db.flush() + for company_id in new_company_ids: + db.add( + NotificationDestinationCompany(destination_id=destination.id, company_id=company_id) + ) + + if pending.make_admin: + new_user.is_admin = True + + await PendingProvisioningRepository(db).delete(pending) + return True diff --git a/apps/api/migrations/versions/2c2a563a5edb_add_pending_provisionings_table.py b/apps/api/migrations/versions/2c2a563a5edb_add_pending_provisionings_table.py new file mode 100644 index 0000000..c128037 --- /dev/null +++ b/apps/api/migrations/versions/2c2a563a5edb_add_pending_provisionings_table.py @@ -0,0 +1,45 @@ +"""add pending_provisionings table + +Revision ID: 2c2a563a5edb +Revises: 9e6c80c11da7 +Create Date: 2026-08-06 22:22:22.348339 + +""" +from __future__ import annotations + +from collections.abc import Sequence + +from alembic import op +import sqlalchemy as sa + + +revision: str = '2c2a563a5edb' +down_revision: str | None = '9e6c80c11da7' +branch_labels: Sequence[str] | str | None = None +depends_on: Sequence[str] | str | None = None + + +def upgrade() -> None: + # Autogenerate also proposed adding a redundant UNIQUE constraint on the + # `id` column of every other existing table (a cosmetic artifact of how + # this project's earlier migrations declared unnamed unique constraints, + # unrelated to this change) - stripped out, keeping only the new table. + op.create_table('pending_provisionings', + sa.Column('email', sa.String(length=320), nullable=False), + sa.Column('source_user_id', sa.Uuid(), nullable=False), + sa.Column('make_admin', sa.Boolean(), nullable=False), + sa.Column('id', sa.Uuid(), nullable=False), + sa.Column('created_at', sa.DateTime(timezone=True), nullable=False), + sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False), + sa.ForeignKeyConstraint(['source_user_id'], ['users.id'], ondelete='CASCADE'), + sa.PrimaryKeyConstraint('id'), + sa.UniqueConstraint('id') + ) + op.create_index(op.f('ix_pending_provisionings_email'), 'pending_provisionings', ['email'], unique=True) + op.create_index(op.f('ix_pending_provisionings_source_user_id'), 'pending_provisionings', ['source_user_id'], unique=False) + + +def downgrade() -> None: + op.drop_index(op.f('ix_pending_provisionings_source_user_id'), table_name='pending_provisionings') + op.drop_index(op.f('ix_pending_provisionings_email'), table_name='pending_provisionings') + op.drop_table('pending_provisionings') diff --git a/apps/api/tests/unit/test_provisioning.py b/apps/api/tests/unit/test_provisioning.py new file mode 100644 index 0000000..c2107ff --- /dev/null +++ b/apps/api/tests/unit/test_provisioning.py @@ -0,0 +1,429 @@ +"""Pending account provisioning: an admin queues up an email in advance, +and the moment that email actually verifies a real account, it gets a +deep copy of another user's API keys and companies (every related table, +not just the company row itself), a notification destination for its own +address, and optionally an admin promotion. See provisioning_service.py.""" + +from __future__ import annotations + +import asyncio +import uuid +from datetime import UTC, datetime, timedelta + +import pytest +from sqlalchemy import select + +from app.core.errors import ConflictError, NotFoundError +from app.core.security import hash_email_code, hash_password +from app.db.session import get_sessionmaker +from app.models.company import Company, CompanyAlias, Competitor +from app.models.company_enrichment import CompanyEnrichment +from app.models.detected_change import DetectedChange +from app.models.email_code import EmailCode +from app.models.enums import ( + ApiKeyProvider, + ChangeStatus, + ChangeType, + CompanyStatus, + EmailCodePurpose, + EnrichmentStatus, + MonitoringFrequency, + MonitoringRunStatus, + MonitoringRunTrigger, + ReportType, + SeverityLevel, + SourceStatus, + SourceType, +) +from app.models.monitor_configuration import MonitorConfiguration +from app.models.monitoring_run import MonitoringRun +from app.models.report import Report +from app.models.snapshot import Snapshot +from app.models.source import Source +from app.models.source_document import SourceDocument +from app.models.user import User +from app.models.user_api_key import UserApiKey +from app.repositories.notification_destination_repository import ( + NotificationDestinationRepository, +) +from app.repositories.pending_provisioning_repository import PendingProvisioningRepository +from app.repositories.user_repository import UserRepository +from app.services import provisioning_service + + +async def _make_user(db_session, *, email: str | None = None, is_admin: bool = False) -> User: + user = User( + id=uuid.uuid4(), + email=(email or f"user-{uuid.uuid4().hex[:10]}@example.com").lower(), + password_hash=hash_password("correct-horse-1"), + display_name="Test User", + timezone="UTC", + is_admin=is_admin, + email_verified=True, + ) + db_session.add(user) + await db_session.flush() + return user + + +async def _build_full_company(db_session, owner: User) -> Company: + """One company with at least one row in every table _clone_company + touches, so the deep-clone test actually exercises every FK remap.""" + company = Company( + user_id=owner.id, + name="Acme Corp", + slug=f"acme-{uuid.uuid4().hex[:6]}", + official_website="https://acme.example.com", + status=CompanyStatus.ACTIVE, + ) + db_session.add(company) + await db_session.flush() + + db_session.add(CompanyAlias(company_id=company.id, alias="Acme")) + db_session.add(Competitor(company_id=company.id, name="Rival Co")) + db_session.add( + MonitorConfiguration( + company_id=company.id, frequency_type=MonitoringFrequency.WEEKLY, enabled=True + ) + ) + db_session.add( + CompanyEnrichment( + company_id=company.id, + status=EnrichmentStatus.COMPLETE, + data={"products": [{"name": "Widget"}]}, + errors={}, + credits_spent=5, + ) + ) + + source = Source( + company_id=company.id, + source_type=SourceType.WEBSITE, + name="Homepage", + status=SourceStatus.ACTIVE, + ) + db_session.add(source) + await db_session.flush() + + db_session.add( + SourceDocument( + source_id=source.id, + company_id=company.id, + url="https://acme.example.com/page", + canonical_url="https://acme.example.com/page", + title="Homepage", + content_text="Hello world", + content_hash="abc123", + retrieved_date=datetime.now(UTC), + extraction_method="static_html", + ) + ) + + run = MonitoringRun( + company_id=company.id, + trigger_type=MonitoringRunTrigger.MANUAL, + status=MonitoringRunStatus.SUCCESSFUL, + ) + db_session.add(run) + await db_session.flush() + + db_session.add( + Report( + company_id=company.id, + monitoring_run_id=run.id, + report_type=ReportType.MANUAL, + title="Acme Report", + executive_summary="Summary", + structured_report={"executive_summary": "Summary"}, + markdown_content="# Acme Report", + model_provider="mock", + model_name="mock", + ) + ) + + snapshot = Snapshot( + company_id=company.id, + source_id=source.id, + monitoring_run_id=run.id, + snapshot_type="structured", + hash="snaphash", + ) + db_session.add(snapshot) + await db_session.flush() + + db_session.add( + DetectedChange( + company_id=company.id, + source_id=source.id, + monitoring_run_id=run.id, + previous_snapshot_id=snapshot.id, + current_snapshot_id=snapshot.id, + change_type=ChangeType.CONTENT_MODIFIED, + severity=SeverityLevel.LOW, + status=ChangeStatus.NEW, + summary="Something changed", + confidence_score=0.5, + significance_score=0.5, + ) + ) + + db_session.add( + UserApiKey( + user_id=owner.id, provider=ApiKeyProvider.ANTHROPIC, encrypted_key="encrypted-blob" + ) + ) + + await db_session.commit() + return company + + +@pytest.mark.asyncio +async def test_create_pending_rejects_an_email_that_already_has_an_account(db_session, settings): + owner = await _make_user(db_session) + existing = await _make_user(db_session, email="taken@example.com") + + with pytest.raises(ConflictError): + await provisioning_service.create_pending(db_session, existing.email, owner.id, False) + + +@pytest.mark.asyncio +async def test_create_pending_rejects_a_duplicate_pending_email(db_session, settings): + owner = await _make_user(db_session) + email = f"pending-{uuid.uuid4().hex[:10]}@example.com" + + await provisioning_service.create_pending(db_session, email, owner.id, False) + + with pytest.raises(ConflictError): + await provisioning_service.create_pending(db_session, email, owner.id, False) + + +@pytest.mark.asyncio +async def test_apply_if_pending_is_a_no_op_without_a_matching_entry(db_session, settings): + new_user = await _make_user(db_session, email="nobody-queued@example.com") + + applied = await provisioning_service.apply_if_pending(db_session, new_user) + + assert applied is False + + +@pytest.mark.asyncio +async def test_apply_if_pending_deep_clones_everything_and_is_consumed_once(db_session, settings): + owner = await _make_user(db_session) + company = await _build_full_company(db_session, owner) + email = f"invitee-{uuid.uuid4().hex[:10]}@example.com" + await provisioning_service.create_pending(db_session, email, owner.id, make_admin=True) + + new_user = await _make_user(db_session, email=email) + applied = await provisioning_service.apply_if_pending(db_session, new_user) + await db_session.commit() + + assert applied is True + assert new_user.is_admin is True + + # The pending entry only ever fires once. + assert await PendingProvisioningRepository(db_session).get_by_email(email) is None + applied_again = await provisioning_service.apply_if_pending(db_session, new_user) + assert applied_again is False + + new_company = ( + await db_session.execute(select(Company).where(Company.user_id == new_user.id)) + ).scalar_one() + assert new_company.id != company.id + assert new_company.name == "Acme Corp" + + assert ( + await db_session.execute( + select(CompanyAlias).where(CompanyAlias.company_id == new_company.id) + ) + ).scalar_one().alias == "Acme" + assert ( + await db_session.execute(select(Competitor).where(Competitor.company_id == new_company.id)) + ).scalar_one().name == "Rival Co" + assert ( + ( + await db_session.execute( + select(MonitorConfiguration).where( + MonitorConfiguration.company_id == new_company.id + ) + ) + ) + .scalar_one() + .enabled + ) + + new_source = ( + await db_session.execute(select(Source).where(Source.company_id == new_company.id)) + ).scalar_one() + assert ( + new_source.id + != (await db_session.execute(select(Source).where(Source.company_id == company.id))) + .scalar_one() + .id + ) + + new_doc = ( + await db_session.execute( + select(SourceDocument).where(SourceDocument.company_id == new_company.id) + ) + ).scalar_one() + assert new_doc.source_id == new_source.id # remapped to the CLONED source, not the original + + new_run = ( + await db_session.execute( + select(MonitoringRun).where(MonitoringRun.company_id == new_company.id) + ) + ).scalar_one() + new_report = ( + await db_session.execute(select(Report).where(Report.company_id == new_company.id)) + ).scalar_one() + assert new_report.monitoring_run_id == new_run.id + + new_snapshot = ( + await db_session.execute(select(Snapshot).where(Snapshot.company_id == new_company.id)) + ).scalar_one() + assert new_snapshot.source_id == new_source.id + assert new_snapshot.monitoring_run_id == new_run.id + + new_change = ( + await db_session.execute( + select(DetectedChange).where(DetectedChange.company_id == new_company.id) + ) + ).scalar_one() + assert new_change.previous_snapshot_id == new_snapshot.id + assert new_change.current_snapshot_id == new_snapshot.id + assert new_change.source_id == new_source.id + assert new_change.monitoring_run_id == new_run.id + + new_enrichment = ( + await db_session.execute( + select(CompanyEnrichment).where(CompanyEnrichment.company_id == new_company.id) + ) + ).scalar_one() + assert new_enrichment.data == {"products": [{"name": "Widget"}]} + + new_key = ( + await db_session.execute(select(UserApiKey).where(UserApiKey.user_id == new_user.id)) + ).scalar_one() + assert new_key.provider == ApiKeyProvider.ANTHROPIC + assert new_key.encrypted_key == "encrypted-blob" # copied as-is, not re-encrypted + + # The original owner's data must be completely untouched. + original_source = ( + await db_session.execute(select(Source).where(Source.company_id == company.id)) + ).scalar_one() + assert original_source.id != new_source.id + + destination = (await NotificationDestinationRepository(db_session).list_for_user(new_user.id))[ + 0 + ] + assert destination.destination_value == email + assert destination.company_links[0].company_id == new_company.id + + +@pytest.mark.asyncio +async def test_apply_if_pending_without_make_admin_leaves_user_a_regular_account( + db_session, settings +): + owner = await _make_user(db_session) + await _build_full_company(db_session, owner) + email = f"invitee-{uuid.uuid4().hex[:10]}@example.com" + await provisioning_service.create_pending(db_session, email, owner.id, make_admin=False) + new_user = await _make_user(db_session, email=email) + + await provisioning_service.apply_if_pending(db_session, new_user) + + assert new_user.is_admin is False + + +@pytest.mark.asyncio +async def test_delete_pending_removes_an_unconsumed_entry(db_session, settings): + owner = await _make_user(db_session) + email = f"pending-{uuid.uuid4().hex[:10]}@example.com" + record = await provisioning_service.create_pending(db_session, email, owner.id, False) + + await provisioning_service.delete_pending(db_session, record.id) + + assert await PendingProvisioningRepository(db_session).get_by_email(email) is None + + +@pytest.mark.asyncio +async def test_delete_pending_404s_for_an_unknown_id(db_session, settings): + with pytest.raises(NotFoundError): + await provisioning_service.delete_pending(db_session, uuid.uuid4()) + + +def test_verify_email_applies_a_matching_pending_entry_end_to_end(client, settings): + """The verify-email endpoint, not registration, is what actually + triggers provisioning - see auth_service.verify_email's docstring.""" + owner_email = f"owner-{uuid.uuid4().hex[:10]}@example.com" + client.post( + "/api/v1/auth/register", + json={"email": owner_email, "password": "correct-horse-1", "display_name": "Owner"}, + ) + invitee_email = f"invitee-{uuid.uuid4().hex[:10]}@example.com" + + async def _setup() -> None: + async with get_sessionmaker()() as db: + owner = await UserRepository(db).get_by_email(owner_email) + await _build_full_company(db, owner) + await provisioning_service.create_pending(db, invitee_email, owner.id, make_admin=True) + + asyncio.run(_setup()) + + register_resp = client.post( + "/api/v1/auth/register", + json={"email": invitee_email, "password": "correct-horse-1", "display_name": "Invitee"}, + ) + assert register_resp.status_code == 201 + assert register_resp.json()["is_admin"] is False # not yet - only verifying triggers it + + code = "424242" + + async def _issue_code() -> None: + async with get_sessionmaker()() as db: + db.add( + EmailCode( + user_id=uuid.UUID(register_resp.json()["id"]), + purpose=EmailCodePurpose.VERIFY_EMAIL, + code_hash=hash_email_code(code), + expires_at=datetime.now(UTC) + timedelta(hours=1), + ) + ) + await db.commit() + + asyncio.run(_issue_code()) + + verify_resp = client.post( + "/api/v1/auth/verify-email", json={"email": invitee_email, "code": code} + ) + assert verify_resp.status_code == 204 + + login = client.post( + "/api/v1/auth/login", json={"email": invitee_email, "password": "correct-horse-1"} + ).json() + headers = {"Authorization": f"Bearer {login['access_token']}"} + + me = client.get("/api/v1/auth/me", headers=headers).json() + assert me["is_admin"] is True + + companies = client.get("/api/v1/companies", headers=headers).json() + assert len(companies) == 1 + assert companies[0]["name"] == "Acme Corp" + + +def test_pending_provisioning_admin_endpoints_require_admin(client): + headers = _register_and_login_plain(client) + resp = client.get("/api/v1/system/pending-provisioning", headers=headers) + assert resp.status_code == 403 + + +def _register_and_login_plain(client) -> dict[str, str]: + email = f"plain-{uuid.uuid4().hex[:10]}@example.com" + client.post( + "/api/v1/auth/register", + json={"email": email, "password": "correct-horse-1", "display_name": "Plain"}, + ) + tokens = client.post( + "/api/v1/auth/login", json={"email": email, "password": "correct-horse-1"} + ).json() + return {"Authorization": f"Bearer {tokens['access_token']}"}