package dnsresolver

import (
	"net"
	"testing"

	"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/stretchr/testify/assert"

	"go.uber.org/zap/zaptest"
)

func TestRegistry(t *testing.T) {
	registry := NewRegistry()
	registry.AddRecord("8.8.8.8", "dns.google")
	registry.AddRecord("8.8.4.4", "dns2.google")
	registry.AddRecord("1.1.1.1", "dns3.notgoogle")
	registry.AddRecord("1.1.1.1", "dns4.notgoogle")

	assert.Equal(t, []string{"dns.google"}, registry.ResolvedIPs["8.8.8.8"], "Registry add check")
	assert.Equal(t, []string{"dns3.notgoogle", "dns4.notgoogle"}, registry.ResolvedIPs["1.1.1.1"], "Registry add check")

	registry2 := CopyRegistry(registry, nil)
	assert.Equal(t, []string{"dns.google"}, registry2.ResolvedIPs["8.8.8.8"], "Registry add check")
	assert.Equal(t, []string{"dns3.notgoogle", "dns4.notgoogle"}, registry2.ResolvedIPs["1.1.1.1"], "Registry copy check")

	registry3 := CopyRegistry(registry, []string{"8.8.4.4", "1.1.1.1"})
	assert.Equal(t, []string([]string(nil)), registry3.ResolvedIPs["8.8.8.8"], "Registry add check")
	assert.Equal(t, []string{"dns3.notgoogle", "dns4.notgoogle"}, registry3.ResolvedIPs["1.1.1.1"], "Registry limited copy check")

	//Let's get back to first test registry
	registry.RemoveIP("8.8.4.4")
	assert.Equal(t, []string([]string(nil)), registry.ResolvedIPs["8.8.4.4"], "Registry RemoveIP check")
	registry.AddRecord("8.8.4.4", "dns2.google")
	registry.RemoveIP("8.8.4.5")
	assert.Equal(t, registry2, registry, "Registry remove nonexistant IP check")
	registry.AddRecord("2.2.2.2", "dns4.notgoogle")
	registry.AddRecord("1.1.1.1", "dns5.notgoogle")
	registry.RemoveRecord("dns4.notgoogle")
	assert.Equal(t, []string{"dns3.notgoogle", "dns5.notgoogle"}, registry.ResolvedIPs["1.1.1.1"], "Registry remove record check")
	registry.AddRecord("2.2.2.2", "dns4.notgoogle")
	registry.AddRecord("1.1.1.1", "dns4.notgoogle")
	registry.RemoveRecordForIP("1.1.1.1", "dns4.notgoogle")
	assert.Equal(t, []string{"dns3.notgoogle", "dns5.notgoogle"}, registry.ResolvedIPs["1.1.1.1"], "Registry remove record for IP check")
	assert.Equal(t, []string{"dns4.notgoogle"}, registry.ResolvedIPs["2.2.2.2"], "Registry remove record for IP check")

	//Merge test
	mergeRegistry1 := NewRegistry()
	mergeRegistry1.AddRecord("8.8.8.8", "dns.google")
	mergeRegistry1.AddRecord("8.8.4.4", "dns2.google")
	mergeRegistry1.AddRecord("1.1.1.1", "dns3.notgoogle")
	mergeRegistry1.AddRecord("1.1.1.1", "dns4.notgoogle")
	mergeRegistry2 := NewRegistry()
	mergeRegistry2.AddRecord("9.9.9.9", "dns.google")
	mergeRegistry2.AddRecord("8.8.4.4", "dns2.google")
	mergeRegistry2.AddRecord("1.1.1.1", "dns3.notgoogle")
	mergeRegistry2.AddRecord("1.1.1.1", "dns6.notgoogle")

	expectedRegistry := NewRegistry()
	expectedRegistry.AddRecord("8.8.8.8", "dns.google")
	expectedRegistry.AddRecord("8.8.4.4", "dns2.google")
	expectedRegistry.AddRecord("1.1.1.1", "dns3.notgoogle")
	expectedRegistry.AddRecord("1.1.1.1", "dns4.notgoogle")
	expectedRegistry.AddRecord("9.9.9.9", "dns.google")
	expectedRegistry.AddRecord("1.1.1.1", "dns6.notgoogle")
	merged := MergeRegistries(mergeRegistry1, mergeRegistry2)
	assert.Equal(t, expectedRegistry, merged, "Registry merge check")

	expectedIpList := []net.IP{net.ParseIP("8.8.8.8").To4(), net.ParseIP("8.8.4.4").To4(), net.ParseIP("1.1.1.1").To4(), net.ParseIP("9.9.9.9").To4()}
	assert.ElementsMatch(t, expectedIpList, expectedRegistry.GetIPList(), "Registry get IP list check")
	expectedRecordList := []string{"dns.google", "dns2.google", "dns3.notgoogle", "dns4.notgoogle", "dns6.notgoogle"}
	assert.ElementsMatch(t, expectedRecordList, expectedRegistry.GetRecordList(), "Registry get IP list check")

	logger := zaptest.NewLogger(t).Sugar()
	merged.AddRecord("192.168.10.10", "dns9.notgoogle")
	merged.AddRecord("10.10.10.10", "dns10.notgoogle")
	err := merged.RemovePrivateAddressesIPv4(logger)
	if err != nil {
		assert.FailNow(t, "Can't remove private addresses")
	}
	assert.Equal(t, expectedRegistry, merged, "RemovePrivateAddressesIPv4 check")

}
func TestIPFunctions(t *testing.T) {
	logger := zaptest.NewLogger(t).Sugar()
	checkprivate, err := CheckIPPrivate(net.ParseIP("192.168.0.1"))
	if err != nil {
		assert.FailNow(t, "IP used for check is incorrect")
	}
	assert.True(t, checkprivate, "Private IP check")
	checkprivate, err = CheckIPPrivate(net.ParseIP("8.8.8.8"))
	if err != nil {
		assert.FailNow(t, "IP used for check is incorrect")
	}
	assert.False(t, checkprivate, "Private IP check")

	_, testIPRange, err := net.ParseCIDR("10.0.0.0/8")
	if err != nil {
		assert.FailNow(t, "IP range used for check is incorrect")
	}
	checkprivate, err = CheckIPRangePrivate(testIPRange)
	assert.True(t, checkprivate, "Private range IP check")

	_, testIPRange, err = net.ParseCIDR("11.0.0.0/8")
	if err != nil {
		assert.FailNow(t, "IP range used for check is incorrect")
	}
	checkprivate, err = CheckIPRangePrivate(testIPRange)
	assert.False(t, checkprivate, "Private range IP check")

	testIPSlice := []net.IP{net.ParseIP("1.2.3.4").To4(), net.ParseIP("4.3.2.1").To4(), net.ParseIP("1.2.3.4").To4(), net.ParseIP("192.168.8.81").To4()}
	expectedIPSlice := []net.IP{net.ParseIP("1.2.3.4").To4(), net.ParseIP("4.3.2.1").To4(), net.ParseIP("192.168.8.81").To4()}
	assert.Equal(t, 3, len(UniqueIPs(testIPSlice)), "Unique on IP list")
	assert.Equal(t, expectedIPSlice, UniqueIPs(testIPSlice), "Unique on IP list")

	expectedIPSlice = []net.IP{net.ParseIP("1.2.3.4").To4(), net.ParseIP("4.3.2.1").To4(), net.ParseIP("1.2.3.4").To4()}
	assessedIPSlice, err := RemovePrivateAddressesIPv4(logger, testIPSlice)
	if err != nil {
		assert.FailNow(t, "RemovePrivateAddressesIPv4 call failed")
	}
	assert.Equal(t, expectedIPSlice, assessedIPSlice, "Unique on IP list")

	testResolveAddrs, err := ResolveHost(logger, "dns.google")
	assert.Equal(t, 2, len(testResolveAddrs), "Host resolve test on google dns")
	assert.Contains(t, testResolveAddrs, "8.8.8.8", "Host resolve test on google dns")

	registry := ResolveCustomRecords(logger, []string{"dns.google", "google.com"})
	assert.Greater(t, len(registry.GetIPList()), 2, "List of custom records")
}

