"""Tests for DNS service.""" import pytest from unittest.mock import patch, MagicMock import dns.resolver import dns.exception from src.dns.service import DNSService from src.dns.schemas import DNSLookupRequest, DNSRecord from src.dns.exceptions import DNSQueryError @pytest.fixture def dns_service(): """Create a DNSService instance.""" return DNSService() class TestDNSServiceInit: """Test DNSService initialization.""" def test_service_has_resolver(self, dns_service): """Service should have resolver configured.""" assert dns_service.resolver is not None def test_service_has_timeout(self, dns_service): """Service should have timeout configured.""" assert dns_service.resolver.timeout == 5.0 assert dns_service.resolver.lifetime == 10.0 def test_supported_record_types(self, dns_service): """Service should have supported record types.""" assert "A" in dns_service.SUPPORTED_RECORD_TYPES assert "AAAA" in dns_service.SUPPORTED_RECORD_TYPES assert "MX" in dns_service.SUPPORTED_RECORD_TYPES assert "TXT" in dns_service.SUPPORTED_RECORD_TYPES assert "CNAME" in dns_service.SUPPORTED_RECORD_TYPES assert "NS" in dns_service.SUPPORTED_RECORD_TYPES class TestDNSServiceLookup: """Test DNS lookup functionality.""" @pytest.mark.asyncio async def test_lookup_validates_record_type(self, dns_service): """lookup should raise error for unsupported record type.""" request = DNSLookupRequest(domain="example.com", record_type="INVALID") with pytest.raises(DNSQueryError) as exc_info: await dns_service.lookup(request) assert "Unsupported record type" in str(exc_info.value) @pytest.mark.asyncio async def test_lookup_returns_response_on_success(self, dns_service): """lookup should return DNSLookupResponse on success.""" request = DNSLookupRequest(domain="example.com", record_type="A") # Mock the resolver mock_answer = MagicMock() mock_rdata = MagicMock() mock_rdata.__str__ = MagicMock(return_value="93.184.216.34") mock_answer.__iter__ = MagicMock(return_value=iter([mock_rdata])) with patch("dns.resolver.Resolver") as mock_resolver_class: mock_resolver = MagicMock() mock_resolver.resolve.return_value = mock_answer mock_resolver.nameservers = ["8.8.8.8"] mock_resolver_class.return_value = mock_resolver response = await dns_service.lookup(request) assert response.success is True assert response.domain == "example.com" assert response.record_type == "A" assert len(response.records) > 0 @pytest.mark.asyncio async def test_lookup_uses_custom_nameserver(self, dns_service): """lookup should use custom nameserver when specified.""" request = DNSLookupRequest( domain="example.com", record_type="A", nameserver="1.1.1.1" ) mock_answer = MagicMock() mock_answer.__iter__ = MagicMock(return_value=iter([])) with patch("dns.resolver.Resolver") as mock_resolver_class: mock_resolver = MagicMock() mock_resolver.resolve.return_value = mock_answer mock_resolver.nameservers = [] mock_resolver_class.return_value = mock_resolver response = await dns_service.lookup(request) # Verify nameserver was set assert mock_resolver.nameservers == ["1.1.1.1"] assert response.nameserver_used == "1.1.1.1" @pytest.mark.asyncio async def test_lookup_handles_nxdomain(self, dns_service): """lookup should handle NXDOMAIN (domain not found).""" request = DNSLookupRequest(domain="nonexistent.invalid", record_type="A") with patch("dns.resolver.Resolver") as mock_resolver_class: mock_resolver = MagicMock() mock_resolver.resolve.side_effect = dns.resolver.NXDOMAIN() mock_resolver.nameservers = ["8.8.8.8"] mock_resolver_class.return_value = mock_resolver response = await dns_service.lookup(request) assert response.success is False assert "Domain not found" in response.error_message @pytest.mark.asyncio async def test_lookup_handles_no_answer(self, dns_service): """lookup should handle NoAnswer (no records of type).""" request = DNSLookupRequest(domain="example.com", record_type="AAAA") with patch("dns.resolver.Resolver") as mock_resolver_class: mock_resolver = MagicMock() mock_resolver.resolve.side_effect = dns.resolver.NoAnswer() mock_resolver.nameservers = ["8.8.8.8"] mock_resolver_class.return_value = mock_resolver response = await dns_service.lookup(request) assert response.success is False assert "No AAAA records found" in response.error_message @pytest.mark.asyncio async def test_lookup_handles_timeout(self, dns_service): """lookup should handle DNS timeout.""" request = DNSLookupRequest(domain="example.com", record_type="A") with patch("dns.resolver.Resolver") as mock_resolver_class: mock_resolver = MagicMock() mock_resolver.resolve.side_effect = dns.resolver.Timeout() mock_resolver.nameservers = ["8.8.8.8"] mock_resolver_class.return_value = mock_resolver response = await dns_service.lookup(request) assert response.success is False assert "timeout" in response.error_message.lower() @pytest.mark.asyncio async def test_lookup_handles_dns_exception(self, dns_service): """lookup should handle generic DNS exceptions.""" request = DNSLookupRequest(domain="example.com", record_type="A") with patch("dns.resolver.Resolver") as mock_resolver_class: mock_resolver = MagicMock() mock_resolver.resolve.side_effect = dns.exception.DNSException("DNS error") mock_resolver.nameservers = ["8.8.8.8"] mock_resolver_class.return_value = mock_resolver response = await dns_service.lookup(request) assert response.success is False assert "DNS error" in response.error_message @pytest.mark.asyncio async def test_lookup_handles_unexpected_exception(self, dns_service): """lookup should handle unexpected exceptions.""" request = DNSLookupRequest(domain="example.com", record_type="A") with patch("dns.resolver.Resolver") as mock_resolver_class: mock_resolver = MagicMock() mock_resolver.resolve.side_effect = Exception("Unexpected error") mock_resolver.nameservers = ["8.8.8.8"] mock_resolver_class.return_value = mock_resolver response = await dns_service.lookup(request) assert response.success is False assert "Unexpected error" in response.error_message class TestDNSServiceParseRecord: """Test record parsing.""" def test_parse_a_record(self, dns_service): """_parse_record should parse A record.""" mock_rdata = MagicMock() mock_rdata.__str__ = MagicMock(return_value="192.168.1.1") record = dns_service._parse_record(mock_rdata, "A") assert record is not None assert record.value == "192.168.1.1" def test_parse_aaaa_record(self, dns_service): """_parse_record should parse AAAA record.""" mock_rdata = MagicMock() mock_rdata.__str__ = MagicMock(return_value="2001:db8::1") record = dns_service._parse_record(mock_rdata, "AAAA") assert record is not None assert record.value == "2001:db8::1" def test_parse_mx_record(self, dns_service): """_parse_record should parse MX record with priority.""" mock_rdata = MagicMock() mock_rdata.exchange = "mail.example.com" mock_rdata.preference = 10 record = dns_service._parse_record(mock_rdata, "MX") assert record is not None assert "mail.example.com" in record.value assert record.priority == 10 def test_parse_txt_record(self, dns_service): """_parse_record should parse TXT record.""" mock_rdata = MagicMock() mock_rdata.strings = [b"v=spf1 include:_spf.google.com ~all"] record = dns_service._parse_record(mock_rdata, "TXT") assert record is not None assert "spf1" in record.value def test_parse_cname_record(self, dns_service): """_parse_record should parse CNAME record.""" mock_rdata = MagicMock() mock_rdata.target = "alias.example.com" record = dns_service._parse_record(mock_rdata, "CNAME") assert record is not None assert "alias.example.com" in record.value def test_parse_ns_record(self, dns_service): """_parse_record should parse NS record.""" mock_rdata = MagicMock() mock_rdata.target = "ns1.example.com" record = dns_service._parse_record(mock_rdata, "NS") assert record is not None assert "ns1.example.com" in record.value def test_parse_soa_record(self, dns_service): """_parse_record should parse SOA record.""" mock_rdata = MagicMock() mock_rdata.mname = "ns1.example.com" mock_rdata.rname = "admin.example.com" mock_rdata.serial = 2024010101 record = dns_service._parse_record(mock_rdata, "SOA") assert record is not None assert "ns1.example.com" in record.value assert "2024010101" in record.value def test_parse_srv_record(self, dns_service): """_parse_record should parse SRV record.""" mock_rdata = MagicMock() mock_rdata.target = "server.example.com" mock_rdata.port = 443 mock_rdata.priority = 10 mock_rdata.weight = 100 record = dns_service._parse_record(mock_rdata, "SRV") assert record is not None assert "server.example.com" in record.value assert "port=443" in record.value assert record.priority == 10 def test_parse_record_returns_none_on_error(self, dns_service): """_parse_record should return None on parsing error.""" mock_rdata = MagicMock() mock_rdata.__str__ = MagicMock(side_effect=Exception("Parse error")) record = dns_service._parse_record(mock_rdata, "A") assert record is None class TestDNSServiceErrorResponse: """Test error response generation.""" def test_error_response_includes_domain(self, dns_service): """_error_response should include domain.""" import time request = DNSLookupRequest(domain="test.example.com", record_type="A") response = dns_service._error_response(request, "8.8.8.8", time.time(), "Test error") assert response.domain == "test.example.com" def test_error_response_includes_record_type(self, dns_service): """_error_response should include record type.""" import time request = DNSLookupRequest(domain="test.example.com", record_type="mx") response = dns_service._error_response(request, "8.8.8.8", time.time(), "Test error") assert response.record_type == "MX" # Should be uppercase def test_error_response_has_empty_records(self, dns_service): """_error_response should have empty records list.""" import time request = DNSLookupRequest(domain="test.example.com", record_type="A") response = dns_service._error_response(request, "8.8.8.8", time.time(), "Test error") assert response.records == [] def test_error_response_has_success_false(self, dns_service): """_error_response should have success=False.""" import time request = DNSLookupRequest(domain="test.example.com", record_type="A") response = dns_service._error_response(request, "8.8.8.8", time.time(), "Test error") assert response.success is False def test_error_response_includes_error_message(self, dns_service): """_error_response should include error message.""" import time request = DNSLookupRequest(domain="test.example.com", record_type="A") response = dns_service._error_response(request, "8.8.8.8", time.time(), "Specific error") assert response.error_message == "Specific error"