import functools import logging from dataclasses import dataclass import anyio.from_thread import anyio.to_thread from anydi import singleton from anyio import CapacityLimiter from fansifter_common.auth.types import Permission from pydantic import TypeAdapter from snowflake.cortex import complete from dmp.adapters.db import ReportingDB from dmp.artists.cache import ArtistCache, cached_artist_data from dmp.artists.dtos import FansSegmentExplained, SegmentCategory, SegmentCrmCampaign from dmp.artists.prompts.get_artist_segments_explained_prompt import ( artist_segments_explained_prompt_template, ) from dmp.artists.requests import ArtistFanReportFilterRequest from dmp.artists.services import ArtistAccessService from dmp.config import Settings from dmp.fandata.enums import FanSegment, SegmentExplainedResponseLanguage logger = logging.getLogger(__name__) @dataclass class GetArtistSegmentsExplainedRequest(ArtistFanReportFilterRequest): language: SegmentExplainedResponseLanguage countries: list[str] | None = None FanSegmentExplainedListValidator = TypeAdapter(list[FansSegmentExplained]) FanSegmentCategoryListValidator = TypeAdapter(list[SegmentCategory]) FanSegmentCrmCampaignListValidator = TypeAdapter(list[SegmentCrmCampaign]) @singleton class GetArtistSegmentsExplainedHandler: permission = Permission("fan_data_list", "view") def __init__( self, db: ReportingDB, artist_access_service: ArtistAccessService, artist_cache: ArtistCache, settings: Settings, ) -> None: self.db = db self.artist_access_service = artist_access_service self.artist_cache = artist_cache self.settings = settings # allow up to 32 parallel blocking LLM calls self.limiter = CapacityLimiter(32) def handle( self, request: GetArtistSegmentsExplainedRequest ) -> list[FansSegmentExplained]: artist_access = self.artist_access_service.check_artist_access( request.identity_id, request.global_participant_id, request.account, permission=self.permission, ) segments_related_features = self._get_segments_related_features( request.global_participant_id, vendor_id=request.vendor_id, subaccount_id=request.subaccount_id, countries=request.countries, identity_id=request.identity_id, is_global=artist_access.is_global, ) segments_related_crm_campaigns = self._get_segments_related_crm_campaigns( request.global_participant_id, vendor_id=request.vendor_id, subaccount_id=request.subaccount_id, countries=request.countries, identity_id=request.identity_id, is_global=artist_access.is_global, ) if not segments_related_features: return [] return self._get_llm_response( request.global_participant_id, segments_related_features, segments_related_crm_campaigns, request.language, ) @cached_artist_data("segments_features") def _get_segments_related_features( self, global_participant_id: str, vendor_id: int | None, subaccount_id: int | None, countries: list[str] | None, identity_id: str, is_global: bool, ) -> list[SegmentCategory]: template_name = self._select_features_query_template_name(countries) query = self.db.query_from_template( template_name, context={ "vendor_id": vendor_id, "subaccount_id": subaccount_id, "global_participant_id": global_participant_id, "identity_id": identity_id, "countries": countries, "is_global": is_global, }, ) return FanSegmentCategoryListValidator.validate_python( self.db.session.execute(query).mappings() ) @cached_artist_data("segments_campaign") def _get_segments_related_crm_campaigns( self, global_participant_id: str, vendor_id: int | None, subaccount_id: int | None, countries: list[str] | None, identity_id: str, is_global: bool, ) -> list[SegmentCrmCampaign]: template_name = self._select_campaigns_query_template_name(countries) query = self.db.query_from_template( template_name, context={ "vendor_id": vendor_id, "subaccount_id": subaccount_id, "global_participant_id": global_participant_id, "identity_id": identity_id, "countries": countries, "is_global": is_global, }, ) return FanSegmentCrmCampaignListValidator.validate_python( self.db.session.execute(query).mappings() ) @cached_artist_data("segments_explained") def _get_llm_response( self, _: str, features: list[SegmentCategory], campaigns: list[SegmentCrmCampaign], language: SegmentExplainedResponseLanguage, ) -> list[FansSegmentExplained]: async def _call() -> list[FansSegmentExplained]: present_segment_names = {f.segment_name for f in features} present_segment_names.update( {c.segment_name for c in campaigns if c.segment_name == "NEW_FANS"} ) input_segments: list[FanSegment] = [ segment for segment in FanSegment if segment.value in present_segment_names ] if not input_segments: return [] prompts = [ ( segment, ( artist_segments_explained_prompt_template.replace( "[SEGMENT_RELATED_FEATURES]", self.inject_segments_related_features( segment.value, features ), ) .replace( "[SEGMENT_RELATED_CAMPAIGNS]", self.inject_segments_related_campaigns( segment.value, campaigns ), ) .replace("[LANGUAGE]", language.name) ), ) for segment in input_segments ] results: list[FansSegmentExplained] = [] async def call_llm(segment: FanSegment, prompt: str) -> None: try: segment_summary = await anyio.to_thread.run_sync( functools.partial( complete, model=self.settings.segment_explain_llm_model, prompt=prompt, session=self.db.create_snowpark_session(), options=self.settings.segment_explain_llm_options, timeout=self.settings.segment_explain_llm_timeout, ), limiter=self.limiter, ) results.append( FansSegmentExplained( segment=segment, explanation=segment_summary ) ) except Exception as e: logger.warning( "Calling Snowflake LLM.", extra={ "segment": segment, "prompt": prompt, "model": (self.settings.segment_explain_llm_model), "error": {e}, }, ) pass async with anyio.create_task_group() as tg: for segment, prompt in prompts: tg.start_soon(call_llm, segment, prompt) return results return anyio.from_thread.run(_call) @staticmethod def _select_features_query_template_name(countries: list[str] | None) -> str: if countries: return "artist/get-artist-features-by-segment-country.sql" return "artist/get-artist-features-by-segment.sql" @staticmethod def _select_campaigns_query_template_name(countries: list[str] | None) -> str: if countries: return "artist/get-artist-campaigns-by-segment-country.sql" return "artist/get-artist-campaigns-by-segment.sql" @staticmethod def inject_segments_related_features( segment_name: str, features: list[SegmentCategory] ) -> str: lines = [ f"{cat.category} {cat.fans_share}" for cat in features if cat.segment_name == segment_name ] return "\n".join(lines) @staticmethod def inject_segments_related_campaigns( segment_name: str, campaigns: list[SegmentCrmCampaign] ) -> str: lines = [ f"{cmp.campaign} {cmp.fans_share}" for cmp in campaigns if cmp.segment_name == segment_name ] return "\n".join(lines)