495 lines
18 KiB
Python
495 lines
18 KiB
Python
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)
|