"""Workspace service with tariff-gated limits for VoIdea.""" from typing import Any, Optional from uuid import UUID from sqlalchemy import select, func from sqlalchemy.ext.asyncio import AsyncSession from app.models.workspace import Workspace, WorkspaceMembership from app.models.user import User from app.services.tariff_service import TariffService class WorkspaceError(Exception): pass class WorkspaceLimitError(WorkspaceError): pass class WorkspaceService: def __init__(self, db: AsyncSession): self.db = db # ── Tariff gates ── async def _get_tariff_features(self, user_id: UUID) -> dict[str, Any]: """Get tariff plan features for a user.""" tariff_svc = TariffService(self.db) sub = await tariff_svc.get_user_subscription(user_id) if not sub: return {} plan = await tariff_svc.get_plan_by_id(sub.plan_id) return plan.features if plan and plan.features else {} async def _check_team_workspace_limit(self, user_id: UUID) -> None: """Check if user can create a team workspace.""" features = await self._get_tariff_features(user_id) max_team = features.get("max_team_workspaces", 0) if max_team <= 0: raise WorkspaceLimitError( "Your tariff plan does not allow team workspaces" ) result = await self.db.execute( select(func.count(Workspace.id)) .where( Workspace.owner_id == user_id, Workspace.type == "team", ) ) current_team_count = result.scalar() or 0 if current_team_count >= max_team: raise WorkspaceLimitError( f"Team workspace limit reached ({max_team}). " "Upgrade your tariff to create more." ) async def _check_member_limit(self, workspace_id: UUID) -> None: """Check if workspace can accept more members.""" ws = await self.db.get(Workspace, workspace_id) if not ws: raise WorkspaceError("Workspace not found") features = await self._get_tariff_features(ws.owner_id) max_members = features.get("max_members_per_workspace", 0) result = await self.db.execute( select(func.count(WorkspaceMembership.id)) .where(WorkspaceMembership.workspace_id == workspace_id) ) current_count = result.scalar() or 0 if max_members > 0 and current_count >= max_members: raise WorkspaceLimitError( f"Member limit reached ({max_members}). " "Upgrade your tariff to add more members." ) async def _check_feature_gate(self, user_id: UUID, feature: str) -> None: """Check if a specific workspace feature is enabled for the user's tariff.""" features = await self._get_tariff_features(user_id) if not features.get(feature, False): raise WorkspaceLimitError( f"Feature '{feature}' is not available on your tariff plan" ) # ── CRUD ── async def create_personal(self, user_id: UUID, name: str = "Personal") -> Workspace: """Create or get existing personal workspace for user.""" result = await self.db.execute( select(Workspace).where( Workspace.owner_id == user_id, Workspace.type == "personal", ) ) existing = result.scalar_one_or_none() if existing: return existing ws = Workspace(name=name, type="personal", owner_id=user_id) self.db.add(ws) await self.db.flush() membership = WorkspaceMembership( workspace_id=ws.id, user_id=user_id, role="owner" ) self.db.add(membership) await self.db.commit() await self.db.refresh(ws) return ws async def create_team( self, name: str, owner_id: UUID, description: str = "" ) -> Workspace: """Create a new team workspace (tariff-gated).""" await self._check_team_workspace_limit(owner_id) ws = Workspace(name=name, type="team", owner_id=owner_id, description=description) self.db.add(ws) await self.db.flush() membership = WorkspaceMembership( workspace_id=ws.id, user_id=owner_id, role="owner" ) self.db.add(membership) await self.db.commit() await self.db.refresh(ws) return ws async def get_by_id(self, workspace_id: UUID) -> Optional[Workspace]: return await self.db.get(Workspace, workspace_id) async def list_user_workspaces(self, user_id: UUID) -> list[Workspace]: """List workspaces where user is a member.""" result = await self.db.execute( select(Workspace) .join(WorkspaceMembership) .where(WorkspaceMembership.user_id == user_id) .order_by(Workspace.type, Workspace.name) ) return list(result.scalars().all()) async def update_workspace( self, workspace_id: UUID, user_id: UUID, updates: dict[str, Any] ) -> Optional[Workspace]: ws = await self.db.get(Workspace, workspace_id) if not ws: return None if ws.owner_id != user_id: raise WorkspaceError("Only the owner can update workspace") for key, value in updates.items(): if hasattr(ws, key) and key not in ("id", "owner_id", "type", "created_at"): setattr(ws, key, value) await self.db.commit() await self.db.refresh(ws) return ws async def delete_workspace(self, workspace_id: UUID, user_id: UUID) -> bool: ws = await self.db.get(Workspace, workspace_id) if not ws: return False if ws.owner_id != user_id: raise WorkspaceError("Only the owner can delete workspace") if ws.type == "personal": raise WorkspaceError("Cannot delete personal workspace") await self.db.delete(ws) await self.db.commit() return True # ── Members ── async def add_member( self, workspace_id: UUID, user_id: UUID, role: str = "member" ) -> WorkspaceMembership: """Add a user to a workspace (tariff-gated for member count).""" await self._check_member_limit(workspace_id) existing = await self.db.execute( select(WorkspaceMembership).where( WorkspaceMembership.workspace_id == workspace_id, WorkspaceMembership.user_id == user_id, ) ) if existing.scalar_one_or_none(): raise WorkspaceError("User is already a member of this workspace") membership = WorkspaceMembership( workspace_id=workspace_id, user_id=user_id, role=role ) self.db.add(membership) await self.db.commit() await self.db.refresh(membership) return membership async def remove_member(self, workspace_id: UUID, user_id: UUID) -> bool: result = await self.db.execute( select(WorkspaceMembership).where( WorkspaceMembership.workspace_id == workspace_id, WorkspaceMembership.user_id == user_id, ) ) membership = result.scalar_one_or_none() if not membership: return False await self.db.delete(membership) await self.db.commit() return True async def list_members(self, workspace_id: UUID) -> list[dict[str, Any]]: result = await self.db.execute( select(WorkspaceMembership, User) .join(User, WorkspaceMembership.user_id == User.id) .where(WorkspaceMembership.workspace_id == workspace_id) ) rows = result.all() return [ { "id": str(m.WorkspaceMembership.id), "user_id": str(m.User.id), "workspace_id": str(m.WorkspaceMembership.workspace_id), "role": m.WorkspaceMembership.role, "email": m.User.email, "display_name": m.User.display_name, "joined_at": m.WorkspaceMembership.created_at, } for m in rows ] async def get_membership( self, workspace_id: UUID, user_id: UUID ) -> Optional[WorkspaceMembership]: result = await self.db.execute( select(WorkspaceMembership).where( WorkspaceMembership.workspace_id == workspace_id, WorkspaceMembership.user_id == user_id, ) ) return result.scalar_one_or_none()