package dnsresolver

import (
	"fmt"
	"net"
	"regexp"
	"strings"
	"time"

	"github.com/aws/aws-sdk-go/aws"
	"github.com/aws/aws-sdk-go/aws/session"
	"github.com/aws/aws-sdk-go/service/route53"
	"github.com/filtr/infrastructure/go/openvpn-route-updater/internal/config"
	"github.com/filtr/infrastructure/go/openvpn-route-updater/internal/logger"
	"github.com/filtr/infrastructure/go/openvpn-route-updater/internal/utility"
	"github.com/pkg/errors"
)

// Try resolving roun robin DNS N time
// Verify that all ips from the list are resolved
const roundRobinResolveAttemptsLimit = 100
const roundRobinVerifyAttempts = 3

// Magic nubmers for dealing with
// RDS RO endpoint behavior
const rdsRequestPeriod = 500
const rdsVerifyAttempts = 20

// Datastruct for key/value
type Registry struct {
	ResolvedIPs map[string][]string
}

func NewRegistry() Registry {
	var c Registry
	c.ResolvedIPs = make(map[string][]string)
	return c
}

func CopyRegistry(original Registry, ipKeys []string) Registry {
	c := NewRegistry()
	if ipKeys == nil {
		for k, v := range original.ResolvedIPs {
			c.ResolvedIPs[k] = v
		}
	} else {
		for _, ipKey := range ipKeys {
			if v, ok := original.ResolvedIPs[ipKey]; ok {
				c.ResolvedIPs[ipKey] = v
			}
		}
	}
	return c
}

func MergeRegistries(r1, r2 Registry) Registry {
	reg := NewRegistry()
	for k, v := range r1.ResolvedIPs {
		reg.ResolvedIPs[k] = utility.Unique(v)
	}
	for k, v := range r2.ResolvedIPs {
		reg.ResolvedIPs[k] = utility.Unique(append(reg.ResolvedIPs[k], v...))
	}
	return reg
}

//AddRecord Add new Unique records to the existing IP key, or add a new IP key
func (c *Registry) AddRecord(ip string, record string) {
	if _, ok := c.ResolvedIPs[ip]; ok {
		c.ResolvedIPs[ip] = utility.Unique(append(c.ResolvedIPs[ip], record))
	} else {
		c.ResolvedIPs[ip] = []string{record}
	}
}

func (c *Registry) RemoveIP(ip string) {
	delete(c.ResolvedIPs, ip)
}

func (c *Registry) RemoveRecord(record string) {
	for key, records := range c.ResolvedIPs {
		recordExist := false
		// Record list should be unique for each IP, so only one remove index is presented
		removeIndex := 0
		for i, r := range records {
			if r == record {
				recordExist = true
				removeIndex = i
				break
			}
		}
		if recordExist {
			c.ResolvedIPs[key] = utility.Remove(c.ResolvedIPs[key], removeIndex)
		}
	}
}

func (c *Registry) RemoveRecordForIP(ip string, record string) {
	for key, records := range c.ResolvedIPs {
		if key == ip {
			recordExist := false
			// Record list should be unique for each IP, so only one remove index is presented
			removeIndex := 0
			for i, r := range records {
				if r == record {
					recordExist = true
					removeIndex = i
					break
				}
			}
			if recordExist {
				c.ResolvedIPs[key] = utility.Remove(c.ResolvedIPs[key], removeIndex)
			}
		}
	}
}

func (c Registry) GetIPList() []net.IP {
	var res []net.IP
	for k := range c.ResolvedIPs {
		res = append(res, net.ParseIP(k).To4())
	}
	return res
}

func (c Registry) GetRecordList() []string {
	var res []string
	for _, v := range c.ResolvedIPs {
		res = utility.Unique(append(res, v...))
	}
	return res
}

func (c *Registry) RemovePrivateAddressesIPv4(log logger.Logger) error {
	var keysToRemove []string
	for key := range c.ResolvedIPs {
		ip := net.ParseIP(key).To4()
		isPrivate, err := CheckIPPrivate(ip)
		if err != nil {
			return err
		}
		if isPrivate {
			keysToRemove = append(keysToRemove, key)
		} else {
			log.Debugf("%v was filtered out", key)
		}
	}
	for _, key := range keysToRemove {
		delete(c.ResolvedIPs, key)
	}
	return nil
}

// FilterRecordSets - filter given record set to include only records which matches recordFilter parameter for specific hosted zone
func FilterRecordSets(c *config.GroupConfigManager, zoneName string, recordSets []*route53.ResourceRecordSet) []*route53.ResourceRecordSet {
	var resultingSet []*route53.ResourceRecordSet
	for _, recordSet := range recordSets {
		if c.CheckRecordType(*recordSet.Type) && c.CheckRecordName(*recordSet.Name, zoneName) {
			resultingSet = append(resultingSet, recordSet)
		}
	}
	return resultingSet
}

