"""pika.data tests."""

import datetime
import decimal
import math
import struct
import unittest
from collections import OrderedDict
from typing import ClassVar

from pika import data, exceptions


class DataTests(unittest.TestCase):

    FIELD_TBL_ENCODED = (
        b'\x00\x00\x00\xf7'
        b'\x05arrayA\x00\x00\x00\x0fI\x00\x00\x00\x01I'
        b'\x00\x00\x00\x02I\x00\x00\x00\x03'
        b'\x07boolvalt\x01'
        b'\x07decimalD\x02\x00\x00\x01:'
        b'\x0bdecimal_tooD\x00\x00\x00\x00d'
        b'\x07dictvalF\x00\x00\x00\x0c\x03fooS\x00\x00\x00\x03bar'
        b'\x06intvalI\x00\x00\x00\x01'
        b'\x06bigintl\x00\x00\x00\x00\x9a\x7e\xc8\x00'
        b'\x07longvalI\x36\x65\x26\x55'
        b'\x07neglongI\xff\xff\xff\xff'
        b'\x04nullV'
        b'\x06strvalS\x00\x00\x00\x04Test'
        b'\x0ctimestampvalT\x00\x00\x00\x00Ec)\x92'
        b'\x07unicodeS\x00\x00\x00\x08utf8=\xe2\x9c\x93')

    FIELD_TBL_ENCODED += b'\x05bytesx\x00\x00\x00\x06foobar'
    FIELD_TBL_ENCODED += b'\x08floatvald' + struct.pack('>d', 2.5)

    FIELD_TBL_VALUE: ClassVar[OrderedDict] = OrderedDict([
        ('array', [1, 2, 3]),
        ('boolval', True),
        ('decimal', decimal.Decimal('3.14')),
        ('decimal_too', decimal.Decimal(100)),
        ('dictval', {
            'foo': 'bar'
        }),
        ('intval', 1),
        ('bigint', 2592000000),
        ('longval', 912598613),
        ('neglong', -1),
        ('null', None),
        ('strval', 'Test'),
        ('timestampval',
         datetime.datetime(2006,
                           11,
                           21,
                           16,
                           30,
                           10,
                           tzinfo=datetime.timezone.utc)),
        ('unicode', 'utf8=✓'),
        ('bytes', b'foobar'),
        ('floatval', 2.5),
    ])

    def test_decode_bytes(self):
        input = (b'\x00\x00\x00\x01'
                 b'\x05bytesx\x00\x00\x00\x06foobar')
        result = data.decode_table(input, 0)
        self.assertEqual(result, ({'bytes': b'foobar'}, 21))

    # b'\x08shortints\x04\xd2'
    # ('shortint', 1234),
    def test_decode_shortint(self):
        input = (b'\x00\x00\x00\x01'
                 b'\x08shortints\x04\xd2')
        result = data.decode_table(input, 0)
        self.assertEqual(result, ({'shortint': 1234}, 16))

    def test_encode_table(self):
        result = []
        data.encode_table(result, self.FIELD_TBL_VALUE)
        self.assertEqual(b''.join(result), self.FIELD_TBL_ENCODED)

    def test_encode_table_bytes(self):
        result = []
        byte_count = data.encode_table(result, self.FIELD_TBL_VALUE)
        self.assertEqual(byte_count, 251)

    def test_decode_table(self):
        value, _byte_count = data.decode_table(self.FIELD_TBL_ENCODED, 0)
        self.assertDictEqual(value, self.FIELD_TBL_VALUE)

    def test_decode_table_bytes(self):
        _value, byte_count = data.decode_table(self.FIELD_TBL_ENCODED, 0)
        self.assertEqual(byte_count, 251)

    def test_decode_signed_long_negative(self):
        """
        Verify that type tag 'l' decodes as signed 64-bit (fixes #1531).

        RabbitMQ encodes negative longs (e.g. x-delay after delivery) with type tag 'l' and signed
        64-bit representation.
        """
        # Table with x-delay = -30000 encoded as signed 64-bit 'l'
        input = (b'\x00\x00\x00\x10'
                 b'\x07x-delayl\xff\xff\xff\xff\xff\xff\x8a\xd0')
        result, _ = data.decode_table(input, 0)
        self.assertEqual(result, {'x-delay': -30000})

    def test_encode_decode_negative_long_roundtrip(self):
        """Verify negative long values round-trip correctly."""
        table = {'x-delay': -30000}
        pieces = []
        data.encode_table(pieces, table)
        encoded = b''.join(pieces)
        decoded, _ = data.decode_table(encoded, 0)
        self.assertEqual(decoded, table)

    def test_encode_raises(self):
        self.assertRaises(exceptions.UnsupportedAMQPFieldException,
                          data.encode_table, [], {'foo': {1, 2, 3}})

    def test_decode_raises(self):
        self.assertRaises(exceptions.InvalidFieldTypeException,
                          data.decode_table,
                          b'\x00\x00\x00\t\x03fooZ\x00\x00\x04\xd2', 0)

    def test_long_repr(self):
        value = 912598613
        self.assertEqual(repr(value), '912598613')

    def test_encode_short_string_too_long(self):
        self.assertRaises(exceptions.ShortStringTooLong,
                          data.encode_short_string, [], 'a' * 256)

    def test_decode_short_string_invalid_utf8(self):
        encoded = b'\x02\xff\xfe'
        value, offset = data.decode_short_string(encoded, 0)
        self.assertIsInstance(value, bytes)
        self.assertEqual(value, b'\xff\xfe')
        self.assertEqual(offset, 3)

    def test_encode_decimal_scale_too_large(self):
        value = decimal.Decimal('0.' + '0' * 300 + '1')
        self.assertRaises(exceptions.UnencodableDecimalError, data.encode_value,
                          [], value)

    def test_encode_decimal_mantissa_out_of_range(self):
        value = decimal.Decimal('99999999999999999999.99')
        self.assertRaises(exceptions.UnencodableDecimalError, data.encode_value,
                          [], value)

    def test_encode_decimal_nan(self):
        self.assertRaises(exceptions.UnencodableDecimalError, data.encode_value,
                          [], decimal.Decimal('NaN'))

    def test_encode_decimal_infinity(self):
        self.assertRaises(exceptions.UnencodableDecimalError, data.encode_value,
                          [], decimal.Decimal('Infinity'))

    def test_encode_decimal_excess_precision_is_not_silently_rounded(self):
        # A finite value with more significant digits than the thread-local
        # decimal context precision (28 by default) must be rejected, not
        # silently rounded down to a small mantissa and encoded as a different
        # number. as_tuple() keeps every digit regardless of context.
        value = decimal.Decimal('1.0000000000000000000000000000005')
        self.assertRaises(exceptions.UnencodableDecimalError, data.encode_value,
                          [], value)

    def test_encode_decimal_context_precision_does_not_affect_result(self):
        # Encoding must not depend on the ambient decimal context, so a low
        # context precision cannot round an in-range value before it is packed.
        value = decimal.Decimal('12345.6789')
        with decimal.localcontext() as ctx:
            ctx.prec = 3
            pieces = []
            data.encode_table(pieces, {'k': value})
        decoded, _ = data.decode_table(b''.join(pieces), 0)
        self.assertEqual(decoded, {'k': value})

    def test_encode_decimal_preserves_trailing_zeros(self):
        # Dropping normalize() keeps the scale exactly as given, so a value
        # with trailing fractional zeros round-trips to that same scale.
        value = decimal.Decimal('1.10')
        pieces = []
        data.encode_table(pieces, {'k': value})
        decoded, _ = data.decode_table(b''.join(pieces), 0)
        self.assertEqual(decoded, {'k': value})

    def test_decode_value_short_short_int(self):
        # b'b' = signed byte
        encoded = b'\x00\x00\x00\x04\x01kb\xff'
        result, _ = data.decode_table(encoded, 0)
        self.assertEqual(result, {'k': -1})

    def test_decode_value_short_short_uint(self):
        # b'B' = unsigned byte
        encoded = b'\x00\x00\x00\x04\x01kB\xff'
        result, _ = data.decode_table(encoded, 0)
        self.assertEqual(result, {'k': 255})

    def test_decode_value_short_int(self):
        # b'U' = signed short
        encoded = b'\x00\x00\x00\x05\x01kU' + struct.pack('>h', -1000)
        result, _ = data.decode_table(encoded, 0)
        self.assertEqual(result, {'k': -1000})

    def test_decode_value_short_uint(self):
        # b'u' = unsigned short
        encoded = b'\x00\x00\x00\x05\x01ku' + struct.pack('>H', 1000)
        result, _ = data.decode_table(encoded, 0)
        self.assertEqual(result, {'k': 1000})

    def test_decode_value_long_uint(self):
        # b'i' = unsigned long
        encoded = b'\x00\x00\x00\x07\x01ki' + struct.pack('>I', 4294967295)
        result, _ = data.decode_table(encoded, 0)
        self.assertEqual(result, {'k': 4294967295})

    def test_decode_value_long_long_int_uppercase(self):
        # b'L' = signed 64-bit int
        encoded = b'\x00\x00\x00\x0b\x01kL' + struct.pack('>q', -30000)
        result, _ = data.decode_table(encoded, 0)
        self.assertEqual(result, {'k': -30000})

    def test_decode_value_float(self):
        # b'f' = 32-bit float
        encoded = b'\x00\x00\x00\x07\x01kf' + struct.pack('>f', 1.5)
        result, _ = data.decode_table(encoded, 0)
        self.assertAlmostEqual(result['k'], 1.5, places=5)

    def test_decode_value_double(self):
        # b'd' = 64-bit double
        encoded = b'\x00\x00\x00\x0b\x01kd' + struct.pack('>d', 3.14)
        result, _ = data.decode_table(encoded, 0)
        self.assertAlmostEqual(result['k'], 3.14, places=10)

    def test_encode_value_double(self):
        # A Python float is a C double, so it encodes as b'd'
        pieces = []
        length = data.encode_value(pieces, 3.14)
        self.assertEqual(b''.join(pieces), b'd' + struct.pack('>d', 3.14))
        self.assertEqual(length, 9)

    def test_encode_decode_float_roundtrip(self):
        # Every float decode_value can yield must also be encodable
        for value in (0.0, 3.14, -1.5, 1e300, float('inf'), float('-inf')):
            table = {'k': value}
            pieces = []
            data.encode_table(pieces, table)
            decoded, _ = data.decode_table(b''.join(pieces), 0)
            self.assertEqual(decoded, table)

    def test_encode_decode_nan_roundtrip(self):
        # NaN is a float decode_value can yield too, but it is never equal to
        # itself, so it cannot ride the assertEqual loop above
        pieces = []
        data.encode_table(pieces, {'k': float('nan')})
        decoded, _ = data.decode_table(b''.join(pieces), 0)
        self.assertTrue(math.isnan(decoded['k']))

    def test_reencode_wire_float(self):
        # A 32-bit b'f' field off the wire must survive a decode/encode cycle,
        # which is what forwarding a message with its headers does
        encoded = b'\x00\x00\x00\x07\x01kf' + struct.pack('>f', 1.5)
        decoded, _ = data.decode_table(encoded, 0)
        pieces = []
        data.encode_table(pieces, decoded)
        redecoded, _ = data.decode_table(b''.join(pieces), 0)
        self.assertEqual(redecoded, decoded)

    def test_decode_value_long_string_invalid_utf8(self):
        # b'S' with non-UTF-8 content stays as bytes
        raw = b'\xff\xfe'
        encoded = b'\x00\x00\x00\x09\x01kS' + struct.pack('>I', len(raw)) + raw
        result, _ = data.decode_table(encoded, 0)
        self.assertIsInstance(result['k'], bytes)
        self.assertEqual(result['k'], raw)

    def test_decode_value_timestamp_max_representable(self):
        encoded = b'\x00\x00\x00\x0b\x01kT' + struct.pack('>Q', 253402300799)
        result, _ = data.decode_table(encoded, 0)
        self.assertEqual(
            result['k'],
            datetime.datetime(9999,
                              12,
                              31,
                              23,
                              59,
                              59,
                              tzinfo=datetime.timezone.utc))

    def test_decode_value_timestamp_out_of_datetime_range(self):
        # `timestamp` is a u64 (AMQP 0-9-1 errata), so the wire carries values
        # `datetime` cannot hold; they must not raise out of the decoder.
        values = (253402300800, 1772000000000, 2**63 - 1, 2**64 - 1)
        decoded = {}
        for seconds in values:
            encoded = (b'\x00\x00\x00\x0b\x01kT' + struct.pack('>Q', seconds))
            result, offset = data.decode_table(encoded, 0)
            self.assertEqual(offset, 15)
            decoded[seconds] = result['k']
        self.assertEqual(decoded, {seconds: seconds for seconds in values})
