import logging
from typing import Optional

from sqlalchemy import select
from sqlalchemy.orm import Session

from src.apps.pricing_control.models.merchant_pricing_assignment import MerchantPricingAssignment
from src.apps.pricing_control.models.pricing_template import PricingType
from src.apps.pricing_control.models.pricing_template_rate import PricingTemplateRate
from src.apps.pricing_control.crud import get_active_assignment

logger = logging.getLogger(__name__)


class PricingService:
    @staticmethod
    def calculate_fee(
        db: Session,
        merchant_id: int,
        transaction_type: str,
        amount_cents: int,
        card_network: Optional[str] = None,
    ) -> int:
        """
        Returns platform fee in cents for the given transaction.
        Returns 0 if no active pricing assignment, no matching rate, or any error.

        Two-pass lookup when card_network is provided:
          1. Try (template_id, transaction_type, card_network = <network>)
          2. Fall back to (template_id, transaction_type, card_network IS NULL)
        """
        try:
            assignment = get_active_assignment(db, merchant_id)
            if not assignment:
                return 0

            template = assignment.template
            if not template:
                return 0

            rate = None

            # Pass 1 — network-specific rate
            if card_network:
                rate = db.execute(
                    select(PricingTemplateRate).where(
                        PricingTemplateRate.template_id == template.id,
                        PricingTemplateRate.transaction_type == transaction_type,
                        PricingTemplateRate.card_network == card_network,
                        PricingTemplateRate.deleted_at.is_(None),
                    )
                ).scalar_one_or_none()

            # Pass 2 — generic fallback (card_network IS NULL)
            if rate is None:
                rate = db.execute(
                    select(PricingTemplateRate).where(
                        PricingTemplateRate.template_id == template.id,
                        PricingTemplateRate.transaction_type == transaction_type,
                        PricingTemplateRate.card_network.is_(None),
                        PricingTemplateRate.deleted_at.is_(None),
                    )
                ).scalar_one_or_none()

            if not rate:
                return 0

            pricing_type = template.pricing_type

            if pricing_type == PricingType.INTERCHANGE_PLUS:
                ibp = rate.interchange_basis_points or 0
                mbp = rate.markup_basis_points or 0
                percentage_fee = round(amount_cents * (ibp + mbp) / 10000)
            else:
                # flat_rate, tiered, custom all use rate_percentage
                pct = rate.rate_percentage or 0.0
                percentage_fee = round(amount_cents * pct / 100)

            return percentage_fee + (rate.fixed_fee_cents or 0)

        except Exception as exc:
            logger.warning("PricingService.calculate_fee failed for merchant %s: %s", merchant_id, exc)
            return 0
