"""Offline metadata fixtures: no real socket, deployment host, cloud API or credentials."""
import copy
import json
import socket
import unittest
from unittest.mock import Mock, patch

from network_provider import MetadataFailure, ORIGIN, PATHS, _fetch, observe, validate_provider
from network_provider_test_fixture import INVENTORY, PUBLIC, VALUES, connection, fixture


class ProviderTests(unittest.TestCase):
    def test_identity_and_private_gateway_route_are_separate(self):
        result, fetch = fixture()
        self.assertEqual(result['status'], 'verified')
        self.assertEqual(result['private_subnet'], '10.116.0.0/20')
        self.assertEqual(result['private_gateway_route'], 'not_observed')
        self.assertEqual([call.args[0] for call in fetch.call_args_list], list(PATHS.values()))
        self.assertEqual(len(fetch.call_args_list), 6)
        inventory = copy.deepcopy(INVENTORY)
        inventory['host_records'].append({'subnet': '10.116.0.1/32', 'kind': 'route_gateway', 'device': 'eth1'})
        self.assertEqual(fixture(inventory=inventory)[0]['private_gateway_route'], 'observed')
        inventory['host_records'][-1]['device'] = 'eth2'
        result = fixture(inventory=inventory)[0]
        self.assertEqual((result['status'], result['private_gateway_route']), ('verified', 'not_observed'))
        self.assertNotIn('policy', json.dumps(result))

    def test_expected_host_requires_literal_ipv4_and_matching_kernel_address(self):
        for expected in ('host.example', '::1', '138.197.036.107', 123, None):
            result, fetch = fixture(expected=expected)
            self.assertEqual(result['reason'], 'expected_host_unverified')
            fetch.assert_not_called()
        result, _ = fixture(expected='138.197.36.108')
        self.assertEqual((result['status'], result['reason']), ('mismatch', 'public_address_mismatch'))
        self.assertEqual(result['droplet_id'], '123456789')
        inventory = copy.deepcopy(INVENTORY)
        inventory['host_records'].pop(0)
        self.assertEqual(fixture(inventory=inventory)[0]['reason'], 'public_address_unobserved')
        for inventory in ({}, dict(INVENTORY, complete=False)):
            result, fetch = fixture(inventory=inventory)
            self.assertEqual(result['reason'], 'host_inventory_unverified')
            fetch.assert_not_called()

    def test_private_identity_requires_one_matching_nonbridge_device(self):
        variants = []
        for index, device in ((1, 'eth2'), (2, 'eth2'), (1, ''), (1, 'br-fixture')):
            inventory = copy.deepcopy(INVENTORY)
            inventory['host_records'][index]['device'] = device
            variants.append(inventory)
        inventory = copy.deepcopy(INVENTORY)
        inventory['host_records'].append(dict(inventory['host_records'][1], device='eth2'))
        variants.append(inventory)
        inventory = copy.deepcopy(INVENTORY)
        inventory['networks'] = [{'Driver': 'bridge', 'Id': 'a' * 64,
                                  'Options': {'com.docker.network.bridge.name': 'eth1'}}]
        variants.append(inventory)
        inventory = copy.deepcopy(INVENTORY)
        inventory['host_records'][2]['subnet'] = '10.116.0.0/24'
        variants.append(inventory)
        for inventory in variants:
            with self.subTest(inventory=inventory):
                self.assertEqual(fixture(inventory=inventory)[0]['reason'], 'private_address_mismatch')

    def test_malformed_leaves_are_closed_and_never_retained(self):
        changes = [('droplet_id', b'0'), ('droplet_id', b'123456789012345678901'),
                   ('region', b'PRIVATE_SENTINEL'), ('region', b'nyc3\x00'),
                   ('public_ipv4', b'138.197.036.107'), ('public_ipv4', b'::1'),
                   ('private_ipv4', b'8.8.8.8'), ('private_ipv4', b'10.116.0.0'),
                   ('private_gateway', b'10.116.16.1'), ('private_gateway', b'10.116.0.2'),
                   ('private_gateway', b'10.116.15.255'), ('netmask', b'255.0.255.0'),
                   ('netmask', b'0.0.15.255'), ('netmask', b'20'), ('netmask', b'255.255.255.255'),
                   ('droplet_id', b'PRIVATE_SENTINEL' * 10), ('region', b'\xff')]
        for field, value in changes:
            with self.subTest(field=field, value=value):
                result, _ = fixture({field: value})
                self.assertEqual(result['status'], 'unverified')
                self.assertEqual(result['droplet_id'], '')
                self.assertNotIn('PRIVATE_SENTINEL', json.dumps(result))
        for error, reason in ((socket.timeout('PRIVATE_SENTINEL'), 'metadata_timeout'),
                              (OSError('PRIVATE_SENTINEL'), 'metadata_unavailable'),
                              (MetadataFailure('PRIVATE_SENTINEL'), 'metadata_invalid')):
            result = observe(INVENTORY, PUBLIC, Mock(side_effect=error))
            validate_provider(result)
            self.assertEqual(result['reason'], reason)
            self.assertNotIn('PRIVATE_SENTINEL', json.dumps(result))

    def test_transport_is_fixed_bounded_and_has_no_redirect_or_proxy_path(self):
        fake = connection(b'HTTP/1.0 200 OK\r\nContent-Length: 3\r\n\r\n123')
        with patch('network_provider.socket.socket', return_value=fake) as factory, \
                patch('network_provider.time.monotonic', return_value=0), \
                patch.dict('os.environ', {'HTTP_PROXY': 'http://PRIVATE_SENTINEL:80'}):
            self.assertEqual(_fetch(PATHS['droplet_id'], 10), b'123')
        factory.assert_called_once_with(socket.AF_INET, socket.SOCK_STREAM)
        fake.connect.assert_called_once_with(ORIGIN)
        fake.sendall.assert_called_once_with(
            b'GET /metadata/v1/id HTTP/1.0\r\nHost: 169.254.169.254\r\nConnection: close\r\n\r\n')
        self.assertTrue(all(0 < call.args[0] <= 2 for call in fake.settimeout.call_args_list))
        self.assertTrue(all(call.args[0] <= 129 for call in fake.recv.call_args_list))
        for path in ('/metadata/v1.json', '/metadata/v1/user-data', '/metadata/v1/auth-token',
                     '/metadata/v1/id?x=1', 'http://other/metadata/v1/id'):
            with patch('network_provider.socket.socket') as factory, self.assertRaises(MetadataFailure):
                _fetch(path, 10)
            factory.assert_not_called()

    def test_http_refusals_headers_and_body_limits_do_not_leak_content(self):
        cases = [(b'HTTP/1.0 404 Missing\r\n\r\nPRIVATE_SENTINEL', 'metadata_missing'),
                 (b'HTTP/1.0 302 Found\r\nLocation: http://other/\r\n\r\n', 'metadata_unavailable'),
                 (b'HTTP/1.0 200 OK\r\nX: ' + b'a' * 8192 + b'\r\n\r\n', 'metadata_response_limit'),
                 (b'HTTP/1.0 200 OK\r\n' + b'X: a\r\n' * 17 + b'\r\n', 'metadata_response_limit'),
                 (b'HTTP/1.0 200 OK\r\n\r\n' + b'a' * 129, 'metadata_response_limit'),
                 (b'HTTP/1.0 200 OK\r\nContent-Length: 129\r\n\r\n', 'metadata_response_limit'),
                 (b'HTTP/1.0 200 OK\r\nContent-Length: 3\r\n\r\n12', 'metadata_invalid'),
                 (b'HTTP/1.0 200 OK\r\nContent-Length: 1\r\nContent-Length: 1\r\n\r\n1', 'metadata_invalid'),
                 (b'HTTP/1.0 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n', 'metadata_invalid'),
                 (b'HTTP/1.0 200 OK\r\nContent-Encoding: gzip\r\n\r\n', 'metadata_invalid')]
        for wire, reason in cases:
            with self.subTest(reason=reason), patch('network_provider.socket.socket', return_value=connection(wire)), \
                    patch('network_provider.time.monotonic', return_value=0):
                result = observe(INVENTORY, PUBLIC)
            validate_provider(result)
            self.assertEqual(result['reason'], reason)
            self.assertNotIn('PRIVATE_SENTINEL', json.dumps(result))

    def test_deadlines_stop_slow_leaves_and_bound_the_whole_sequence(self):
        clock = [0]
        def slow(path, deadline):
            clock[0] += 2.1
            return b'123'
        with patch('network_provider.time.monotonic', side_effect=lambda: clock[0]):
            fetch = Mock(side_effect=slow)
            self.assertEqual(observe(INVENTORY, PUBLIC, fetch)['reason'], 'metadata_timeout')
        self.assertEqual(fetch.call_count, 1)
        clock[0] = 0
        def steady(path, deadline):
            clock[0] += 1.9
            return VALUES[next(key for key in PATHS if PATHS[key] == path)]
        with patch('network_provider.time.monotonic', side_effect=lambda: clock[0]):
            fetch = Mock(side_effect=steady)
            self.assertEqual(observe(INVENTORY, PUBLIC, fetch)['reason'], 'metadata_timeout')
        self.assertEqual(fetch.call_count, 6)
        self.assertEqual(fetch.call_args.args[1], 10)

    def test_receipt_rejects_extra_fields_bad_types_and_contradictory_states(self):
        result, _ = fixture()
        changes = ({'secret': 'PRIVATE_SENTINEL'}, {'droplet_id': 123}, {'region': []},
                   {'private_gateway': 1}, {'private_subnet': '10.116.0.2/20'},
                   {'provider': 'other'}, {'reason': 'PRIVATE_SENTINEL'}, {'status': 'unverified'},
                   {'private_gateway_route': 'unverified'}, {'public_ipv4': '138.197.036.107'})
        for changed in changes:
            with self.subTest(changed=changed), self.assertRaises(ValueError):
                validate_provider(dict(result, **changed))
        failure = observe(INVENTORY, 'hostname', Mock())
        for changed in ({'droplet_id': '123'}, {'private_gateway_route': 'not_observed'}, {'public_ipv4': None}):
            with self.assertRaises(ValueError):
                validate_provider(dict(failure, **changed))
        mismatch, _ = fixture(expected='138.197.36.108')
        with self.assertRaises(ValueError):
            validate_provider(dict(mismatch, private_gateway_route='observed'))


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