package routingcodes

import (
	"fmt"
	"regexp"
	"strings"

	"github.com/niels/banking-app/backend/internal/domain"
)

const (
	TypeABARoutingNumber = "aba_routing_number"
	TypeSortCode         = "sort_code"
	TypeBankCode         = "bank_code"
	TypeClearingCode     = "clearing_code"
	TypeBSB              = "bsb"
	TypeTransitNumber    = "transit_number"
	TypeInstitution      = "institution_number"
	TypeIFSC             = "ifsc"
	TypeBranchCode       = "branch_code"
)

var (
	countryPattern      = regexp.MustCompile(`^[A-Z]{2}$`)
	alphaNumericPattern = regexp.MustCompile(`^[A-Z0-9]+$`)
	clearingPattern     = regexp.MustCompile(`^[A-Z0-9-]+$`)
	ifscPattern         = regexp.MustCompile(`^[A-Z]{4}0[A-Z0-9]{6}$`)

	allowedTypes = map[string]struct{}{
		TypeABARoutingNumber: {},
		TypeSortCode:         {},
		TypeBankCode:         {},
		TypeClearingCode:     {},
		TypeBSB:              {},
		TypeTransitNumber:    {},
		TypeInstitution:      {},
		TypeIFSC:             {},
		TypeBranchCode:       {},
	}

	allowedNetworks = map[string]struct{}{
		"ach":     {},
		"fedwire": {},
		"fps":     {},
		"bacs":    {},
		"sepa":    {},
		"swift":   {},
		"local":   {},
		"eft":     {},
		"interac": {},
		"imps":    {},
		"rtgs":    {},
		"neft":    {},
	}
)

func NormalizeAndValidate(input []domain.RoutingCode) ([]domain.RoutingCode, error) {
	if len(input) > 12 {
		return nil, fmt.Errorf("%w: routing_codes may contain at most 12 entries", domain.ErrValidation)
	}

	seen := map[string]struct{}{}
	output := make([]domain.RoutingCode, 0, len(input))
	for _, raw := range input {
		code, err := NormalizeAndValidateOne(raw)
		if err != nil {
			return nil, err
		}
		key := strings.Join([]string{code.Country, code.Network, code.CodeType, code.Code}, "|")
		if _, exists := seen[key]; exists {
			return nil, fmt.Errorf("%w: duplicate routing code", domain.ErrValidation)
		}
		seen[key] = struct{}{}
		output = append(output, code)
	}
	return output, nil
}

func InferFromIBAN(iban string) []domain.RoutingCode {
	normalized := strings.ToUpper(strings.ReplaceAll(strings.TrimSpace(iban), " ", ""))
	if len(normalized) >= 8 && strings.HasPrefix(normalized, "NL") {
		return []domain.RoutingCode{
			{CodeType: TypeBankCode, Country: "NL", Network: "sepa", Code: normalized[4:8]},
		}
	}
	return nil
}

func MergeInferred(explicit, inferred []domain.RoutingCode) []domain.RoutingCode {
	if len(inferred) == 0 {
		return explicit
	}
	merged := append([]domain.RoutingCode{}, explicit...)
	existing := map[string]struct{}{}
	for _, code := range explicit {
		key := strings.Join([]string{
			strings.ToUpper(strings.TrimSpace(code.Country)),
			strings.ToLower(strings.TrimSpace(code.Network)),
			strings.ToLower(strings.TrimSpace(code.CodeType)),
		}, "|")
		existing[key] = struct{}{}
	}
	for _, code := range inferred {
		key := strings.Join([]string{
			strings.ToUpper(strings.TrimSpace(code.Country)),
			strings.ToLower(strings.TrimSpace(code.Network)),
			strings.ToLower(strings.TrimSpace(code.CodeType)),
		}, "|")
		if _, ok := existing[key]; ok {
			continue
		}
		merged = append(merged, code)
		existing[key] = struct{}{}
	}
	return merged
}

