import operator from typing import List, Type, Union from sqlalchemy import asc, desc, or_ from sqlalchemy.orm import Query from delphi_api.errors import Codes, InvalidInputError from delphi_api.v3.constants import SORT_DESC from delphi_api.v3.data_models.postgres_db import ( Artist, BaseModel, Chart, Dsp, Label, Playlist, Product, Region, Track, db, ) from delphi_api.v3.view_models.params import Params class QueryBuilder: """ Mostly SQLAlchemy-specific (Postgres DB) query builder """ @staticmethod def param_as_list(param: Union[List[str], str]): if isinstance(param, list): return param return [param] @classmethod def sort_results(cls, results: List[dict], params: Params): """Agnostic in-place sorting for a list of dictionaries. Sorts by dictionary key.""" if not params.sort_by: return results if len(results) <= 1: return results reverse = params.sort_order == SORT_DESC sort_by = params.sort_by.split('.') if len(sort_by) == 1: key_fn = operator.itemgetter(sort_by[0]) elif len(sort_by) == 2: key_fn = lambda x: x[sort_by[0]][sort_by[1]] else: raise InvalidInputError({ 'code': Codes.invalid_input, 'description': 'sort_by parameter is only supported up to two levels deep.' f'You provided: {len(sort_by)}' }) try: results.sort(key=key_fn, reverse=reverse) return results except TypeError as e: # some of our sort key values are None if len(sort_by) > 1: # don't attempt to do nested sorting with nulls return results # sort the ones that can be sorted, and append nulls at the end return cls._sort_with_nulls(results, sort_by[0], key_fn, reverse) @classmethod def _sort_with_nulls(cls, results: List[dict], key: str, key_fn: callable, reverse: bool): """Objects may request sorting with null values for some keys. This method will sort the ones that can be sorted, and append the others (null values) to the end """ items, unsorted = [], [] for obj in results: if obj and obj.get(key) is None: unsorted.append(obj) else: items.append(obj) items.sort(key=key_fn, reverse=reverse) items.extend(unsorted) return items @staticmethod def limit_offset_results(results: List[dict], params: Params): return results[params.offset:params.limit + params.offset] @staticmethod def add_query_limit_offset(q: Query, params: Params) -> Query: return q.limit(params.limit).offset(params.offset) @staticmethod def add_sort_clause(q: Query, params: Params, parent: BaseModel) -> Query: """Adds a relational database sorting clause based on parameters""" if params.sort_by and hasattr(parent, params.sort_by): direction = desc if params.sort_order == SORT_DESC else asc sort_fields = [] if params.group_by and hasattr(parent, params.group_by): sort_fields.append(direction(getattr(parent, params.group_by))) sort_fields.append(direction(getattr(parent, params.sort_by))) return q.order_by(*sort_fields) return q @staticmethod def or_filters_from_params(q: Query, params: Params, parent: BaseModel) -> Query: or_clauses = [] # --- Artist --- # if params.artist_id and parent is not Artist: artist_join = parent.artists if hasattr(parent, 'artists') else parent.artist q = q.join(Artist, artist_join) for artist_id in QueryBuilder.param_as_list(params.artist_id): or_clauses.append(Artist.artist_id == artist_id) # --- DSP --- # if params.dsp or params.dsp_id and parent is not Dsp: q = q.join(Dsp) if params.dsp: for dsp_id in QueryBuilder.param_as_list(params.dsp): or_clauses.append(Dsp.dsp_id == dsp_id) if params.dsp_id: for dsp_id in QueryBuilder.param_as_list(params.dsp_id): or_clauses.append(Dsp.dsp_id == dsp_id) # --- Label --- # if params.label_id and parent is not Label: q = q.join(Label) for label_id in QueryBuilder.param_as_list(params.label_id): or_clauses.append(Label.label_id == label_id) # --- Playlist --- # if params.playlist_id and parent is not Playlist: playlist_join = parent.playlists if hasattr(parent, 'playlists') else parent.playlist q = q.join(Playlist, playlist_join) for playlist_id in QueryBuilder.param_as_list(params.playlist_id): or_clauses.append(Playlist.playlist_id == playlist_id) # --- Chart --- # if params.chart_id and parent is not Chart: chart_join = parent.charts if hasattr(parent, 'charts') else parent.chart q = q.join(Chart, chart_join) for chart_id in QueryBuilder.param_as_list(params.chart_id): or_clauses.append(Chart.chart_id == chart_id) # --- Product --- # if params.product_id and parent is not Product: product_join = parent.products if hasattr(parent, 'products') else parent.product q = q.join(Product, product_join) for product_id in QueryBuilder.param_as_list(params.product_id): or_clauses.append(Product.product_id == product_id) # --- Region --- # if params.region_id and parent is not Region: q = q.join(Region) or_clauses.append(Region.region_id == params.region_id) if parent is Region: if params.country_code: or_clauses.append(Region.country_code == params.country_code.upper()) # --- Track --- # if params.track_id or params.isrc and parent is not Track: track_join = parent.tracks if hasattr(parent, 'tracks') else parent.track q = q.join(Track, track_join) if params.track_id: for track_id in QueryBuilder.param_as_list(params.track_id): or_clauses.append(Track.track_id == track_id) if params.isrc: for isrc in QueryBuilder.param_as_list(params.isrc): or_clauses.append(Track.isrc == isrc) return q.filter(or_(*or_clauses)) @classmethod def get_values_from_objs(cls, key: str, items: List[dict]): try: return list(set(map(operator.itemgetter(key), items))) except KeyError as e: return [] @classmethod def get_related_data_by_objs(cls, klass: Type[BaseModel], id_key: str, items: List[dict]) -> Query: """Pass a list of objects containing the primary keys of the ``klass`` you want queried""" ids = cls.get_values_from_objs(id_key, items) klass_id = getattr(klass, id_key) return db.session.query(klass).filter(klass_id.in_(ids))