// GetRoute53ZoneRecords - get all route53 zone records sets
func GetRoute53ZoneRecords(zoneID string, svc *route53.Route53) ([]*route53.ResourceRecordSet, error) {
	listParams := &route53.ListResourceRecordSetsInput{
		HostedZoneId: aws.String(zoneID), // Required
	}

	var res []*route53.ResourceRecordSet
	err := svc.ListResourceRecordSetsPages(listParams,
		func(page *route53.ListResourceRecordSetsOutput, lastPage bool) bool {
			res = append(res, page.ResourceRecordSets...)
			return true
		})
	if err != nil {
		return nil, errors.Wrap(err, "Unable to get list of records in the Route53 zone")
	}
	return res, nil
}

// GetZoneIDByName get AWS Route 53 zone id by its name
func GetZoneIDByName(log logger.Logger, svc *route53.Route53, zoneName string) (string, error) {
	zoneList, err := svc.ListHostedZones(&route53.ListHostedZonesInput{})
	if err != nil {
		return "", errors.Wrapf(err, "Unable to get list of hosted zones %v", svc)
	}
	for _, zone := range zoneList.HostedZones {
		if *zone.Name == zoneName {
			return *zone.Id, nil
		}
	}
	return "", errors.New(fmt.Sprintf("Zone hasn't been found %v", svc))
}

// GetRoute53RecordsByGroupConfig - returns list of records for zones from the config with applied filters.Creates its own AWS client
func GetRoute53RecordsByGroupConfig(log logger.Logger, c *config.GroupConfigManager) ([]*route53.ResourceRecordSet, error) {
	var DNSRecords []*route53.ResourceRecordSet

	for _, zone := range c.Group.Zones {
		sess, err := session.NewSessionWithOptions(session.Options{
			Profile:           zone.AwsProfile,
			SharedConfigState: session.SharedConfigEnable,
		})
		if err != nil {
			return nil, err
		}
		svc := route53.New(sess)

		zoneID, err := GetZoneIDByName(log, svc, zone.Name)
		if err != nil {
			return nil, err
		}
		zoneRecords, err := GetRoute53ZoneRecords(zoneID, svc)
		if err != nil {
			return nil, err
		}
		DNSRecords = append(DNSRecords, FilterRecordSets(c, zone.Name, zoneRecords)...)
	}
	log.Debug(DNSRecords)
	return DNSRecords, nil
}

// ResolveRoute53Records - gets a list of recordsets from Route53 and resolve it to list of IP addresses
func ResolveRoute53Records(log logger.Logger, recordSets []*route53.ResourceRecordSet) Registry {
	registry := NewRegistry()
	for _, recordSet := range recordSets {
		if recordSet.ResourceRecords != nil {
			for _, record := range recordSet.ResourceRecords {
				registry.AddRecord(net.ParseIP(*record.Value).String(), *recordSet.Name)
			}
		}
		if recordSet.AliasTarget != nil {
			//ignore cloudfront records
			if strings.Contains(*recordSet.AliasTarget.DNSName, ".cloudfront.net") {
				log.Debugf("Ignoring cloudfront record %v", *recordSet.AliasTarget.DNSName)
			} else {
				addrs, err := ResolveHost(log, *recordSet.AliasTarget.DNSName)
				if err != nil {
					log.Warnf("Unable to resolve record %v", *recordSet.AliasTarget.DNSName)
				} else {
					for _, addr := range addrs {
						registry.AddRecord(net.ParseIP(addr).String(), *recordSet.Name)
					}
				}
			}
		}
	}
	return registry
}

// ResolveCustomRecords - resolves list of records to ip addresses
func ResolveCustomRecords(log logger.Logger, records []string) Registry {
	registry := NewRegistry()
	for _, record := range records {
		addrs, err := ResolveHost(log, record)
		if err != nil {
			log.Warnf("Unable to resolve record %v", record)
		} else {
			for _, addr := range addrs {
				registry.AddRecord(net.ParseIP(addr).String(), record)
			}
		}
	}
	return registry
}

