241 lines
8.5 KiB
Python
241 lines
8.5 KiB
Python
"""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()
|