feat: add workspace tenant foundations
This commit is contained in:
@@ -16,6 +16,7 @@ from app.models.dns_cache import DNSRecordChange, DNSRecordSnapshot
|
||||
from app.models.domain import Domain
|
||||
from app.models.setting import Setting
|
||||
from app.services.dns_resolver import CloudflareDNSProvider, extract_dmarc_policy
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
PROVIDER_NAME = "cloudflare"
|
||||
|
||||
@@ -90,6 +91,7 @@ async def import_cloudflare_domains(
|
||||
) -> Dict[str, Any]:
|
||||
"""Create Domain rows for Cloudflare zones, returning imported and existing names."""
|
||||
zones = await discover_cloudflare_zones(db)
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db)
|
||||
requested = {domain.strip().lower() for domain in requested_domains or [] if domain.strip()}
|
||||
imported: List[str] = []
|
||||
existing: List[str] = []
|
||||
@@ -100,9 +102,13 @@ async def import_cloudflare_domains(
|
||||
if requested and name not in requested:
|
||||
skipped.append(name)
|
||||
continue
|
||||
domain = db.query(Domain).filter(Domain.name == name).first()
|
||||
domain = (
|
||||
db.query(Domain)
|
||||
.filter(Domain.name == name, Domain.workspace_id == workspace.id)
|
||||
.first()
|
||||
)
|
||||
if domain is None:
|
||||
db.add(Domain(name=name, active=True, verified=True))
|
||||
db.add(Domain(name=name, active=True, verified=True, workspace_id=workspace.id))
|
||||
imported.append(name)
|
||||
else:
|
||||
existing.append(name)
|
||||
|
||||
@@ -7,6 +7,7 @@ from sqlalchemy.orm import Session
|
||||
from app.models.domain import Domain
|
||||
from app.models.report import ForensicReport
|
||||
from app.services.forensic_redaction import ForensicRedactionPolicy, redact_forensic_value
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
from app.utils.domain_validator import DomainValidationError, validate_domain
|
||||
|
||||
|
||||
@@ -28,9 +29,14 @@ def _domain_for_report(db: Session, domain_name: Optional[str]) -> Optional[Doma
|
||||
if not is_valid and error_code != DomainValidationError.DNS_RESOLUTION_FAILED:
|
||||
return None
|
||||
|
||||
domain = db.query(Domain).filter(Domain.name == normalized).first()
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
domain = (
|
||||
db.query(Domain)
|
||||
.filter(Domain.name == normalized, Domain.workspace_id == workspace.id)
|
||||
.first()
|
||||
)
|
||||
if domain is None:
|
||||
domain = Domain(name=normalized)
|
||||
domain = Domain(name=normalized, workspace_id=workspace.id)
|
||||
db.add(domain)
|
||||
db.flush()
|
||||
return domain
|
||||
|
||||
@@ -7,6 +7,7 @@ from sqlalchemy.orm import Session, selectinload
|
||||
from app.models.domain import Domain
|
||||
from app.models.report import DMARCReport, ReportRecord
|
||||
from app.services.report_store import ReportStore
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
|
||||
|
||||
def _parse_timestamp(value: Any) -> int:
|
||||
@@ -98,10 +99,15 @@ def save_parsed_report(db: Session, report: Dict[str, Any]) -> tuple[DMARCReport
|
||||
domain_name = report.get("domain") or "unknown"
|
||||
report_id = report.get("report_id") or ""
|
||||
policy = _policy_parts(report)
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
|
||||
domain = db.query(Domain).filter(Domain.name == domain_name).first()
|
||||
domain = (
|
||||
db.query(Domain)
|
||||
.filter(Domain.name == domain_name, Domain.workspace_id == workspace.id)
|
||||
.first()
|
||||
)
|
||||
if domain is None:
|
||||
domain = Domain(name=domain_name, dmarc_policy=policy["p"])
|
||||
domain = Domain(name=domain_name, dmarc_policy=policy["p"], workspace_id=workspace.id)
|
||||
db.add(domain)
|
||||
db.flush()
|
||||
elif policy.get("p"):
|
||||
|
||||
@@ -12,9 +12,9 @@ from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from app.models.domain import Domain
|
||||
from app.models.report import TLSReport, TLSReportFailure
|
||||
from app.services.workspaces import assign_default_workspace_to_unscoped_rows
|
||||
from app.utils.domain_validator import DomainValidationError, validate_domain
|
||||
|
||||
|
||||
TLS_REPORT_PRIVACY_CONTROLS = {
|
||||
"retention": (
|
||||
"TLS reports store aggregate session counts, reporting organization metadata, "
|
||||
@@ -69,9 +69,14 @@ def _domain_for_report(db: Session, domain_name: Optional[str]) -> Optional[Doma
|
||||
if not is_valid and error_code != DomainValidationError.DNS_RESOLUTION_FAILED:
|
||||
return None
|
||||
|
||||
domain = db.query(Domain).filter(Domain.name == normalized).first()
|
||||
workspace = assign_default_workspace_to_unscoped_rows(db, commit=False)
|
||||
domain = (
|
||||
db.query(Domain)
|
||||
.filter(Domain.name == normalized, Domain.workspace_id == workspace.id)
|
||||
.first()
|
||||
)
|
||||
if domain is None:
|
||||
domain = Domain(name=normalized)
|
||||
domain = Domain(name=normalized, workspace_id=workspace.id)
|
||||
db.add(domain)
|
||||
db.flush()
|
||||
return domain
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
"""Workspace/tenant helpers for MSP mode foundations."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Query, Session
|
||||
|
||||
from app.models.domain import Domain
|
||||
from app.models.mail_source import MailSource
|
||||
from app.models.user import User
|
||||
from app.models.workspace import Workspace
|
||||
|
||||
DEFAULT_WORKSPACE_SLUG = "default"
|
||||
DEFAULT_WORKSPACE_NAME = "Default Workspace"
|
||||
|
||||
|
||||
def normalize_workspace_slug(value: str) -> str:
|
||||
"""Normalize a workspace slug for stable lookups."""
|
||||
slug = (value or "").strip().lower()
|
||||
cleaned = []
|
||||
previous_dash = False
|
||||
for char in slug:
|
||||
if char.isalnum():
|
||||
cleaned.append(char)
|
||||
previous_dash = False
|
||||
elif not previous_dash:
|
||||
cleaned.append("-")
|
||||
previous_dash = True
|
||||
return "".join(cleaned).strip("-")
|
||||
|
||||
|
||||
def get_or_create_default_workspace(db: Session, *, commit: bool = True) -> Workspace:
|
||||
"""Return the single-tenant default workspace, creating it when needed."""
|
||||
workspace = db.query(Workspace).filter(Workspace.slug == DEFAULT_WORKSPACE_SLUG).first()
|
||||
if workspace:
|
||||
return workspace
|
||||
|
||||
workspace = Workspace(
|
||||
slug=DEFAULT_WORKSPACE_SLUG,
|
||||
name=DEFAULT_WORKSPACE_NAME,
|
||||
description="Automatically created for existing single-tenant installs.",
|
||||
active=True,
|
||||
)
|
||||
db.add(workspace)
|
||||
if commit:
|
||||
db.commit()
|
||||
db.refresh(workspace)
|
||||
else:
|
||||
db.flush()
|
||||
return workspace
|
||||
|
||||
|
||||
def assign_default_workspace_to_unscoped_rows(
|
||||
db: Session,
|
||||
*,
|
||||
commit: bool = True,
|
||||
) -> Workspace:
|
||||
"""Attach legacy unscoped rows to the default workspace."""
|
||||
workspace = get_or_create_default_workspace(db, commit=commit)
|
||||
for model in (Domain, MailSource, User):
|
||||
db.query(model).filter(model.workspace_id.is_(None)).update(
|
||||
{model.workspace_id: workspace.id},
|
||||
synchronize_session=False,
|
||||
)
|
||||
if commit:
|
||||
db.commit()
|
||||
else:
|
||||
db.flush()
|
||||
return workspace
|
||||
|
||||
|
||||
def resolve_workspace(
|
||||
db: Session,
|
||||
*,
|
||||
workspace_id: Optional[int] = None,
|
||||
slug: Optional[str] = None,
|
||||
) -> Workspace:
|
||||
"""Resolve a workspace, defaulting to the single-tenant workspace."""
|
||||
if workspace_id is not None:
|
||||
workspace = (
|
||||
db.query(Workspace)
|
||||
.filter(Workspace.id == workspace_id, Workspace.active.is_(True))
|
||||
.first()
|
||||
)
|
||||
if workspace:
|
||||
return workspace
|
||||
raise ValueError("Workspace not found")
|
||||
|
||||
if slug:
|
||||
normalized = normalize_workspace_slug(slug)
|
||||
workspace = (
|
||||
db.query(Workspace)
|
||||
.filter(Workspace.slug == normalized, Workspace.active.is_(True))
|
||||
.first()
|
||||
)
|
||||
if workspace:
|
||||
return workspace
|
||||
raise ValueError("Workspace not found")
|
||||
|
||||
return assign_default_workspace_to_unscoped_rows(db)
|
||||
|
||||
|
||||
def workspace_domain_query(db: Session, workspace: Workspace) -> Query:
|
||||
"""Return the default scoped domain query for a workspace."""
|
||||
return db.query(Domain).filter(Domain.workspace_id == workspace.id)
|
||||
|
||||
|
||||
def workspace_mail_source_query(db: Session, workspace: Workspace) -> Query:
|
||||
"""Return the default scoped mail-source query for a workspace."""
|
||||
return db.query(MailSource).filter(MailSource.workspace_id == workspace.id)
|
||||
|
||||
|
||||
def workspace_user_query(db: Session, workspace: Workspace) -> Query:
|
||||
"""Return the default scoped user query for a workspace."""
|
||||
return db.query(User).filter(User.workspace_id == workspace.id)
|
||||
Reference in New Issue
Block a user