from .interface import AnalyticsModuleInterface from .params import MultiParam, Param, ParamSchema class AttributeRFMModelScoreComponents(AnalyticsModuleInterface): description = """ Calculate average/share of each Superfan model attribute in each segment created by Superfan model """ inputs = [ ParamSchema, Param("field_name", list), MultiParam("attr_", ptype=int), Param("attributes", tuple), Param("attribute_labels", dict), Param("collection_id", required=False), Param("collection_ids", list, required=False), ] outputs = [Param("group", idx=0), Param("label", idx=1), Param("value", int, idx=2)] @staticmethod def get_query( # type: ignore schema: str, attributes: tuple, collection_id: str, collection_ids: list, field_name: list = None, **kwargs, ) -> str: result_query = f""" WITH fan_segment AS ( SELECT fa.fan_id, fa.value AS segment FROM {schema}.fan_attribute fa JOIN {schema}.attribute a ON a.id = fa.attribute_id WHERE {f"fa.collection_id IN ({', '.join(map(str, collection_ids))}) AND" if collection_ids else ""} {f"fa.fan_id IN ( SELECT fan_id FROM {schema}.collection_fan WHERE collection_id = %(collection_id)s ) AND" if collection_id else ""} a.name IN ('enrUserRFM')), rfm_attributes AS ( SELECT fa.fan_id, a.name AS attribute_name, MAX(fa.value::numeric) AS value FROM {schema}.fan_attribute fa JOIN {schema}.attribute a ON a.id = fa.attribute_id WHERE {f"fa.collection_id IN ({', '.join(map(str, collection_ids))}) AND" if collection_ids else ""} {f"fa.fan_id IN ( SELECT fan_id FROM {schema}.collection_fan WHERE collection_id = %(collection_id)s ) AND" if collection_id else ""} a.name IN %(attributes)s GROUP BY fa.fan_id, a.name ) SELECT fan_segment.segment, rfm_attributes.attribute_name, ROUND(AVG(rfm_attributes.value), 1) AS value, 'numeric' AS type FROM rfm_attributes JOIN fan_segment ON fan_segment.fan_id = rfm_attributes.fan_id GROUP BY fan_segment.segment, rfm_attributes.attribute_name ORDER BY segment; """ return result_query def _apply_sort_list(self, result): if self.params.get("sort_list") is not None: idx = self.get_output_map()[self.params["sort_column"]].idx result.sort(key=lambda x: self.params.get("sort_list").index(x[idx])) return result @staticmethod def populate_attribute_ids(conf, available_fields_list): field_names = conf["inputs"]["field_name"] attribute_ids = {} for field_name in field_names: if field_name in available_fields_list: attribute_ids[f"attr_{field_name}"] = available_fields_list.get( field_name ) return attribute_ids