from datetime import date from enum import Enum from flask import current_app from wtforms import fields, validators from wtforms_alchemy.fields import QuerySelectField, QuerySelectMultipleField from atlas_um import pgdb from atlas_um.consts import Products from atlas_um.dna_accounts.services import ImportDNAAccountFromAuth0Service from atlas_um.helpers.fields import ( ClaimValueSelectField, ClaimValueSelectMultipleField, ) from atlas_um.helpers.forms import BaseForm from atlas_um.pgdb.dna_account import DNAAccountStatuses class DNAAccountForm(BaseForm): class Actions(Enum): SAVE = "Save" SUSPEND = "Suspend" action = fields.SelectField( choices=[(action.name, action.value) for action in Actions], default=Actions.SAVE.name, ) given_name = fields.StringField( "First Name", validators=[validators.Optional()] ) family_name = fields.StringField( "Last Name", validators=[validators.Optional()] ) email = fields.StringField( "Email", validators=[validators.Email(), validators.Optional()] ) is_sony_employee = fields.BooleanField("Sony Employee") is_vip = fields.BooleanField("VIP") job_category = QuerySelectField( "Job Category", render_kw={"class": "ui dropdown"}, query_factory=lambda: pgdb.JobCategory.query.active(), allow_blank=True, validators=[validators.Optional()], blank_text="N/A", ) business_unit = QuerySelectField( "Business Unit", render_kw={"class": "ui dropdown"}, query_factory=lambda: pgdb.BusinessUnit.query.active(), allow_blank=True, validators=[validators.Optional()], blank_text="N/A", ) personnel_type = QuerySelectField( "Personnel Type", render_kw={"class": "ui dropdown"}, query_factory=lambda: pgdb.PersonnelType.query.active(), allow_blank=True, validators=[validators.Optional()], blank_text="N/A", ) expiration_date = fields.DateField( "Expiration Date", format="%m/%d/%y", validators=[validators.Optional()], ) supervisor_email = fields.StringField( "Supervisor Email", validators=[validators.Email(), validators.Optional()], render_kw={"placeholder": "Search email"}, ) supervisor_name = fields.HiddenField(validators=[validators.Optional()]) job_title = fields.StringField( "Job title", validators=[validators.Optional()], render_kw={"readonly": True}, ) location = fields.StringField( "Location", validators=[validators.Optional()], render_kw={"readonly": True}, ) preferred_username = fields.StringField( "Sony username", validators=[validators.Optional()], render_kw={"readonly": True}, ) tags = QuerySelectMultipleField( "Tags", query_factory=lambda: pgdb.Tag.query.active(), validators=[validators.Optional()], ) no_mfa = fields.BooleanField("No MFA", validators=[validators.Optional()]) def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) # if the account was suspended, init date blank to reactivate it if ( self.obj is not None and self.obj.status == DNAAccountStatuses.SUSPENDED and self["expiration_date"].data is not None and self["expiration_date"].data <= date.today() ): self["expiration_date"].data = None def validate(self, *args, **kwargs): is_valid = super().validate(*args, **kwargs) required_extra_fields = [] if ( self.obj is None or self.obj is not None and not self.obj.usm_account ): required_extra_fields.extend( ("given_name", "family_name", "email") ) if ( self["email"].data and self.obj is None and pgdb.DNAAccount.query.filter( pgdb.DNAAccount.email.ilike(self["email"].data) ).first() is not None ): self["email"].errors.append("The account is already exists") is_valid = False elif ( self["email"].data and self.obj is not None and pgdb.DNAAccount.query.filter( pgdb.DNAAccount.email.ilike(self["email"].data) ) .filter(pgdb.DNAAccount.id != self.obj.id) .first() is not None ): self["email"].errors.append("The account is already exists") is_valid = False for field in required_extra_fields: if not self[field].data: self[field].errors.append(self.FIELD_REQUIRED_MESSAGE) is_valid = False return is_valid @property def no_mfa_fields_are_available(self): return current_app.config.get("NO_MFA_ENABLED") class BaseProductForm(BaseForm): class Actions(Enum): SAVE = "Save" DISABLE = "Disable" action = fields.SelectField( choices=[(action.name, action.value) for action in Actions], default=Actions.SAVE.name, ) @property def claim_values(self): values = [] for value in self.data.values(): if isinstance(value, pgdb.ClaimValue): values.append(value) elif isinstance(value, list): values.extend( v for v in value if isinstance(v, pgdb.ClaimValue) ) return values @property def claim_fields(self): return [ field for field in self if isinstance( field, (ClaimValueSelectField, ClaimValueSelectMultipleField) ) ] def validate(self, *args, **kwargs): is_valid = super().validate(*args, **kwargs) if self["action"].data == self.Actions.DISABLE.name: is_valid = True return is_valid def product_form_factory(product, **kwargs): """ Create dynamic form with fields based on claims for specific product. """ data = {} form_fields = {} for claim_name, claim_value in product.claims.items(): field_name = claim_name.internal_name if claim_name.multiple_values_allowed: field = ClaimValueSelectMultipleField(claim_name) data[field_name] = [v for v in claim_value] else: field = ClaimValueSelectField(claim_name) data[field_name] = ( [v for v in claim_value][0] if len(claim_value) else None ) form_fields[field_name] = field cls = type("ProductForm", (BaseProductForm,), form_fields) return cls(data=data, **kwargs) class ImportAuth0UsersForm(BaseForm): product = fields.SelectField( choices=[ (product, product.capitalize()) for product in ImportDNAAccountFromAuth0Service.MIGRATION_SETTINGS if product in (Products.RTI.value, Products.APOLLO.value) ] ) user_ids = fields.StringField( label="Filter ids", validators=[validators.Optional()], filters=[lambda val: val.split(",") if val else None], ) import_all = fields.BooleanField( label="Import all users", validators=[validators.Optional()], default=False, ) batch_from = fields.IntegerField( validators=[validators.Optional()], default=None, ) batch_to = fields.IntegerField( validators=[validators.Optional()], default=None, ) def validate(self, *args, **kwargs): is_valid = super().validate(*args, **kwargs) if not self["user_ids"].data and not self["import_all"].data: self["user_ids"].errors.append(self.FIELD_REQUIRED_MESSAGE) is_valid = False return is_valid