"""SQLAlchemy model for Territory storage.""" from sqlalchemy import Column from sqlalchemy import Integer from sqlalchemy import String from sqlalchemy import types from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.orm import aliased from territories.constants import response as const from territories.response import create_error_response from territories.response import create_fatal_response from territories.response import Response from territories.util import model_util from territories.util.handler_util import add_pagination CONTINENTS = [ 'Africa', 'Antarctica', 'Asia', 'Europe', 'North America', 'Oceania', 'South America'] # TODO Remove this when RDS_OWS_TERRITORIES flag deleted BaseModel = model_util.get_base_model() session_scope = model_util.get_session_scope() class Territory(BaseModel): """Territory class.""" __tablename__ = 'territory' id = Column(Integer, primary_key=True) orch_id = Column(Integer, nullable=True) standard = Column(String(255), nullable=False) territory_name = Column(String(255), nullable=False) territory_code_a2 = Column(String(2), nullable=False) territory_code_a3 = Column(String(3), nullable=False) territory_code_numeric = Column(Integer, nullable=False) continent = Column(types.Enum(*CONTINENTS), nullable=True) def __init__( self, territory_id, orch_id, standard, territory_name, territory_code_a2, territory_code_a3, territory_code_numeric, continent=None): """Initialize object Territory. Args: territory_id (int): unique ID of record. orch_id (int): ID assigned by Orchard list. standard (string): type of standart ISO_3166_1_2016 | Orch_1_2016. territory_name (string): full territory name. territory_code_a2: territory alpha-2-code. territory_code_a3: territory alpha-2-code. territory_code_numeric: territory numerical-code. continent (str): optional continent """ self.id = territory_id self.orch_id = orch_id self.standard = standard self.territory_name = territory_name self.territory_code_a2 = territory_code_a2 self.territory_code_a3 = territory_code_a3 self.territory_code_numeric = territory_code_numeric if continent: # Ignore an empty string or any other non-truthy value self.continent = continent class TerritoryRelationship(BaseModel): """Map one territory ID to another.""" __tablename__ = 'territory_relationship' source = Column(Integer, nullable=False, primary_key=True) target = Column(Integer, nullable=False, primary_key=True) def __init__(self, source, target): """Initialize TerritoryRelationship object. Args: source (int): ID of source territory. target (int): ID of target territory. """ self.source = source self.target = target territory_input = aliased(Territory, name='input_territory') territory_output = aliased(Territory, name='output_territory') def get_standards(): """Get list of all available standards. Returns: Response: standards list """ try: with session_scope() as db_session: result = db_session.query(Territory.standard).distinct().all() except SQLAlchemyError as ex: return create_fatal_response(ex) paginated = add_pagination([standard for (standard,) in result]) return Response(paginated) def get_territories(standard): """Get all territories for a given standard. Args: standard (str): name of the standard Returns: Response: list of all territories for a given standard """ try: with session_scope() as db_session: result = ( db_session.query(Territory). filter(Territory.standard == standard).all()) except SQLAlchemyError as ex: return create_fatal_response(ex) if not result: return Response(const.NOT_FOUND, status=404) paginated = add_pagination([row.to_dict() for row in result]) return Response(paginated) def convert_territories(input_standard, output_standard, territories): """Convert given list of territories from input to output standard. Args: input_standard (str): name of the standard to convert from output_standard (str): name of the standard to convert to territories (list): list of territories to be converted Returns: Response: list of territories converted to output standard """ try: with session_scope() as db_session: query = ( db_session.query(territory_output). join( TerritoryRelationship, territory_output.id == TerritoryRelationship.target). join( territory_input, TerritoryRelationship.source == territory_input.id). filter( territory_input.standard == input_standard, territory_output.standard == output_standard, territory_input.territory_code_a2.in_(territories))) result = query.all() except SQLAlchemyError as ex: return create_fatal_response(ex) if not result: return Response(const.NOT_FOUND, status=404) paginated = add_pagination([row.to_dict() for row in result]) return Response(paginated) def get_complement(standard, territories): """Get complement for a given list of territories of given standard. Args: standard (str): name of the standard territories (list): list of territories to be converted Returns: Response: complement for a given list of territories """ all_territories = get_territories(standard) if not all_territories: return create_error_response( const.INVALID_STANDARD, const.STANDARD_NOT_FOUND) territory_names = ( item['territory_code_a2'] for item in all_territories.message['items']) if territories and not set(territories).issubset(set(territory_names)): return create_error_response( const.INVALID_COUNTRY_CODE, const.INVALID_TERRITORY) try: with session_scope() as db_session: result = ( db_session.query(Territory). filter( Territory.standard == standard, ~Territory.territory_code_a2.in_(territories) ).all()) except SQLAlchemyError as ex: return create_fatal_response(ex) paginated = add_pagination([row.to_dict() for row in result]) return Response(paginated)