func TestRoute53Functions(t *testing.T) {
	logger := zaptest.NewLogger(t).Sugar()
	groupConfig := &config.GroupConfig{
		Name:               "Test group 1",
		AllowedRecordTypes: []string{"A"},
		SingleRecords:      []string{"dev-apollo-aurora-cep.cluster-custom-cic9zqsxd2y6.us-east-1.rds.amazonaws.com"},
		Zones: []config.Zone{
			{Name: "delphiplatform.io.", RecordFilters: []string{"^slz-api"}, AwsProfile: "gdb-delphi-dev"}},
	}
	testZoneID := "/hostedzone/Z2YYPPL1K0RHLS"
	groupManager := &config.GroupConfigManager{Group: groupConfig}
	groupManager.SetLogger(logger)

	// Create AWS route53 cli
	sess, err := session.NewSessionWithOptions(session.Options{
		Profile:           "gdb-delphi-dev",
		SharedConfigState: session.SharedConfigEnable,
	})
	if err != nil {
		assert.FailNowf(t, "Unable to create AWS client", err.Error())
	}
	svc := route53.New(sess)

	records, err := GetRoute53ZoneRecords(testZoneID, svc)
	if err != nil {
		assert.FailNowf(t, "Can't get dns records from GetRoute53ZoneRecords", err.Error())
	}
	assert.Greater(t, len(records), 1)

	zoneID, err := GetZoneIDByName(logger, svc, "delphiplatform.io.")
	if err != nil {
		assert.FailNowf(t, "Can't get zone id from GetZoneIDByName", err.Error())
	}
	assert.Equal(t, zoneID, testZoneID)

	records = FilterRecordSets(groupManager, groupConfig.Zones[0].Name, records)
	assert.GreaterOrEqual(t, len(records), 1)

	records, err = GetRoute53RecordsByGroupConfig(logger, groupManager)
	if err != nil {
		assert.FailNowf(t, "Can't get dns records from GetRoute53RecordsByGroupConfig: %v", err.Error())
	}
	assert.Equal(t, len(records), 1)
	assert.Equal(t, "slz-api.delphiplatform.io.", *records[0].Name)

	registry := ResolveRoute53Records(logger, records)
	assert.Greater(t, len(registry.GetIPList()), 2)

	//FilterRecordSets()
	//ResolveRoute53Records()
}