func NormalizeAndValidateOne(raw domain.RoutingCode) (domain.RoutingCode, error) {
	code := domain.RoutingCode{
		CodeType: strings.ToLower(strings.TrimSpace(raw.CodeType)),
		Country:  strings.ToUpper(strings.TrimSpace(raw.Country)),
		Network:  strings.ToLower(strings.TrimSpace(raw.Network)),
		Code:     strings.ToUpper(strings.TrimSpace(raw.Code)),
	}
	code.Code = strings.ReplaceAll(code.Code, " ", "")

	if _, ok := allowedTypes[code.CodeType]; !ok {
		return domain.RoutingCode{}, fmt.Errorf("%w: routing code type is invalid", domain.ErrValidation)
	}
	if !countryPattern.MatchString(code.Country) {
		return domain.RoutingCode{}, fmt.Errorf("%w: routing code country must be ISO-3166 alpha-2", domain.ErrValidation)
	}
	if _, ok := allowedNetworks[code.Network]; !ok {
		return domain.RoutingCode{}, fmt.Errorf("%w: routing code network is invalid", domain.ErrValidation)
	}

	switch code.CodeType {
	case TypeABARoutingNumber:
		code.Code = strings.ReplaceAll(code.Code, "-", "")
		if code.Country != "US" {
			return domain.RoutingCode{}, fmt.Errorf("%w: aba_routing_number requires country US", domain.ErrValidation)
		}
		if code.Network != "ach" && code.Network != "fedwire" {
			return domain.RoutingCode{}, fmt.Errorf("%w: aba_routing_number requires ach or fedwire network", domain.ErrValidation)
		}
		if !isDigits(code.Code, 9) || !validABAChecksum(code.Code) {
			return domain.RoutingCode{}, fmt.Errorf("%w: aba_routing_number checksum is invalid", domain.ErrValidation)
		}
	case TypeSortCode:
		digits := strings.ReplaceAll(code.Code, "-", "")
		if code.Country != "GB" {
			return domain.RoutingCode{}, fmt.Errorf("%w: sort_code requires country GB", domain.ErrValidation)
		}
		if code.Network != "fps" && code.Network != "bacs" {
			return domain.RoutingCode{}, fmt.Errorf("%w: sort_code requires fps or bacs network", domain.ErrValidation)
		}
		if !isDigits(digits, 6) {
			return domain.RoutingCode{}, fmt.Errorf("%w: sort_code must be 6 digits", domain.ErrValidation)
		}
		code.Code = digits[0:2] + "-" + digits[2:4] + "-" + digits[4:6]
	case TypeBSB:
		digits := strings.ReplaceAll(code.Code, "-", "")
		if code.Country != "AU" {
			return domain.RoutingCode{}, fmt.Errorf("%w: bsb requires country AU", domain.ErrValidation)
		}
		if code.Network != "local" {
			return domain.RoutingCode{}, fmt.Errorf("%w: bsb requires local network", domain.ErrValidation)
		}
		if !isDigits(digits, 6) {
			return domain.RoutingCode{}, fmt.Errorf("%w: bsb must be 6 digits", domain.ErrValidation)
		}
		code.Code = digits[0:3] + "-" + digits[3:6]
	case TypeTransitNumber:
		if code.Country != "CA" {
			return domain.RoutingCode{}, fmt.Errorf("%w: transit_number requires country CA", domain.ErrValidation)
		}
		if code.Network != "eft" && code.Network != "interac" {
			return domain.RoutingCode{}, fmt.Errorf("%w: transit_number requires eft or interac network", domain.ErrValidation)
		}
		if !isDigits(code.Code, 5) {
			return domain.RoutingCode{}, fmt.Errorf("%w: transit_number must be 5 digits", domain.ErrValidation)
		}
	case TypeInstitution:
		if code.Country != "CA" {
			return domain.RoutingCode{}, fmt.Errorf("%w: institution_number requires country CA", domain.ErrValidation)
		}
		if code.Network != "eft" && code.Network != "interac" {
			return domain.RoutingCode{}, fmt.Errorf("%w: institution_number requires eft or interac network", domain.ErrValidation)
		}
		if !isDigits(code.Code, 3) {
			return domain.RoutingCode{}, fmt.Errorf("%w: institution_number must be 3 digits", domain.ErrValidation)
		}
	case TypeIFSC:
		if code.Country != "IN" {
			return domain.RoutingCode{}, fmt.Errorf("%w: ifsc requires country IN", domain.ErrValidation)
		}
		if code.Network != "imps" && code.Network != "rtgs" && code.Network != "neft" {
			return domain.RoutingCode{}, fmt.Errorf("%w: ifsc requires imps, rtgs or neft network", domain.ErrValidation)
		}
		if !ifscPattern.MatchString(code.Code) {
			return domain.RoutingCode{}, fmt.Errorf("%w: ifsc format is invalid", domain.ErrValidation)
		}
	case TypeBankCode:
		if code.Network != "sepa" && code.Network != "swift" && code.Network != "local" {
			return domain.RoutingCode{}, fmt.Errorf("%w: bank_code requires sepa, swift or local network", domain.ErrValidation)
		}
		if len(code.Code) < 2 || len(code.Code) > 12 || !alphaNumericPattern.MatchString(code.Code) {
			return domain.RoutingCode{}, fmt.Errorf("%w: bank_code must be 2 to 12 letters or digits", domain.ErrValidation)
		}
	case TypeClearingCode, TypeBranchCode:
		if len(code.Code) < 2 || len(code.Code) > 20 || !clearingPattern.MatchString(code.Code) {
			return domain.RoutingCode{}, fmt.Errorf("%w: routing code must be 2 to 20 letters, digits or hyphens", domain.ErrValidation)
		}
	}

	return code, nil
}

func isDigits(value string, length int) bool {
	if len(value) != length {
		return false
	}
	for _, r := range value {
		if r < '0' || r > '9' {
			return false
		}
	}
	return true
}

func validABAChecksum(value string) bool {
	weights := []int{3, 7, 1, 3, 7, 1, 3, 7, 1}
	sum := 0
	for i, r := range value {
		sum += int(r-'0') * weights[i]
	}
	return sum%10 == 0
}
