"""Repository for formal risk assessments and internal profile projections.""" from datetime import datetime from typing import cast from sqlalchemy import func, select, update from sqlalchemy.ext.asyncio import AsyncSession from app.model.fund import FundRiskAssessment as RiskAssessment from app.model.memory import MemorySyncOutbox from app.model.profile_tag import AdvisorProfileDriftReview, AdvisorProfileTag from app.model.risk_questionnaire import ProfileSnapshot class RiskQuestionnaireRepository: def __init__(self, session: AsyncSession) -> None: self.session = session async def has_valid_assessment(self, customer_id: int, now: datetime) -> bool: return await self.session.scalar( select(RiskAssessment.id) .where(RiskAssessment.customer_id == customer_id, RiskAssessment.valid_until > now) .limit(1) ) is not None async def assessments_since(self, customer_id: int, since: datetime) -> int: value = await self.session.scalar( select(func.count()) .select_from(RiskAssessment) .where(RiskAssessment.customer_id == customer_id, RiskAssessment.assessed_at >= since) ) return int(value or 0) async def next_profile_version(self, customer_id: int) -> int: value = await self.session.scalar( select(func.coalesce(func.max(ProfileSnapshot.version), 0)).where( ProfileSnapshot.customer_id == customer_id ) ) return int(value or 0) + 1 async def deactivate_current_profile(self, customer_id: int, now: datetime) -> None: await self.session.execute( update(ProfileSnapshot) .where(ProfileSnapshot.customer_id == customer_id, ProfileSnapshot.is_current.is_(True)) .values(is_current=False, updated_at=now) ) def add_assessment(self, assessment: RiskAssessment) -> None: self.session.add(assessment) def add_profile(self, profile: ProfileSnapshot) -> None: self.session.add(profile) def add_sync_event(self, event: MemorySyncOutbox) -> None: self.session.add(event) async def latest_assessment(self, customer_id: int) -> RiskAssessment | None: return cast(RiskAssessment | None, await self.session.scalar( select(RiskAssessment) .where(RiskAssessment.customer_id == customer_id) .order_by(RiskAssessment.assessed_at.desc(), RiskAssessment.id.desc()) .limit(1) )) async def current_profile(self, customer_id: int) -> ProfileSnapshot | None: return cast(ProfileSnapshot | None, await self.session.scalar( select(ProfileSnapshot) .where( ProfileSnapshot.customer_id == customer_id, ProfileSnapshot.is_current.is_(True), ) .order_by(ProfileSnapshot.version.desc(), ProfileSnapshot.id.desc()) .limit(1) )) async def active_tags( self, customer_id: int, *, lock: bool = False ) -> list[AdvisorProfileTag]: statement = ( select(AdvisorProfileTag) .where( AdvisorProfileTag.customer_id == customer_id, AdvisorProfileTag.status == "active", ) .order_by(AdvisorProfileTag.tag_key.asc(), AdvisorProfileTag.id.desc()) ) if lock: statement = statement.with_for_update() return list(await self.session.scalars(statement)) async def tags(self, customer_id: int, *, limit: int = 100) -> list[AdvisorProfileTag]: return list(await self.session.scalars( select(AdvisorProfileTag) .where(AdvisorProfileTag.customer_id == customer_id) .order_by(AdvisorProfileTag.tag_key.asc(), AdvisorProfileTag.id.desc()) .limit(limit) )) async def pending_drift_review( self, customer_id: int, *, lock: bool = False ) -> AdvisorProfileDriftReview | None: statement = ( select(AdvisorProfileDriftReview) .where( AdvisorProfileDriftReview.customer_id == customer_id, AdvisorProfileDriftReview.status == "pending_review", ) .order_by( AdvisorProfileDriftReview.created_at.desc(), AdvisorProfileDriftReview.id.desc(), ) .limit(1) ) if lock: statement = statement.with_for_update() return cast(AdvisorProfileDriftReview | None, await self.session.scalar(statement)) async def drift_review( self, review_id: int, *, lock: bool = False ) -> AdvisorProfileDriftReview | None: statement = select(AdvisorProfileDriftReview).where( AdvisorProfileDriftReview.id == review_id ) if lock: statement = statement.with_for_update() return cast(AdvisorProfileDriftReview | None, await self.session.scalar(statement)) async def pending_reviews(self, *, limit: int) -> list[AdvisorProfileDriftReview]: return list(await self.session.scalars( select(AdvisorProfileDriftReview) .where(AdvisorProfileDriftReview.status == "pending_review") .order_by( AdvisorProfileDriftReview.created_at.asc(), AdvisorProfileDriftReview.id.asc(), ) .limit(limit) )) async def tags_for_review( self, review_id: int, *, lock: bool = False ) -> list[AdvisorProfileTag]: statement = ( select(AdvisorProfileTag) .where(AdvisorProfileTag.drift_review_id == review_id) .order_by(AdvisorProfileTag.tag_key.asc(), AdvisorProfileTag.id.asc()) ) if lock: statement = statement.with_for_update() return list(await self.session.scalars(statement)) async def supersede_active_tags( self, customer_id: int, tag_keys: tuple[str, ...], now: datetime ) -> None: if not tag_keys: return await self.session.execute( update(AdvisorProfileTag) .where( AdvisorProfileTag.customer_id == customer_id, AdvisorProfileTag.tag_key.in_(tag_keys), AdvisorProfileTag.status == "active", ) .values(status="superseded", active_customer_tag=None, updated_at=now) ) def add_tag(self, tag: AdvisorProfileTag) -> None: self.session.add(tag) def add_drift_review(self, review: AdvisorProfileDriftReview) -> None: self.session.add(review)