from abc import ABC, abstractmethod from typing import Any, Dict, List, Mapping, Iterable, Type from sqlalchemy import values, column, Integer, String from sqlalchemy.orm import Query from server.db.constants import ALL_DSPS from server.db.es_connector import ES_INSTANCE from server.dna.utils import DataQueryWrapper from .params_parser import ParamsParser __all__ = ["DiscoveryHelper"] class DiscoveryHelper(ABC): _es_connector = ES_INSTANCE _params_parser_cls: Type[ParamsParser] = ParamsParser @property @abstractmethod def _index_name(self) -> str: pass @property @abstractmethod def _data_model(self): pass @property @abstractmethod def _dsp_sort_mapping(self) -> Mapping[str, str]: pass @property @abstractmethod def _dsp_default_charts(self) -> Mapping[str, str]: pass def _get_query(self, params: Dict[str, Any]) -> Dict[str, Any]: return {"bool": {"must": list(self._params_parser_cls.parse(params))}} def _get_chart_types(self, params: Dict[str, Any]) -> Dict[str, str]: chart_types = {} for dsp in ALL_DSPS: chart_types[dsp] = params.get(dsp, {}).get("chart_type", self._dsp_default_charts[dsp]).lower() return chart_types @abstractmethod def _get_sort(self, params: Dict[str, Any]) -> Dict[str, Any]: pass async def _get_results(self, params: Dict[str, Any], *, limit: int, offset: int) -> List[str]: query = self._get_query(params) sort = self._get_sort(params) _, results, _ = await self._es_connector.search( index=self._index_name, query=query, sort=sort, from_=offset, size=limit ) return [row["_id"] for row in results] async def _get_data_by_ids(self, ids: List[str], *, ordering) -> List[Dict[str, Any]]: query = Query(self._data_model.value).where(self._data_model.id.in_(ids)) if ordering: ordering = values(column("index", Integer()), column("id", String()), name="ordering").data( [(i, id_) for i, id_ in enumerate(ids)] ) query = query.join(ordering, self._data_model.id == ordering.c.id).order_by(ordering.c.index) return list(await DataQueryWrapper.select(query, many=True)) async def _get_data_by_id(self, id_: str) -> Dict[str, Any]: query = Query(self._data_model.value).where(self._data_model.id == id_) return await DataQueryWrapper.select(query) def _hydrate_search( self, results: List[Dict[str, Any]], *, chart_types: Dict[str, str] ) -> Iterable[Dict[str, Any]]: for row in results: for dsp, chart_type in chart_types.items(): if dsp not in row: continue row[dsp]["chart_data"] = row[dsp][f"{chart_type}_chart"] row[dsp]["chart_type"] = chart_type yield row async def get_data(self, ids: List[str], *, params: Dict[str, Any], ordering: bool = False) -> List[Dict[str, Any]]: results = await self._get_data_by_ids(ids, ordering=ordering) chart_types = self._get_chart_types(params) return list(self._hydrate_search(results, chart_types=chart_types)) async def search(self, params: Dict[str, Any]) -> List[Dict[str, Any]]: limit, offset = params.pop("limit"), params.pop("offset") ids = await self._get_results(params, limit=limit, offset=offset) if not ids: return [] return await self.get_data(ids, params=params, ordering=True) async def count(self, params: Dict[str, Any]) -> int: query = self._get_query(params) result = await self._es_connector.count(index=self._index_name, query=query) return result["count"] def _hydrate_get(self, data: Dict[str, Any]) -> Dict[str, Any]: return data.copy() async def get(self, id_: str) -> Dict[str, Any]: data = await self._get_data_by_id(id_) if not data: return {} return self._hydrate_get(data)