// test in different zone
func TestRoute53FunctionsV2(t *testing.T) {
	logger := zaptest.NewLogger(t).Sugar()
	groupConfig := &config.GroupConfig{
		Name:               "Test group 1",
		AllowedRecordTypes: []string{"A"},
		SingleRecords:      []string{},
		Zones: []config.Zone{
			{Name: "apollo.stream.", RecordFilters: []string{"^uat-static", "^uat-proxy"}, AwsProfile: "gdb-delphi-dev"}},
	}
	testZoneID := "/hostedzone/Z46LEVA0ZF8PP"
	groupManager := &config.GroupConfigManager{Group: groupConfig}
	groupManager.SetLogger(logger)

	// Create AWS route53 cli
	sess, err := session.NewSessionWithOptions(session.Options{
		Profile:           "gdb-delphi-dev",
		SharedConfigState: session.SharedConfigEnable,
	})
	if err != nil {
		assert.FailNowf(t, "Unable to create AWS client", err.Error())
	}
	svc := route53.New(sess)

	records, err := GetRoute53ZoneRecords(testZoneID, svc)
	if err != nil {
		assert.FailNowf(t, "Can't get dns records from GetRoute53ZoneRecords", err.Error())
	}
	assert.Greater(t, len(records), 1)

	zoneID, err := GetZoneIDByName(logger, svc, "apollo.stream.")
	if err != nil {
		assert.FailNowf(t, "Can't get zone id from GetZoneIDByName", err.Error())
	}
	assert.Equal(t, zoneID, testZoneID)

	records = FilterRecordSets(groupManager, groupConfig.Zones[0].Name, records)
	assert.GreaterOrEqual(t, len(records), 1)

	records, err = GetRoute53RecordsByGroupConfig(logger, groupManager)
	if err != nil {
		assert.FailNowf(t, "Can't get dns records from GetRoute53RecordsByGroupConfig: %v", err.Error())
	}
	assert.Equal(t, len(records), 2)
	assert.Equal(t, "uat-static.apollo.stream.", *records[1].Name)

	registry := ResolveRoute53Records(logger, records)
	assert.Greater(t, len(registry.GetIPList()), 1)

	//FilterRecordSets()
	//ResolveRoute53Records()
}
