"""Track Pricing Override Model. This model represents a Track Pricing Override """ import datetime from oto import response from pricing.connectors import mysql from pricing.constants import error from pricing.models import track_pricing_override_store from pricing.models import track_pricing_override_territory from pricing.models.track_pricing_override_store import TrackPricingOverrideStore from pricing.models.track_pricing_override_territory import TrackPricingOverrideTerritory import sqlalchemy class TrackPricingOverride(mysql.BaseModel): """Track Pricing Override model.""" __tablename__ = 'track_pricing_override' track_pricing_override_id = sqlalchemy.Column( sqlalchemy.BIGINT, primary_key=True) orchard_pricing_tier_id = sqlalchemy.Column(sqlalchemy.BIGINT) track_id = sqlalchemy.Column(sqlalchemy.Integer) custom_price = sqlalchemy.Column(sqlalchemy.VARCHAR(30)) custom_currency_code = sqlalchemy.Column(sqlalchemy.VARCHAR(10)) start_date = sqlalchemy.Column(sqlalchemy.DateTime) end_date = sqlalchemy.Column(sqlalchemy.DateTime) applies_worldwide = sqlalchemy.Column(sqlalchemy.Boolean, default=False) territory_list_include = sqlalchemy.Column( sqlalchemy.Boolean, default=False) price_code = sqlalchemy.Column(sqlalchemy.VARCHAR(45)) sort_order = sqlalchemy.Column(sqlalchemy.Integer) activated = sqlalchemy.Column(sqlalchemy.Boolean, default=True) created_date = sqlalchemy.Column(sqlalchemy.DateTime) created_by = sqlalchemy.Column(sqlalchemy.VARCHAR(45)) updated_date = sqlalchemy.Column(sqlalchemy.DateTime) updated_by = sqlalchemy.Column(sqlalchemy.VARCHAR(45)) def to_dict(self): """Convert Track Pricing Override to dict.""" start_date = None if isinstance(self.start_date, datetime.datetime): start_date = self.start_date.timestamp() end_date = None if isinstance(self.end_date, datetime.datetime): end_date = self.end_date.timestamp() created_date = None if isinstance(self.created_date, datetime.datetime): created_date = self.created_date.timestamp() updated_date = None if isinstance(self.updated_date, datetime.datetime): updated_date = self.updated_date.timestamp() return dict( track_pricing_override_id=self.track_pricing_override_id, orchard_pricing_tier_id=self.orchard_pricing_tier_id, track_id=self.track_id, custom_price=self.custom_price, custom_currency_code=self.custom_currency_code, start_date=start_date, end_date=end_date, applies_worldwide=self.applies_worldwide, territory_list_include=self.territory_list_include, price_code=self.price_code, sort_order=self.sort_order, activated=self.activated, created_date=created_date, updated_date=updated_date ) @mysql.autosession() def create_track_pricing_override(track_id, data, session): """Create a new track pricing override. Args: track_id (int): the track id data (dict): the data from which to create the pricing override. session (Session): the mysql session. Returns: response.Response: containing the created pricing override dict. """ if not data: return response.create_error_response( 400, error.ERROR_MESSAGE_EMPTY_BODY) data['created_date'] = datetime.datetime.now() data['track_id'] = track_id track_pricing_override = TrackPricingOverride(**data) session.add(track_pricing_override) session.commit() return response.Response(track_pricing_override.to_dict()) def _get_by_track_pricing_override_id(track_pricing_override_id, session): """Get a track pricing override by its ID. Args: track_pricing_override_id (int): track pricing override id session (Session): the mysql session Returns: TrackPricingOverride: the found track pricing override """ found_track_pricing_override = session.query( TrackPricingOverride).get(track_pricing_override_id) return found_track_pricing_override def _get_by_track_id_and_track_pricing_override_id( track_id, track_pricing_override_id, session): """Get a track pricing override by its ID. Args: track_pricing_override_id (int): track pricing override id session (Session): the mysql session Returns: TrackPricingOverride: the found track pricing override """ found_track_pricing_override = session.query( TrackPricingOverride).filter_by( track_id=track_id, track_pricing_override_id=track_pricing_override_id).one_or_none() return found_track_pricing_override @mysql.autosession() def get_by_track_pricing_override_id(track_pricing_override_id, session): """Get a track pricing override by its ID. Args: track_pricing_override_id (int): track pricing override id session (Session): the mysql session Returns: TrackPricingOverride: the found track pricing override """ found_track_pricing_override = _get_by_track_pricing_override_id( track_pricing_override_id, session) if not found_track_pricing_override: return response.create_not_found_response() result = found_track_pricing_override.to_dict() return response.Response(result) def _get_territories_by_track_id(track_id, session): """Get territories by track ID. Args: track_id (int): track ID session (Session): the mysql session Returns: array: the found track pricing override territories """ territories = ( session.query(TrackPricingOverrideTerritory) .join( TrackPricingOverride, TrackPricingOverrideTerritory.track_pricing_override_id == TrackPricingOverride.track_pricing_override_id ) .where(TrackPricingOverride.track_id == track_id) ).all() return [row.to_dict() for row in territories] def _get_stores_by_track_id(track_id, session): """Get stores by track ID. Args: track_id (int): track ID session (Session): the mysql session Returns: array: the found track pricing override stores """ stores = ( session.query(TrackPricingOverrideStore) .join( TrackPricingOverride, TrackPricingOverrideStore.track_pricing_override_id == TrackPricingOverride.track_pricing_override_id ) .where(TrackPricingOverride.track_id == track_id) ).all() return [row.to_dict() for row in stores] @mysql.autosession() def get_by_track_id(track_id, session): """Get track pricing overrides with the track_id. Args: track_id (int): the ID of a track session (Session): the mysql session Returns: response.Response: The track pricing overrides """ track_pricing_overrides = session.query( TrackPricingOverride).filter_by( track_id=track_id).order_by(TrackPricingOverride.sort_order.desc()).all() territory_rows = _get_territories_by_track_id(track_id, session) store_rows = _get_stores_by_track_id(track_id, session) result = [row.to_dict() for row in track_pricing_overrides] for override in result: override['territories'] = [] override['stores'] = [] for row in territory_rows: if override['track_pricing_override_id'] == row[ 'track_pricing_override_id']: override['territories'].append(row['territory_code']) for row in store_rows: if override['track_pricing_override_id'] == row[ 'track_pricing_override_id']: override['stores'].append(row['store_id']) return response.Response( { 'items': result }) @mysql.autosession() def update_track_pricing_override( track_id, track_pricing_override_id, data, session): """Update a track pricing override. Args: track_id (int): the id of the track track_pricing_override_id (int): the track pricing override id data (dict): The data from which to update the track pricing override session: the mysql session Returns: response.Response: containing the updated track pricing override. """ if not data: return response.create_error_response( 400, error.ERROR_MESSAGE_EMPTY_BODY) data['updated_date'] = datetime.datetime.now() found_track_pricing_override = \ _get_by_track_id_and_track_pricing_override_id( track_id, track_pricing_override_id, session) if not found_track_pricing_override: return response.create_not_found_response() for key, value in data.items(): setattr(found_track_pricing_override, key, value) session.commit() return response.Response(found_track_pricing_override.to_dict()) @mysql.autosession() def find_duplicate(track_id, data, session): """Find a duplicate track pricing override. Args: track_id (int): the ID of the track. data (dict): the data to use to find the duplicate. session (Session): the mysql session. Returns: response.Response: containing True if there's a duplicate and False if there isn't. """ if not data: return response.create_error_response( 400, error.ERROR_MESSAGE_EMPTY_BODY) data['track_id'] = track_id data['activated'] = True match = session.query(TrackPricingOverride).filter_by(**data).first() if match is None: return response.Response({'found': False}) else: return response.Response({'found': True, 'item': match.to_dict()}) @mysql.autosession() def delete_by_track_id(track_id, session): """Delete all overrides for a track. Args: track_id (int): the ID of the track. session (Session): the mysql session. Returns: response.Response: containing the number of deleted overrides """ track_pricing_overrides = session.query(TrackPricingOverride)\ .filter_by(track_id=track_id) result = [] for override in track_pricing_overrides: track_pricing_override_territory.\ delete_by_track_pricing_override_id( override.track_pricing_override_id) track_pricing_override_store.\ delete_by_track_pricing_override_id( override.track_pricing_override_id) session.delete(override) result.append(override) return response.Response( { 'deleted': len(result) })