from __future__ import annotations from dataclasses import dataclass from uuid import uuid4 from sqlalchemy import select from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from app.security.database import SecurityDatabase from app.security.errors import ForbiddenError, UnauthorizedError from app.security.models import APIKey, APIKeyPrincipal, User, Workspace, WorkspaceMembership @dataclass(frozen=True, slots=True) class TenantPrincipal: """The persisted tenant authority resolved for an authenticated API key.""" workspace_id: str user_id: str membership_id: str membership_role: str class TenantService: """Owns authoritative tenant records and API-key principal bindings. API keys remain authentication credentials. They never become workspace IDs, and all tenant selection happens server-side through this service. """ def __init__(self, database: SecurityDatabase) -> None: self.database = database async def resolve_api_key(self, api_key_id: str) -> TenantPrincipal: async with self.database.session() as session: binding = await session.scalar( select(APIKeyPrincipal) .where(APIKeyPrincipal.api_key_id == api_key_id) .with_for_update() ) if binding is not None: return await self._validated_principal(session, binding) # Existing installs predate native tenancy. Provisioning a # one-time private owner/workspace binding preserves access while # ensuring the API-key ID itself can no longer act as a tenant. key = await session.get(APIKey, api_key_id) if key is None: raise UnauthorizedError user_id = str(uuid4()) workspace_id = str(uuid4()) membership_id = str(uuid4()) user = User( id=user_id, subject=f"legacy-api-key-principal:{api_key_id}", display_name=key.name, metadata_json={"provisioned_from": "api_key"}, ) workspace = Workspace( id=workspace_id, slug=f"legacy-{api_key_id}", name=f"{key.name} workspace", metadata_json={"provisioned_from": "api_key"}, ) membership = WorkspaceMembership( id=membership_id, workspace_id=workspace_id, user_id=user_id, role="owner", status="active", ) binding = APIKeyPrincipal( api_key_id=api_key_id, workspace_id=workspace_id, user_id=user_id, membership_id=membership_id, status="active", ) session.add_all([user, workspace]) await session.flush() session.add(membership) await session.flush() session.add(binding) try: await session.commit() except IntegrityError: # A concurrent first request won the unique binding race. await session.rollback() binding = await session.scalar( select(APIKeyPrincipal).where(APIKeyPrincipal.api_key_id == api_key_id) ) if binding is None: raise return await self._validated_principal(session, binding) return TenantPrincipal( workspace_id=workspace_id, user_id=user_id, membership_id=membership_id, membership_role=membership.role, ) async def bind_api_key( self, *, api_key_id: str, workspace_id: str, user_id: str ) -> TenantPrincipal: """Bind a newly created credential to an active persisted membership.""" async with self.database.session() as session: membership = await session.scalar( select(WorkspaceMembership).where( WorkspaceMembership.workspace_id == workspace_id, WorkspaceMembership.user_id == user_id, ) ) if membership is None or membership.status != "active": raise ForbiddenError key = await session.get(APIKey, api_key_id) if key is None: raise UnauthorizedError existing = await session.scalar( select(APIKeyPrincipal).where(APIKeyPrincipal.api_key_id == api_key_id) ) if existing is None: existing = APIKeyPrincipal( api_key_id=api_key_id, workspace_id=workspace_id, user_id=user_id, membership_id=membership.id, status="active", ) session.add(existing) await session.commit() return await self._validated_principal(session, existing) async def ensure_all_api_key_principals(self) -> None: """Provision legacy bindings before accepting requests after upgrade.""" async with self.database.session() as session: key_ids = list( ( await session.scalars( select(APIKey.id) .outerjoin(APIKeyPrincipal, APIKeyPrincipal.api_key_id == APIKey.id) .where(APIKeyPrincipal.id.is_(None)) ) ).all() ) for key_id in key_ids: await self.resolve_api_key(key_id) async def list_principals(self) -> list[tuple[str, TenantPrincipal]]: """Return server-side legacy-key mappings for one-time tenant adoption.""" async with self.database.session() as session: bindings = list((await session.scalars(select(APIKeyPrincipal))).all()) result: list[tuple[str, TenantPrincipal]] = [] for binding in bindings: try: result.append( (binding.api_key_id, await self._validated_principal(session, binding)) ) except ForbiddenError: # Disabled memberships must not migrate or retain access. continue return result @staticmethod async def _validated_principal( session: AsyncSession, binding: APIKeyPrincipal ) -> TenantPrincipal: membership = await session.get(WorkspaceMembership, binding.membership_id) user = await session.get(User, binding.user_id) workspace = await session.get(Workspace, binding.workspace_id) if ( binding.status != "active" or membership is None or membership.status != "active" or membership.workspace_id != binding.workspace_id or membership.user_id != binding.user_id or user is None or user.status != "active" or workspace is None or workspace.status != "active" ): raise ForbiddenError return TenantPrincipal( workspace_id=binding.workspace_id, user_id=binding.user_id, membership_id=membership.id, membership_role=membership.role, )