"""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']}"}