from __future__ import annotations from dataclasses import dataclass from datetime import date, datetime, timedelta import hashlib import re import secrets from typing import Any from sqlalchemy import select from sqlalchemy.orm import Session from app.core.security.jwt_tokens import utcnow from app.core.security.passwords import hash_password from app.core.settings import get_settings from app.modules.core.iam.models import User from app.modules.core.iam.password_flows_models import InviteToken from app.modules.core.rbac.models import Role, UserRole from app.modules.core.tenancy.models import Branch, FinancialYear, Tenant from app.modules.core.tenancy.settings_models import BranchSettings from app.modules.employees.models import Employee class FirmWizardError(ValueError): """Raised when the firm creation wizard receives invalid data.""" @dataclass(slots=True) class FirmWizardResult: tenant: Tenant branch: Branch firm_admin: User financial_year: FinancialYear | None invite_url: str invite_token: str def normalize_code(value: str, *, upper: bool = True) -> str: value = (value or "").strip() value = re.sub(r"\s+", "_", value) value = re.sub(r"[^A-Za-z0-9_\-]", "", value) return value.upper() if upper else value def clean_text(value: str | None) -> str: return (value or "").strip() def parse_bool(value: Any) -> bool: if isinstance(value, bool): return value return str(value or "").strip().lower() in {"1", "true", "yes", "on", "y"} def parse_int(value: Any, default: int) -> int: try: return int(str(value).strip()) except Exception: return default def parse_iso_date(value: str | None) -> date | None: value = clean_text(value) if not value: return None return date.fromisoformat(value) def default_ay_from_fy(year_code: str) -> str: year_code = clean_text(year_code) try: start_year = int(year_code.split("-", 1)[0]) except Exception: return "" ay_start = start_year + 1 return f"{ay_start}-{str(ay_start + 1)[-2:]}" def default_dates_from_fy(year_code: str) -> tuple[date | None, date | None]: year_code = clean_text(year_code) try: start_year = int(year_code.split("-", 1)[0]) except Exception: return None, None return date(start_year, 4, 1), date(start_year + 1, 3, 31) def public_invite_url(invite_token: str) -> str: base = (get_settings().ERP_PUBLIC_BASE_URL or "").strip().rstrip("/") or "http://localhost:8000" return f"{base}/invite/accept?token={invite_token}" def build_firm_wizard_payload(form: dict[str, Any]) -> dict[str, Any]: """Return a normalized payload used by preview and confirm pages.""" tenant_code = normalize_code(str(form.get("tenant_code") or form.get("firm_code") or "")) branch_code = normalize_code(str(form.get("branch_code") or "HO")) admin_email = clean_text(str(form.get("admin_email") or "")).lower() fy_enabled = parse_bool(form.get("create_financial_year")) fy_code = clean_text(str(form.get("fy_year_code") or "")) ay_code = clean_text(str(form.get("fy_assessment_year") or "")) fy_start_date = clean_text(str(form.get("fy_start_date") or "")) fy_end_date = clean_text(str(form.get("fy_end_date") or "")) if fy_enabled and fy_code: if not ay_code: ay_code = default_ay_from_fy(fy_code) if not fy_start_date or not fy_end_date: start, end = default_dates_from_fy(fy_code) fy_start_date = fy_start_date or (start.isoformat() if start else "") fy_end_date = fy_end_date or (end.isoformat() if end else "") return { "tenant_code": tenant_code, "tenant_name": clean_text(str(form.get("tenant_name") or form.get("firm_name") or "")), "firm_type": clean_text(str(form.get("firm_type") or "partnership")) or "partnership", "default_timezone": clean_text(str(form.get("default_timezone") or "Asia/Kolkata")) or "Asia/Kolkata", "default_session_duration_minutes": parse_int(form.get("default_session_duration_minutes"), 480), "default_otp_required_roles_csv": clean_text(str(form.get("default_otp_required_roles_csv") or "Partner,System Admin")) or "Partner,System Admin", "default_storage_mode": clean_text(str(form.get("default_storage_mode") or "local_only")) or "local_only", "branch_code": branch_code, "branch_name": clean_text(str(form.get("branch_name") or "Head Office")) or "Head Office", "branch_timezone": clean_text(str(form.get("branch_timezone") or form.get("default_timezone") or "Asia/Kolkata")) or "Asia/Kolkata", "branch_address_line1": clean_text(str(form.get("branch_address_line1") or "")), "branch_address_line2": clean_text(str(form.get("branch_address_line2") or "")), "branch_city": clean_text(str(form.get("branch_city") or "")), "branch_state": clean_text(str(form.get("branch_state") or "")), "branch_pin_code": clean_text(str(form.get("branch_pin_code") or "")), "branch_gstin": clean_text(str(form.get("branch_gstin") or "")).upper(), "branch_pan": clean_text(str(form.get("branch_pan") or "")).upper(), "admin_full_name": clean_text(str(form.get("admin_full_name") or "")), "admin_email": admin_email, "admin_mobile": clean_text(str(form.get("admin_mobile") or "")), "admin_designation": clean_text(str(form.get("admin_designation") or "Firm Admin")) or "Firm Admin", "create_financial_year": fy_enabled, "fy_year_code": fy_code, "fy_assessment_year": ay_code, "fy_start_date": fy_start_date, "fy_end_date": fy_end_date, "fy_is_current": parse_bool(form.get("fy_is_current")) if fy_enabled else False, } def validate_firm_wizard_payload(db: Session, payload: dict[str, Any]) -> list[str]: errors: list[str] = [] if not payload["tenant_code"]: errors.append("Firm code is required.") if not payload["tenant_name"]: errors.append("Firm name is required.") if not payload["branch_code"]: errors.append("Primary branch code is required.") if not payload["branch_name"]: errors.append("Primary branch name is required.") if not payload["admin_full_name"]: errors.append("Primary Firm Admin full name is required.") if not payload["admin_email"]: errors.append("Primary Firm Admin email is required.") elif "@" not in payload["admin_email"]: errors.append("Primary Firm Admin email is invalid.") if payload["default_session_duration_minutes"] < 15: errors.append("Session duration must be at least 15 minutes.") allowed_firm_types = {"partnership", "proprietorship", "individual"} if payload["firm_type"] not in allowed_firm_types: errors.append("Invalid firm type.") allowed_storage_modes = {"local_only", "cloud_only", "hybrid"} if payload["default_storage_mode"] not in allowed_storage_modes: errors.append("Invalid default storage mode.") if payload["tenant_code"]: existing_tenant = db.execute(select(Tenant).where(Tenant.code == payload["tenant_code"])).scalar_one_or_none() if existing_tenant: errors.append("Firm code already exists.") if payload["admin_email"]: existing_user = db.execute(select(User).where(User.email == payload["admin_email"])).scalar_one_or_none() if existing_user: errors.append("Primary Firm Admin email already exists as a user.") role = db.execute(select(Role).where(Role.name == "Firm Admin", Role.is_active.is_(True))).scalar_one_or_none() if not role: errors.append("Firm Admin role is missing or inactive. Please seed roles before using this wizard.") if payload["create_financial_year"]: if not payload["fy_year_code"]: errors.append("Financial year code is required when default FY is enabled.") if not payload["fy_assessment_year"]: errors.append("Assessment year is required when default FY is enabled.") try: start = parse_iso_date(payload["fy_start_date"]) end = parse_iso_date(payload["fy_end_date"]) if not start or not end: errors.append("Financial year start and end date are required.") elif end <= start: errors.append("Financial year end date must be after start date.") except Exception: errors.append("Financial year dates must be valid ISO dates, for example 2026-04-01.") return errors def _hash_token(token: str) -> str: return hashlib.sha256(token.encode("utf-8")).hexdigest() def create_invite_token_without_commit(db: Session, user: User) -> str: plain = secrets.token_urlsafe(32) now = utcnow() db.add( InviteToken( user_id=user.id, token_hash=_hash_token(plain), created_at_utc=now.replace(tzinfo=None), expires_at_utc=(now + timedelta(hours=get_settings().INVITE_TOKEN_HOURS)).replace(tzinfo=None), used_at_utc=None, ) ) user.must_change_password = True return plain def create_firm_from_wizard(db: Session, payload: dict[str, Any]) -> FirmWizardResult: errors = validate_firm_wizard_payload(db, payload) if errors: raise FirmWizardError(" ".join(errors)) role = db.execute(select(Role).where(Role.name == "Firm Admin", Role.is_active.is_(True))).scalar_one() temp_password = secrets.token_urlsafe(18) tenant = Tenant( code=payload["tenant_code"], name=payload["tenant_name"], display_name=payload["tenant_name"], is_active=True, firm_type=payload["firm_type"], default_timezone=payload["default_timezone"], default_session_duration_minutes=payload["default_session_duration_minutes"], default_otp_required_roles_csv=payload["default_otp_required_roles_csv"], default_storage_mode=payload["default_storage_mode"], contact_email=payload["admin_email"], contact_mobile=payload["admin_mobile"] or None, ) db.add(tenant) db.flush() branch = Branch( tenant_id=tenant.id, code=payload["branch_code"], name=payload["branch_name"], is_active=True, timezone=payload["branch_timezone"], allow_login=True, allow_new_assignments=True, is_head_office=True, smtp_use_tls=True, ) db.add(branch) db.flush() branch_settings = BranchSettings( branch_id=branch.id, address_line1=payload["branch_address_line1"] or None, address_line2=payload["branch_address_line2"] or None, city=payload["branch_city"] or None, state=payload["branch_state"] or None, pin_code=payload["branch_pin_code"] or None, gstin=payload["branch_gstin"] or None, pan=payload["branch_pan"] or None, storage_mode=payload["default_storage_mode"], otp_required_roles_csv=payload["default_otp_required_roles_csv"], session_duration_minutes=payload["default_session_duration_minutes"], ) db.add(branch_settings) db.flush() firm_admin = User( email=payload["admin_email"], full_name=payload["admin_full_name"], password_hash=hash_password(temp_password), tenant_id=tenant.id, branch_id=branch.id, is_active=True, allow_login=True, is_locked=False, deleted_at=None, must_change_password=True, password_changed_at_utc=None, mobile=payload["admin_mobile"] or None, designation=payload["admin_designation"] or "Firm Admin", ) db.add(firm_admin) db.flush() db.add(UserRole(user_id=firm_admin.id, role_id=role.id)) # Auto-create and link Employee Master for the primary Firm Admin. # Employee Portal requires employees.user_id to match the logged-in user. employee_code = normalize_code(f"{payload['tenant_code']}_ADMIN")[:50] firm_admin_employee = Employee( tenant_id=tenant.id, branch_id=branch.id, user_id=firm_admin.id, employee_code=employee_code, full_name=payload["admin_full_name"], email=payload["admin_email"], mobile=payload["admin_mobile"] or None, date_of_joining=date.today(), employment_type="full_time", status="active", is_active=True, department="Administration", designation=payload["admin_designation"] or "Firm Admin", created_by_user_id=firm_admin.id, updated_by_user_id=firm_admin.id, notes="Auto-created by Firm Creation Wizard for primary Firm Admin.", ) db.add(firm_admin_employee) db.flush() financial_year = None if payload["create_financial_year"]: financial_year = FinancialYear( tenant_id=tenant.id, year_code=payload["fy_year_code"], assessment_year=payload["fy_assessment_year"], start_date=parse_iso_date(payload["fy_start_date"]), end_date=parse_iso_date(payload["fy_end_date"]), is_current=payload["fy_is_current"], is_locked=False, created_at_utc=datetime.utcnow(), updated_at_utc=datetime.utcnow(), ) db.add(financial_year) db.flush() invite_token = create_invite_token_without_commit(db, firm_admin) invite_url = public_invite_url(invite_token) return FirmWizardResult( tenant=tenant, branch=branch, firm_admin=firm_admin, financial_year=financial_year, invite_url=invite_url, invite_token=invite_token, )