"""Feature Model.""" from collections import defaultdict from typing import Any from ddtrace import tracer from owsresponse import response from sqlalchemy import Column, ForeignKey, Integer, String, text from sqlalchemy.orm import Session from account.connectors import mysql from account.constants import error from account.constants.features import FEATURES, FEATURES_SET from account.models.types import Feature as FeatureType GET_ENABLED_FEATURES_FOR_VENDOR = """ SELECT f.feature_id, f.feature_name FROM features f WHERE f.feature_id NOT IN ( SELECT vrf.feature_id FROM vendor_restricted_features vrf WHERE vrf.vendor_id = {vendor_id} ) """ GET_ENABLED_FEATURES_FOR_VENDORS = """SELECT v.column_0 as vendor_id, f.feature_id, f.feature_name FROM ( VALUES {row_values} ) AS v INNER JOIN vendor vnd ON vnd.vendor_id = v.column_0 CROSS JOIN features f WHERE NOT EXISTS ( SELECT 1 FROM vendor_restricted_features vrf WHERE vrf.vendor_id = v.column_0 AND vrf.feature_id = f.feature_id ) ORDER BY v.column_0, f.feature_id""" GET_RESTRICTED_FEATURES_FOR_VENDOR = """ SELECT f.feature_id, f.feature_name FROM features f WHERE f.feature_id IN ( SELECT vrf.feature_id FROM vendor_restricted_features vrf WHERE vrf.vendor_id = {vendor_id} ) """ GET_ENABLED_FEATURE_IDS_FOR_VENDOR_UUIDS = """ SELECT v.vendor_uuid, f.feature_id FROM vendor v CROSS JOIN features f WHERE v.vendor_uuid IN ({placeholders}) AND NOT EXISTS ( SELECT 1 FROM vendor_restricted_features vrf WHERE vrf.vendor_id = v.vendor_id AND vrf.feature_id = f.feature_id ) ORDER BY v.vendor_uuid, f.feature_id""" GET_ENABLED_FEATURE_IDS_FOR_SUBACCOUNT_UUIDS = """ SELECT sa.subaccount_uuid, f.feature_id FROM subaccount sa CROSS JOIN features f WHERE sa.subaccount_uuid IN ({placeholders}) AND NOT EXISTS ( SELECT 1 FROM vendor_restricted_features vrf WHERE vrf.vendor_id = sa.vendor_id AND vrf.feature_id = f.feature_id ) ORDER BY sa.subaccount_uuid, f.feature_id""" GET_VENDORS_WITH_FEATURE = """ SELECT {select_query} FROM vendor v WHERE v.vendor_id NOT IN ( SELECT vendor_id FROM vendor_restricted_features vrf WHERE vrf.feature_id = {feature_id} ) AND EXISTS ( SELECT 1 FROM features f WHERE f.feature_id = {feature_id} ) {limit_offset} """ class VendorRestrictedFeatures(mysql.BaseModel): """VendorRestrictedFeatures DB Model.""" __tablename__ = 'vendor_restricted_features' vendor_restricted_features_id = Column(Integer, primary_key=True) vendor_id = Column(Integer) feature_id = Column(Integer, ForeignKey('features.feature_id')) def to_dict(self): """Get a dict representation of VendorRestrictedFeatures.""" return {'vendor_id': self.vendor_id, 'feature_id': self.feature_id} class Features(mysql.BaseModel): """Features DB Model.""" __tablename__ = 'features' feature_id = Column(Integer, primary_key=True) feature_name = Column(String) def to_dict(self): """Get a dict representation of Features.""" return {'feature_id': self.feature_id, 'feature_name': self.feature_name} @tracer.wrap() def get_feature(feature_id): """Get feature for a given feature_id. Args: feature_id (int): unique identifier for the feature. Returns: response.Response: containing dict feature or error. """ with mysql.session_scope(read_only=True) as session: row = session.query(Features).filter(Features.feature_id == feature_id).first() if row: return response.Response(row.to_dict()) return response.create_not_found_response( error.ERROR_MESSAGE_FEATURE_NOT_FOUND.format(feature_id=feature_id) ) @tracer.wrap() def is_valid_features(feature_ids: list): """Check feature_ids to be existing Features. Args: feature_ids (list[int]): unique identifier for the feature. Returns: True: If all features are valid. response.Response: 404 Error response with invalid feature ids. """ with mysql.session_scope(read_only=True) as session: valid_feature_ids_response = ( session.query(Features.feature_id).filter(Features.feature_id.in_(feature_ids)).all() ) valid_feature_ids = [feature_id for (feature_id,) in valid_feature_ids_response] invalid_feature_ids = [x for x in feature_ids if x not in valid_feature_ids] if not invalid_feature_ids: return response.Response(True) return response.create_not_found_response( error.ERROR_MESSAGE_FEATURE_NOT_FOUND.format(feature_id=invalid_feature_ids) ) @tracer.wrap() def get_enabled_feature_ids_for_vendor_uuids( vendor_uuids: list[str], session: Session ) -> dict[str, list[int]]: """Get enabled feature IDs for multiple vendors by UUID. Args: vendor_uuids (list[str]): list of vendor UUIDs. session (Session): SQLAlchemy session. Returns: dict[str, list[int]]: mapping of vendor_uuid to list of enabled feature IDs. Only UUIDs present in the vendor table are included. """ if not vendor_uuids: return {} placeholders = ', '.join(f':vendor_uuid_{i}' for i in range(len(vendor_uuids))) params = {f'vendor_uuid_{i}': uuid for i, uuid in enumerate(vendor_uuids)} query = text(GET_ENABLED_FEATURE_IDS_FOR_VENDOR_UUIDS.format(placeholders=placeholders)) rows = session.execute(query, params).fetchall() result = defaultdict(list) for row in rows: result[row[0]].append(row[1]) return dict(result) @tracer.wrap() def get_enabled_feature_ids_for_subaccount_uuids( subaccount_uuids: list[str], session: Session ) -> dict[str, list[int]]: """Get enabled feature IDs for multiple subaccounts by UUID. Args: subaccount_uuids (list[str]): list of subaccount UUIDs. session (Session): SQLAlchemy session. Returns: dict[str, list[int]]: mapping of subaccount_uuid to list of enabled feature IDs. Only UUIDs present in the subaccount table are included. """ if not subaccount_uuids: return {} placeholders = ', '.join(f':subaccount_uuid_{i}' for i in range(len(subaccount_uuids))) params = {f'subaccount_uuid_{i}': uuid for i, uuid in enumerate(subaccount_uuids)} query = text(GET_ENABLED_FEATURE_IDS_FOR_SUBACCOUNT_UUIDS.format(placeholders=placeholders)) rows = session.execute(query, params).fetchall() result = defaultdict(list) for row in rows: result[row[0]].append(row[1]) return dict(result) @tracer.wrap() def get_enabled_features_for_vendor_with_session( vendor_id: int, session: Session ) -> list[FeatureType]: """Get a list of enabled feature controls for a vendor with a session. Args: vendor_id (int): unique identifier for a vendor. session (Session): SQLAlchemy session. Returns: list[FeatureType]: list of enabled feature controls. """ rows = session.execute(GET_ENABLED_FEATURES_FOR_VENDOR.format(vendor_id=vendor_id)).fetchall() items = [] for row in rows: items.append({'feature_id': row[0], 'feature_name': row[1]}) return items @tracer.wrap() def get_enabled_features_for_vendors_with_session( vendor_ids: list[int], session: Session ) -> dict[int, list[FeatureType]]: """Get enabled feature controls for multiple vendors with a session. Args: vendor_ids (list[int]): list of vendor identifiers. session (Session): SQLAlchemy session. Returns: dict[int, list[FeatureType]]: mapping of vendor_id to list of enabled feature controls. """ if not vendor_ids: return {} # Use bound parameters to prevent SQL injection placeholders = ', '.join(f'ROW(:vendor_id_{i})' for i in range(len(vendor_ids))) params = {f'vendor_id_{i}': vendor_id for i, vendor_id in enumerate(vendor_ids)} query = text(GET_ENABLED_FEATURES_FOR_VENDORS.format(row_values=placeholders)) rows = session.execute(query, params).fetchall() # Group results by vendor_id result = defaultdict(list) for row in rows: vendor_id = row[0] result[vendor_id].append({'feature_id': row[1], 'feature_name': row[2]}) return result @tracer.wrap() def get_enabled_features_for_vendor(vendor_id): """Get a list of enabled feature controls for a vendor. Args: vendor_id (int): unique identifier for a vendor. Returns: response.response: containing the list of enabled feature controls. """ with mysql.session_scope(read_only=True) as session: return response.Response( {'items': get_enabled_features_for_vendor_with_session(vendor_id, session)} ) @tracer.wrap() def get_vendors_with_feature( feature_id: int, page_offset: int, page_limit: int, ) -> dict[str, Any]: """Get a list of vendors with a specific feature enabled. Args: feature_id (int): unique identifier for a feature. page_offset (int): offset for pagination. page_limit (int): limit for pagination. Returns: response.response: containing the list of vendors with the feature enabled. """ with mysql.session_scope(read_only=True) as session: query = GET_VENDORS_WITH_FEATURE.format( select_query='v.vendor_id', feature_id=feature_id, limit_offset=f'LIMIT {page_limit} OFFSET {page_offset}' if page_limit != 0 else '', ) count_query = GET_VENDORS_WITH_FEATURE.format( select_query='COUNT(*)', feature_id=feature_id, limit_offset='' ) query_rows = session.execute(query).fetchall() query_count = session.execute(count_query).fetchone() or 0 pagination = { 'page_offset': page_offset, 'page_limit': page_limit, 'count': query_count[0], } return { 'items': [row[0] for row in query_rows], 'pagination': pagination, } @tracer.wrap() def get_restricted_features_for_vendor(vendor_id): """Get a list of restricted feature controls for a vendor. Args: vendor_id (int): unique identifier for a vendor. Returns: response.response: containing the list of restricted feature controls. """ with mysql.session_scope(read_only=True) as session: rows = session.execute( GET_RESTRICTED_FEATURES_FOR_VENDOR.format(vendor_id=vendor_id) ).fetchall() items = [] for row in rows: items.append({'feature_id': row[0], 'feature_name': row[1]}) return response.Response({'items': items}) @tracer.wrap() def bulk_add_restricted_features_for_vendor(vendor_id, feature_ids): """Bulk add restricted features for vendor. Function is idempotent. Multiple identical calls will not create duplicate entries or errors. Existing entries will not be overwritten, missing entries from the `feature_ids` will be created. Previously existing and created entries from the `feature_ids` will be returned in the response. Args: vendor_id (int): unique identifier for vendor. feature_ids (list[int]): feature ids Returns: response.Response: list of created vendor restricted feature. """ with mysql.session_scope() as session: existing_restricted_features = ( session.query(VendorRestrictedFeatures) .filter(VendorRestrictedFeatures.vendor_id == vendor_id) .filter(VendorRestrictedFeatures.feature_id.in_(feature_ids)) .all() ) existing_restricted_features_ids = [x.feature_id for x in existing_restricted_features] for feature_id in feature_ids: if feature_id not in existing_restricted_features_ids: restricted_feature = VendorRestrictedFeatures( vendor_id=vendor_id, feature_id=feature_id ) session.add(restricted_feature) existing_restricted_features.append(restricted_feature) session.commit() return response.Response([feature.to_dict() for feature in existing_restricted_features]) @tracer.wrap() def bulk_remove_restricted_features_for_vendor(vendor_id, feature_ids): """Bulk remove restricted features for vendor. Args: vendor_id (int): unique identifier for vendor. feature_ids (list[int]): feature ids Returns: response.Response: list of deleted vendor restricted features ids. """ with mysql.session_scope() as session: existing_restricted_features = ( session.query(VendorRestrictedFeatures) .filter(VendorRestrictedFeatures.vendor_id == vendor_id) .filter(VendorRestrictedFeatures.feature_id.in_(feature_ids)) .all() ) deleted_feature_ids = [] for restricted_feature in existing_restricted_features: if restricted_feature.feature_id in feature_ids: session.delete(restricted_feature) deleted_feature_ids.append(restricted_feature.feature_id) session.commit() return response.Response(deleted_feature_ids) @tracer.wrap() def convert_features_to_list(features: list[FeatureType]) -> list[FEATURES]: """ Convert a list of Feature objects into a list of FEATURES enum values. Args: features (list[Feature]): A list of Feature objects. Returns: list[FEATURES]: A list of FEATURES enum values. """ return [ FEATURES(feature['feature_id']) for feature in features if feature['feature_id'] in FEATURES_SET ]