from datetime import datetime, timedelta from flask_sqlalchemy import BaseQuery from sqlalchemy import or_, func, distinct from sqlalchemy.orm import joinedload from atlas_um import pgdb from atlas_um.pgdb.views import dna_accounts_global_claims from atlas_um.settings import Settings class DNAAccountQuery(BaseQuery): def by_claim_external_ids( self, resource_group_id: str, claim_name_id: str, claim_id: str ): return ( self.join(dna_accounts_global_claims) .join(pgdb.ClaimValue) .join( pgdb.ClaimName, pgdb.ClaimName.id == dna_accounts_global_claims.c.claim_name_id, ) .join(pgdb.ResourceGroup) .filter(dna_accounts_global_claims.c.is_disabled == False) # noqa .filter(pgdb.ResourceGroup.external_id == resource_group_id) .filter(pgdb.ClaimName.external_id == claim_name_id) .filter(pgdb.ClaimValue.external_id == claim_id) ) def by_claim_name_claim_value( self, claim_name: pgdb.ClaimName, claim_value: pgdb.ClaimValue ): return ( self.join(dna_accounts_global_claims) .join(pgdb.ClaimValue) .filter(dna_accounts_global_claims.c.is_disabled == False) # noqa .filter( dna_accounts_global_claims.c.claim_name_id == claim_name.id ) .filter(pgdb.ClaimValue.id == claim_value.id) ) def with_relations(self): return self.options( joinedload(pgdb.DNAAccount.enabled_resource_groups), ).distinct() def search(self, search_term): search_expr = f"%{search_term.strip()}%" query = self.with_relations().filter( or_( pgdb.DNAAccount.sub.ilike(search_expr), pgdb.DNAAccount.given_name.ilike(search_expr), pgdb.DNAAccount.family_name.ilike(search_expr), pgdb.DNAAccount.name.ilike(search_expr), pgdb.DNAAccount.email.ilike(search_expr), ), ) return query def by_status(self, status): return self.filter(pgdb.DNAAccount.status == status) def by_ids(self, ids: list): return self.with_relations().filter(pgdb.DNAAccount.id.in_(ids)) def by_sub(self, sub: str): return self.filter(pgdb.DNAAccount.sub.ilike(sub)) def by_resource_group_id(self, resource_group_id): return self.join( pgdb.dna_accounts_enabled_resource_groups, pgdb.dna_accounts_enabled_resource_groups.c.dna_account_id == pgdb.DNAAccount.id, ).filter( pgdb.dna_accounts_enabled_resource_groups.c.resource_group_id == resource_group_id ) def by_tag_id(self, tag_id): return self.join( pgdb.dna_account_tag_table, pgdb.dna_account_tag_table.c.dna_account_id == pgdb.DNAAccount.id, ).filter(pgdb.dna_account_tag_table.c.tag_id == tag_id) def no_mfa_by_email(self, email): return self.with_relations().filter( pgdb.DNAAccount.email == email, pgdb.DNAAccount.no_mfa == True, # noqa ) def activated(self): return self.filter( pgdb.DNAAccount.status != pgdb.DNAAccountStatuses.SUSPENDED.value ) def affected_by(self, field, value) -> tuple: instance = None query = self.join( pgdb.dna_accounts_global_claims, pgdb.dna_accounts_global_claims.c.dna_account_id == pgdb.DNAAccount.id, ).join( pgdb.ClaimName, pgdb.ClaimName.id == pgdb.dna_accounts_global_claims.c.claim_name_id, ) if field == "claim_value_id": query = query.filter( pgdb.dna_accounts_global_claims.c.claim_value_id == value ) instance = pgdb.ClaimValue.query.filter_by(id=value).first() elif field == "claim_name_id": query = query.filter(pgdb.ClaimName.id == value) instance = pgdb.ClaimName.query.filter_by(id=value).first() elif field == "resource_group_id": query = query.filter(pgdb.ClaimName.resource_group_id == value) instance = pgdb.ResourceGroup.query.filter_by(id=value).first() elif field == "claim_values_source_id": query = query.filter( pgdb.ClaimName.claim_values_source_id == value ) instance = pgdb.ClaimName.query.filter_by(id=value).first() else: query = query.filter(False) return query, instance def affected_by_count(self, field): if field == "claim_value_id": query = ( pgdb.pgdb.session.query( pgdb.dna_accounts_global_claims.c.claim_value_id.label( "instance_id" ), func.Count(distinct(pgdb.DNAAccount.id)), ) .select_from(pgdb.DNAAccount) .group_by(pgdb.dna_accounts_global_claims.c.claim_value_id) ) elif field == "claim_name_id": query = ( pgdb.pgdb.session.query( pgdb.ClaimName.id.label("instance_id"), func.Count(distinct(pgdb.DNAAccount.id)), ) .select_from(pgdb.DNAAccount) .group_by(pgdb.ClaimName.id) ) elif field == "resource_group_id": query = ( pgdb.pgdb.session.query( pgdb.ClaimName.resource_group_id.label("instance_id"), func.Count(distinct(pgdb.DNAAccount.id)), ) .select_from(pgdb.DNAAccount) .group_by(pgdb.ClaimName.resource_group_id) ) elif field == "claim_values_source_id": query = ( pgdb.pgdb.session.query( pgdb.ClaimName.claim_values_source_id.label("instance_id"), func.Count(distinct(pgdb.DNAAccount.id)), ) .select_from(pgdb.DNAAccount) .group_by(pgdb.ClaimName.claim_values_source_id) ) else: query = pgdb.pgdb.session.query( func.Max(None).label("instance_id"), func.Count(pgdb.DNAAccount.id), ) query = query.join( pgdb.dna_accounts_global_claims, pgdb.dna_accounts_global_claims.c.dna_account_id == pgdb.DNAAccount.id, ).join( pgdb.ClaimName, pgdb.ClaimName.id == pgdb.dna_accounts_global_claims.c.claim_name_id, ) return query def token_size_violation(self): ids = [] for account in self.all(): if account.account_violates_limit: ids.append(account.id) return self.filter(pgdb.DNAAccount.id.in_(ids)) def expire_soon(self): return self.filter( pgdb.DNAAccount.expiration_date != None, # noqa pgdb.DNAAccount.expiration_date <= datetime.now().date() + timedelta(days=Settings.ACCOUNTS_EXPIRATION_NOTIFY_DAYS), pgdb.DNAAccount.expiration_date > datetime.now().date(), )