import hmac import hashlib import logging from typing import Dict, Any, Optional, Tuple from sqlalchemy.orm import Session import app.db.all_models # noqa: F401 from app.core.settings import settings from app.modules.auth.models.user_model import User from app.modules.tenant.models.tenant_model import Tenant from app.modules.auth.models.role_model import Role from app.modules.billing.models.plan_model import Plan from app.modules.auth.models.saas_models import SaaSUserMapping, SaaSTenantMapping, SaaSRoleMapping from app.modules.auth.models.role_access_model import RoleAccess from app.modules.auth.models.access_model import Access from app.core.security import create_access_token import uuid logger = logging.getLogger(__name__) class SaaSService: @staticmethod def verify_signature(signature: str, canonical_string: str) -> bool: """ Verify HMAC signature from SaaS. """ if not settings.SAAS_TRUST_SECRET: logger.error("SAAS_TRUST_SECRET is not configured") return False expected_signature = hmac.new( settings.SAAS_TRUST_SECRET.encode(), canonical_string.encode(), hashlib.sha256 ).hexdigest() if not hmac.compare_digest(signature, expected_signature): logger.warning(f"Signature mismatch. Expected: {expected_signature}, Got: {signature}") return False return True @staticmethod def ensure_saas_user(db: Session, saas_data: Dict[str, Any]) -> Tuple[User, SaaSUserMapping]: from sqlalchemy import text try: db.execute(text("SELECT set_config('docqube.bypass', 'on', true)")) except Exception: pass saas_user_id = str(saas_data.get("id")) email = saas_data.get("email") name = saas_data.get("name") saas_tenant_id = saas_data.get("company_id") subscription = saas_data.get("metadata", {}).get("subscription", {}) or {} tenant = None if saas_tenant_id: tenant_name = saas_data.get("tenant_name") or "SaaS Tenant" tenant = SaaSService.ensure_saas_tenant( db, saas_tenant_id, tenant_name, max_users_allowed=subscription.get("max_users_allowed"), is_active=bool(subscription.get("is_active", True)), ) from sqlalchemy import func role_name = saas_data.get("role") role_id_str = saas_data.get("role_id") target_role = None if role_id_str: try: target_role = db.query(Role).filter(Role.id == uuid.UUID(role_id_str)).first() except Exception: pass if not target_role and role_name and tenant: target_role = db.query(Role).filter( Role.tenant_id == tenant.id, func.lower(Role.name) == role_name.strip().lower() ).first() if not target_role and role_name and tenant: try: target_role = Role( id=uuid.UUID(role_id_str) if role_id_str else uuid.uuid4(), name=role_name.strip(), tenant_id=tenant.id, is_default=False, is_system=False ) db.add(target_role) db.flush() except Exception as e: logger.warning(f"Could not create role {role_name} for tenant {tenant.id}: {e}") permissions = saas_data.get("metadata", {}).get("permissions", []) if target_role and permissions: try: access_rows = db.query(Access).filter(Access.access_code.in_(permissions)).all() existing_access_ids = { ra.access_id for ra in db.query(RoleAccess).filter(RoleAccess.role_id == target_role.id).all() } new_access_ids = {a.id for a in access_rows} to_delete = existing_access_ids - new_access_ids if to_delete: db.query(RoleAccess).filter( RoleAccess.role_id == target_role.id, RoleAccess.access_id.in_(to_delete) ).delete(synchronize_session=False) to_add = new_access_ids - existing_access_ids for aid in to_add: db.add(RoleAccess(role_id=target_role.id, access_id=aid)) db.flush() except Exception as e: logger.warning(f"Could not sync role accesses for {target_role.name}: {e}") mapping = db.query(SaaSUserMapping).filter(SaaSUserMapping.saas_user_id == saas_user_id).first() if mapping: user = mapping.user if tenant and user and user.tenant_id != tenant.id: user.tenant_id = tenant.id if target_role and user and user.role_id != target_role.id: user.role_id = target_role.id if user: db.flush() db.refresh(user) return mapping.user, mapping user = db.query(User).filter(User.email == email).first() if not user: logger.info(f"Creating new local user for SaaS user {email}") user = User( email=email, name=name or email.split('@')[0], tenant_id=tenant.id if tenant else None, role_id=target_role.id if target_role else None, password_hash="saas_authenticated", is_active=True ) db.add(user) db.flush() try: from app.modules.drive.services.drive_service import DriveService drive_service = DriveService(db) drive_service.get_or_create_root_folder(user) except Exception as e: logger.warning(f"Could not init root folder for SaaS user {email}: {e}") try: from app.modules.storage.models.storage_model import UserStorageUsage if not db.query(UserStorageUsage).filter(UserStorageUsage.user_id == user.id).first(): storage_usage = UserStorageUsage(user_id=user.id) db.add(storage_usage) db.flush() except Exception as e: logger.warning(f"Could not init storage usage for SaaS user {email}: {e}") else: if target_role and user.role_id != target_role.id: user.role_id = target_role.id db.flush() db.refresh(user) mapping = SaaSUserMapping( saas_user_id=saas_user_id, docqube_user_id=user.id, metadata_=saas_data.get("metadata") ) db.add(mapping) db.flush() return user, mapping @staticmethod def ensure_saas_tenant( db: Session, saas_tenant_id: str, name: str, max_users_allowed: Optional[int] = None, is_active: Optional[bool] = None, ) -> Tenant: """ Ensures a local tenant exists for the given SaaS tenant ID. """ mapping = db.query(SaaSTenantMapping).filter(SaaSTenantMapping.saas_tenant_id == saas_tenant_id).first() if mapping: tenant = mapping.tenant changed = False if name and tenant.name != name: tenant.name = name changed = True if tenant.max_users_allowed != max_users_allowed: tenant.max_users_allowed = max_users_allowed changed = True if is_active is not None and tenant.is_active != is_active: tenant.is_active = is_active changed = True if changed: db.flush() db.refresh(tenant) return tenant slug = f"saas-{saas_tenant_id[:8]}" tenant = db.query(Tenant).filter(Tenant.slug == slug).first() if not tenant: logger.info(f"Creating new local tenant for SaaS tenant {saas_tenant_id}") tenant = Tenant( name=name, slug=slug, is_active=True if is_active is None else is_active, max_users_allowed=max_users_allowed, ) db.add(tenant) db.flush() mapping = SaaSTenantMapping( saas_tenant_id=saas_tenant_id, docqube_tenant_id=tenant.id ) db.add(mapping) db.flush() return tenant @staticmethod def provision_user(db: Session, data: dict): saas_user_id = data.get("user_id") or data.get("id") email = data.get("email") first_name = data.get("first_name") last_name = data.get("last_name") phone = data.get("phone_number") saas_tenant_id = data.get("tenant_id") saas_role_id = data.get("role_id") status = data.get("status", "active") if not saas_user_id or not email: logger.error("Missing user_id or email in provision payload") return tenant_id = None if saas_tenant_id: tenant_mapping = db.query(SaaSTenantMapping).filter(SaaSTenantMapping.saas_tenant_id == saas_tenant_id).first() if tenant_mapping: tenant_id = tenant_mapping.docqube_tenant_id role_id = None if saas_role_id: role_mapping = db.query(SaaSRoleMapping).filter(SaaSRoleMapping.saas_role_id == saas_role_id).first() if role_mapping: role_id = role_mapping.docqube_role_id mapping = db.query(SaaSUserMapping).filter(SaaSUserMapping.saas_user_id == saas_user_id).first() name = first_name or "" if last_name: name += f" {last_name}" name = name.strip() or email.split('@')[0] if mapping: user = mapping.user if user: user.email = email user.name = name user.tenant_id = tenant_id or user.tenant_id user.role_id = role_id or user.role_id user.is_active = (status == "active") db.flush() return user, mapping user = db.query(User).filter(User.email == email).first() if not user: user = User( email=email, name=name, password_hash="SAAS_MANAGED_ACCOUNT_DO_NOT_USE_PASSWORD", is_active=(status == "active"), tenant_id=tenant_id, role_id=role_id ) db.add(user) db.flush() else: user.tenant_id = tenant_id or user.tenant_id user.role_id = role_id or user.role_id user.is_active = (status == "active") db.flush() mapping = SaaSUserMapping( saas_user_id=saas_user_id, docqube_user_id=user.id, metadata_=data ) db.add(mapping) db.flush() return user, mapping @staticmethod def deprovision_user(db: Session, data: dict): saas_user_id = data.get("user_id") or data.get("id") if not saas_user_id: return mapping = db.query(SaaSUserMapping).filter(SaaSUserMapping.saas_user_id == saas_user_id).first() if mapping and mapping.user: mapping.user.is_active = False db.flush() @staticmethod def provision_role(db: Session, data: dict): saas_role_id = data.get("role_id") or data.get("id") name = data.get("role_name") or data.get("name") description = data.get("description") saas_tenant_id = data.get("tenant_id") targets = data.get("targets", []) if not saas_role_id or not name: logger.error("Missing role_id or name in role provision payload") return tenant_id = None if saas_tenant_id: tenant_mapping = db.query(SaaSTenantMapping).filter(SaaSTenantMapping.saas_tenant_id == saas_tenant_id).first() if tenant_mapping: tenant_id = tenant_mapping.docqube_tenant_id # Extract docqube permissions from targets docqube_permissions = [] for target in targets: # Module ID might be a UUID string or 'docqube'. Fallback to processing all permissions for now if module check is tricky, # but safer to check if there's any permission array. permissions = target.get("permissions") if isinstance(permissions, list): docqube_permissions.extend(permissions) mapping = db.query(SaaSRoleMapping).filter(SaaSRoleMapping.saas_role_id == saas_role_id).first() if mapping: role = mapping.role if role: role.name = name role.description = description role.tenant_id = tenant_id # Sync Role Access if docqube_permissions: # Remove old access db.query(RoleAccess).filter(RoleAccess.role_id == role.id).delete() # Assign new access access_records = db.query(Access).filter(Access.access_code.in_(docqube_permissions)).all() for access in access_records: db.add(RoleAccess(role_id=role.id, access_id=access.id)) db.flush() return role, mapping role = Role( name=name, description=description, tenant_id=tenant_id ) db.add(role) db.flush() # Sync Role Access if docqube_permissions: access_records = db.query(Access).filter(Access.access_code.in_(docqube_permissions)).all() for access in access_records: db.add(RoleAccess(role_id=role.id, access_id=access.id)) db.flush() mapping = SaaSRoleMapping( saas_role_id=saas_role_id, docqube_role_id=role.id, metadata_=data ) db.add(mapping) db.flush() return role, mapping @staticmethod def deprovision_role(db: Session, data: dict): saas_role_id = data.get("role_id") or data.get("id") if not saas_role_id: return mapping = db.query(SaaSRoleMapping).filter(SaaSRoleMapping.saas_role_id == saas_role_id).first() if mapping: if mapping.role: db.delete(mapping.role) db.delete(mapping) db.flush() # ------------------------------------------------------------------ # # Plan sync — SaaS is the catalogue owner, DocQube is a consumer. # # ------------------------------------------------------------------ # @staticmethod def _apply_plan_fields(plan, data: dict) -> None: """Copy SaaS webhook payload fields onto a DocQube Plan row.""" plan.saas_plan_id = str(data["saas_plan_id"]) plan.code = data.get("code") or str(data["saas_plan_id"]) plan.name = data.get("name") or plan.code plan.description = data.get("description") plan.price_amount = data.get("price_amount") plan.currency = data.get("currency") or "USD" plan.interval = data.get("interval") or "monthly" plan.is_public = bool(data.get("is_public", True)) plan.grace_period_days = int(data.get("grace_period_days") or 0) if getattr(plan, "sort_order", None) is None: plan.sort_order = int(data.get("sort_order") or 0) if getattr(plan, "notify_days_before_expiry", None) is None: plan.notify_days_before_expiry = int(data.get("notify_days_before_expiry") or 7) if getattr(plan, "is_popular", None) is None: plan.is_popular = bool(data.get("is_popular", False)) @staticmethod def ensure_saas_plan(db: Session, data: dict): """ Upsert a Plan row from a PLAN_PROVISION_REQUESTED webhook event. Lookup order: 1. by saas_plan_id (most reliable — survives plan renames) 2. by code (fallback for first-time sync) Returns the Plan row. """ from app.modules.billing.models.plan_model import Plan saas_plan_id = str(data.get("saas_plan_id", "")) code = data.get("code") or saas_plan_id plan = ( db.query(Plan).filter(Plan.saas_plan_id == saas_plan_id).first() if saas_plan_id else None ) if plan is None and code: plan = db.query(Plan).filter(Plan.code == code).first() if plan is None: plan = Plan() db.add(plan) SaaSService._apply_plan_fields(plan, data) db.flush() logger.info("Upserted DocQube plan %r from SaaS plan %s", plan.code, saas_plan_id) return plan @staticmethod def update_saas_plan(db: Session, data: dict): """Update an existing Plan from a PLAN_UPDATED webhook event.""" return SaaSService.ensure_saas_plan(db, data) @staticmethod def deprovision_plan(db: Session, data: dict): """ Handle PLAN_DEPROVISION_REQUESTED from SaaS. We never hard-delete here — a plan may still be referenced by TenantSubscription rows. Instead we mark it non-public so it disappears from the pricing page and cannot be assigned to new tenants. A superadmin can clean it up manually once it's empty. """ from app.modules.billing.models.plan_model import Plan, TenantSubscription saas_plan_id = str(data.get("saas_plan_id", "")) code = data.get("code") or saas_plan_id plan = ( db.query(Plan).filter(Plan.saas_plan_id == saas_plan_id).first() if saas_plan_id else None ) if plan is None and code: plan = db.query(Plan).filter(Plan.code == code).first() if plan is None: logger.warning("PLAN_DEPROVISION_REQUESTED for unknown plan %s — ignoring", saas_plan_id) return in_use = db.query(TenantSubscription).filter(TenantSubscription.plan_id == plan.id).count() if in_use: # Can't delete — just hide it from new assignments plan.is_public = False db.flush() logger.info("Plan %r hidden (still used by %d tenant(s))", plan.code, in_use) else: db.delete(plan) db.flush() logger.info("Plan %r deleted", plan.code)