package domain

import (
	"fmt"
	"math/big"
)

const defaultMinorUnit = 2

type RoundingMode string

const (
	RoundingModeFloor    RoundingMode = "floor"
	RoundingModeHalfUp   RoundingMode = "half_up"
	RoundingModeHalfEven RoundingMode = "half_even"
)

type RoundingPolicy struct {
	Currency                        string
	MinorUnit                       int
	Mode                            RoundingMode
	CashRoundingIncrementMinorUnits int64
}

func DefaultRoundingPolicy(currency string) RoundingPolicy {
	currency = NormalizeCurrency(currency)
	minorUnit := defaultMinorUnit
	switch currency {
	case "BHD", "JOD", "KWD", "OMR", "TND":
		minorUnit = 3
	case "CLP", "ISK", "JPY", "KRW", "PYG", "VND":
		minorUnit = 0
	}
	return RoundingPolicy{
		Currency:                        currency,
		MinorUnit:                       minorUnit,
		Mode:                            RoundingModeHalfUp,
		CashRoundingIncrementMinorUnits: 1,
	}
}

func NormalizeRoundingPolicy(policy RoundingPolicy) (RoundingPolicy, error) {
	policy.Currency = NormalizeCurrency(policy.Currency)
	if err := ValidateCurrency(policy.Currency); err != nil {
		return RoundingPolicy{}, err
	}
	if policy.MinorUnit < 0 || policy.MinorUnit > 4 {
		return RoundingPolicy{}, fmt.Errorf("%w: minor_unit must be between 0 and 4", ErrValidation)
	}
	switch policy.Mode {
	case RoundingModeFloor, RoundingModeHalfUp, RoundingModeHalfEven:
	default:
		return RoundingPolicy{}, fmt.Errorf("%w: invalid rounding mode", ErrValidation)
	}
	if policy.CashRoundingIncrementMinorUnits <= 0 {
		policy.CashRoundingIncrementMinorUnits = 1
	}
	return policy, nil
}

func ConvertMinorUnits(amountMinor int64, rateMicros int64, fromPolicy, toPolicy RoundingPolicy) (int64, error) {
	if err := ValidateAmount(amountMinor); err != nil {
		return 0, err
	}
	if rateMicros <= 0 {
		return 0, fmt.Errorf("%w: rate_micros must be greater than zero", ErrValidation)
	}
	fromPolicy, err := NormalizeRoundingPolicy(fromPolicy)
	if err != nil {
		return 0, err
	}
	toPolicy, err = NormalizeRoundingPolicy(toPolicy)
	if err != nil {
		return 0, err
	}

	numerator := big.NewInt(amountMinor)
	numerator.Mul(numerator, big.NewInt(rateMicros))
	numerator.Mul(numerator, big.NewInt(pow10(toPolicy.MinorUnit)))

	denominator := big.NewInt(1_000_000)
	denominator.Mul(denominator, big.NewInt(pow10(fromPolicy.MinorUnit)))

	return RoundRationalMinorUnits(numerator, denominator, toPolicy)
}

func RoundRationalMinorUnits(numerator, denominator *big.Int, policy RoundingPolicy) (int64, error) {
	if numerator == nil || denominator == nil || numerator.Sign() < 0 || denominator.Sign() <= 0 {
		return 0, fmt.Errorf("%w: invalid rounding ratio", ErrValidation)
	}
	policy, err := NormalizeRoundingPolicy(policy)
	if err != nil {
		return 0, err
	}

	quotient, remainder := new(big.Int), new(big.Int)
	quotient.QuoRem(numerator, denominator, remainder)
	if shouldRoundUp(quotient, remainder, denominator, policy.Mode) {
		quotient.Add(quotient, big.NewInt(1))
	}
	if !quotient.IsInt64() {
		return 0, fmt.Errorf("%w: rounded amount exceeds supported maximum", ErrValidation)
	}
	rounded := quotient.Int64()
	if policy.CashRoundingIncrementMinorUnits > 1 {
		rounded, err = roundToIncrement(rounded, policy.CashRoundingIncrementMinorUnits, policy.Mode)
		if err != nil {
			return 0, err
		}
	}
	if rounded <= 0 {
		return 0, fmt.Errorf("%w: rounded amount is too small", ErrValidation)
	}
	return rounded, nil
}

func shouldRoundUp(quotient, remainder, denominator *big.Int, mode RoundingMode) bool {
	if remainder.Sign() == 0 || mode == RoundingModeFloor {
		return false
	}
	twiceRemainder := new(big.Int).Mul(remainder, big.NewInt(2))
	cmp := twiceRemainder.Cmp(denominator)
	if cmp > 0 {
		return true
	}
	if cmp < 0 {
		return false
	}
	if mode == RoundingModeHalfUp {
		return true
	}
	return quotient.Bit(0) == 1
}

func roundToIncrement(amount, increment int64, mode RoundingMode) (int64, error) {
	if amount < 0 || increment <= 0 {
		return 0, fmt.Errorf("%w: invalid cash rounding increment", ErrValidation)
	}
	quotient := amount / increment
	remainder := amount % increment
	if remainder == 0 || mode == RoundingModeFloor {
		return quotient * increment, nil
	}
	cmp := remainder * 2
	if cmp > increment || (cmp == increment && (mode == RoundingModeHalfUp || quotient%2 == 1)) {
		quotient++
	}
	return quotient * increment, nil
}

func pow10(exp int) int64 {
	value := int64(1)
	for i := 0; i < exp; i++ {
		value *= 10
	}
	return value
}