// ResolveHost - resolve host ipv4 addresses. Even if LookupHost returns only one record per run, it will check that all round robin addresses will be returned
func ResolveHost(log logger.Logger, host string) ([]string, error) {
	var resultAddrs []string
	// Due to the different resolve mechanics in diffirent OS
	// Magic Number, maximum amounts of attempts to resolve dns record
	// If we haven't got all records for hostname after 100 attemps, something is clearly wrong
	limit := roundRobinResolveAttemptsLimit
	for proceed := true; proceed && limit != 0; limit-- {
		addrs, err := net.LookupHost(host)
		if err != nil {
			return nil, err
		}
		oldAddrLen := len(resultAddrs)
		resultAddrs = append(resultAddrs, addrs...)
		resultAddrs = utility.Unique(resultAddrs)
		if len(resultAddrs) == oldAddrLen {
			proceed = false
		}
	}
	if limit == 0 {
		return nil, fmt.Errorf("can't correctly resolve hostname %v", host)
	}
	// Verify round robin DNS

	var verifyLimit int
	if len(resultAddrs) >= roundRobinVerifyAttempts {
		verifyLimit = len(resultAddrs)
	} else {
		verifyLimit = roundRobinVerifyAttempts
	}

	for i := 0; i < verifyLimit; i++ {
		addrs, err := net.LookupHost(host)
		if err != nil {
			return nil, err
		}
		resultAddrs = append(resultAddrs, addrs...)
	}

	// Bug in resolving RDS Aurora readonly endpoints
	// These are not usual round robin DNS records as expected
	// AWS manages round robin on by itself
	// As a result, to get all correct values you need to do pause between requests
	// experimently found that 20 requests with 500ms pause between each work the best
	isRdsRO, err := regexp.MatchString(`.*cluster-ro.*rds\.amazonaws\.com`, host)
	if err != nil {
		return nil, fmt.Errorf("can't check with RDS ro regex %v", host)
	}
	if isRdsRO {
		for i := 0; i < rdsVerifyAttempts; i++ {
			time.Sleep(time.Duration(rdsRequestPeriod) * time.Millisecond)
			addrs, err := net.LookupHost(host)
			if err != nil {
				return nil, err
			}
			resultAddrs = append(resultAddrs, addrs...)
		}
	}
	resultAddrs = utility.Unique(resultAddrs)
	return resultAddrs, nil
}

// RemovePrivateAddressesIPv4 - return incoming IPs massive without private addresses
func RemovePrivateAddressesIPv4(log logger.Logger, ips []net.IP) ([]net.IP, error) {
	var filteredIps []net.IP
	for _, ip := range ips {
		ip = ip.To4()
		isPrivate, err := CheckIPPrivate(ip)
		if err != nil {
			return nil, err
		}
		if !isPrivate {
			filteredIps = append(filteredIps, ip)
		} else {
			log.Debugf("%v was filtered out", ip.String())
		}
	}
	return filteredIps, nil
}

// CheckIPPrivate check whether IP range is private or public
func CheckIPPrivate(ip net.IP) (bool, error) {
	for _, cidr := range []string{
		"127.0.0.0/8",    // IPv4 loopback
		"10.0.0.0/8",     // RFC1918
		"172.16.0.0/12",  // RFC1918
		"192.168.0.0/16", // RFC1918
		"169.254.0.0/16", // RFC3927 link-local
		"::1/128",        // IPv6 loopback
		"fe80::/10",      // IPv6 link-local
		"fc00::/7",       // IPv6 unique local addr
	} {
		_, block, err := net.ParseCIDR(cidr)
		if err != nil {
			return false, errors.Wrapf(err, "parse error on %q: %v", cidr, err)
		}
		if block.Contains(ip) {
			return true, nil
		}
	}
	return false, nil
}

// CheckIPRangePrivate check whether IP range is private or public
func CheckIPRangePrivate(ipRange *net.IPNet) (bool, error) {
	for _, cidr := range []string{
		"127.0.0.0/8",    // IPv4 loopback
		"10.0.0.0/8",     // RFC1918
		"172.16.0.0/12",  // RFC1918
		"192.168.0.0/16", // RFC1918
		"169.254.0.0/16", // RFC3927 link-local
		"::1/128",        // IPv6 loopback
		"fe80::/10",      // IPv6 link-local
		"fc00::/7",       // IPv6 unique local addr
	} {
		_, block, err := net.ParseCIDR(cidr)
		if err != nil {
			return false, errors.Wrapf(err, "parse error on %q: %v", cidr, err)
		}
		if block.Contains(ipRange.IP) {
			return true, nil
		}
	}
	return false, nil
}

// UniqueIPs - returns slice without duplicates
func UniqueIPs(ips []net.IP) []net.IP {
	//convert to string slice
	var str []string
	for _, ip := range ips {
		str = append(str, ip.To4().String())
	}
	str = utility.Unique(str)
	var res []net.IP
	for _, ip := range str {
		res = append(res, net.ParseIP(ip).To4())
	}
	return res
}
