#!/usr/bin/env python3
"""Unit tests."""

import unittest

from bind9_records import Record, TXTData

class TestAddress(unittest.TestCase):
    """Test address records."""

    maxDiff = None

    def test_a_record_parse(self) -> None:
        tests = [
            'test\tA\t127.0.0.1',
            'test\tAAAA\t::1',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_caa_record_parse(self) -> None:
        tests = [
            'test\tCAA\t0 iodef ";"',
            'test\tCAA\t1 iodef example.com.',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_cert_record_parse(self) -> None:
        tests = [
            'test\tCERT\tPGP 0 0 MTIzNDU2',
            'test\tCERT\t37 0 42 MTIzNDU2',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_cname_record_parse(self) -> None:
        tests = [
            'test\tCNAME\texample.com.',
            'test\tCNAME\texample\\.dot.bar.',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_ds_record_parse(self) -> None:
        tests = [
            'test\tDS\t123 1 2 0123456789ABCDEF',
            # 'test\tDS\t123 ECDSAP256SHA256 SHA256 0123456789ABCDEF',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_loc_record_parse(self) -> None:
        tests = [
            'test\tLOC\t42 21 28.764 N 71 0 51.617 W -44.4m 2000.1m',
            'test\tLOC\t42 S 71 E 44m',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_mx_record_parse(self) -> None:
        tests = [
            'test\tMX\t0 example.com.',
            'test\tMX\t0 example\\.dot.bar.',
            'test\tMX\t100 .',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_naptr_record_parse(self) -> None:
        tests = [
            'test\tNAPTR\t100 50 "a" "z3950+N2L+N2C" "" cidserver.example.com.',
            ('test\tNAPTR\t' r'100 10 "" "" "!^urn:cid:.+@([^\.]+\.)(.*)$!\2!i" .'),
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_srv_record_parse(self) -> None:
        tests = [
            'test\tSRV\t0 0 443 example.com.',
            'test\tSRV\t0 0 8443 example\\.dot.bar.',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_tlsa_record_parse(self) -> None:
        tests = [
            'test\tTLSA\t0 0 0 0123456789ABCDEF',
            'test\tTLSA\t0 0 1 0123456789ABCDEF',
            'test\tTLSA\t0 0 2 0123456789ABCDEF',
        ]

        for test in tests:
            r = Record(test)
            self.assertEqual(str(r), test)

    def test_txt_data_create(self) -> None:
        tests = [
            ({'version': 'spf', 'mechanisms': 'include:example.com.'},
             'v=SPF1 include:example.com.'),
            ({'version': 'spf', 'mechanisms': ['a:smtp.example.com', '-all']},
             'v=SPF1 a:smtp.example.com -all'),
            ({'version': 'dmarc', 'policy': 'reject', 'percent': 100},
             'v=DMARC1; p=reject; pct=100;'),
        ]

        for kwargs, expected in tests:
            d = TXTData.create(**kwargs)
            self.assertEqual(str(d), expected)


if __name__ == '__main__':
    unittest.main()
