from trex_stl_lib.api import *


#
# Physical / outer Ethernet
#
DST_MAC = "b8:3f:d2:1c:61:09"
SRC_MAC = "00:11:17:40:AD:4F"
OUTER_VLAN = 21

#
# Burst parameters
#
#PKTS_PER_BURST = 100000
#IBG_US = 100000.0       # 100 ms
#RATE_PERCENT = 100
#BURST_COUNT = 100000

PKTS_PER_BURST = 25000
IBG_US = 2000000.0      # 1 second
RATE_PERCENT = 100
BURST_COUNT = 100000

def chain(records):
    """Chain Scapy DNSRR objects together."""
    if not records:
        return None

    result = records[0]

    for rr in records[1:]:
        result = result / rr

    return result


class STLSuriMicroburst(object):

    def build_dns(self):

        answers = []

        #
        # Many A records
        #
        for i in range(10):
            answers.append(
                DNSRR(
                    rrname="bulk.suri.test.",
                    type="A",
                    ttl=300,
                    rdata="10.20.0.{}".format(i + 1)
                )
            )

        #
        # Several AAAA records
        #
        for i in range(4):
            answers.append(
                DNSRR(
                    rrname="bulk.suri.test.",
                    type="AAAA",
                    ttl=300,
                    rdata="2001:db8::{:x}".format(i + 1)
                )
            )

        #
        # CNAME chain
        #
        answers += [
            DNSRR(
                rrname="alias1.suri.test.",
                type="CNAME",
                ttl=300,
                rdata="alias2.suri.test."
            ),

            DNSRR(
                rrname="alias2.suri.test.",
                type="CNAME",
                ttl=300,
                rdata="bulk.suri.test."
            ),
        ]

        #
        # TXT records with substantial payload.
        #
        for i in range(3):
            txt = (
                "suricata-microburst-record-{:02d}-".format(i)
                + ("X" * 128)
            )

            answers.append(
                DNSRR(
                    rrname="txt{}.suri.test.".format(i),
                    type="TXT",
                    ttl=300,
                    rdata=txt.encode("ascii")
                )
            )

        #
        # Authority section
        #
        authority = [
            DNSRR(
                rrname="suri.test.",
                type="NS",
                ttl=600,
                rdata="ns1.suri.test."
            ),

            DNSRR(
                rrname="suri.test.",
                type="NS",
                ttl=600,
                rdata="ns2.suri.test."
            ),
        ]

        #
        # Additional/glue records
        #
        additional = [
            DNSRR(
                rrname="ns1.suri.test.",
                type="A",
                ttl=600,
                rdata="10.30.0.1"
            ),

            DNSRR(
                rrname="ns2.suri.test.",
                type="A",
                ttl=600,
                rdata="10.30.0.2"
            ),

            DNSRR(
                rrname="ns1.suri.test.",
                type="AAAA",
                ttl=600,
                rdata="2001:db8:30::1"
            ),

            DNSRR(
                rrname="ns2.suri.test.",
                type="AAAA",
                ttl=600,
                rdata="2001:db8:30::2"
            ),
        ]

        dns = DNS(
            id=0x4242,

            qr=1,       # response
            aa=1,       # authoritative
            rd=1,
            ra=1,

            qdcount=1,
            ancount=len(answers),
            nscount=len(authority),
            arcount=len(additional),

            qd=DNSQR(
                qname="bulk.suri.test.",
                qtype="ANY"
            ),

            an=chain(answers),
            ns=chain(authority),
            ar=chain(additional),
        )

        return dns


    def get_streams(self, tunables, **kwargs):

        dns = self.build_dns()

        #
        # Inner tenant packet.
        #
        # Another VLAN is deliberately included INSIDE VXLAN.
        #
        inner = (
            Ether(
                src="02:00:00:00:00:53",
                dst="02:00:00:00:00:10"
            )
            / Dot1Q(
                vlan=301
            )
            / IP(
                src="10.10.0.53",
                dst="10.10.0.10",
                ttl=63
            )
            / UDP(
                sport=53,
                dport=53000
            )
            / dns
        )

        #
        # Physical packet:
        #
        # Ethernet
        #   -> VLAN 21
        #      -> IPv4
        #         -> UDP/4789
        #            -> VXLAN
        #               -> Ethernet
        #                  -> VLAN 301
        #                     -> IPv4
        #                        -> UDP
        #                           -> DNS
        #
        pkt = (
            Ether(
                src=SRC_MAC,
                dst=DST_MAC
            )
            / Dot1Q(
                vlan=OUTER_VLAN
            )
            / IP(
                src="198.18.0.1",
                dst="198.18.0.2",
                ttl=64
            )
            / UDP(
                sport=55000,
                dport=4789
            )
            / VXLAN(
                vni=4242
            )
            / inner
        )

        #
        # Standard Ethernet MTU sanity check.
        #
        # Scapy/TRex packet length excludes Ethernet FCS.
        #
        packet_len = len(pkt)

        print(
            "Heavy Suricata packet: {} bytes without FCS, "
            "{} bytes including FCS".format(
                packet_len,
                packet_len + 4
            )
        )

        if packet_len > 1514:
            raise Exception(
                "Packet is {} bytes without FCS; exceeds normal "
                "1518-byte Ethernet frame size".format(packet_len)
            )

        return [
            STLStream(
                name="suricata_heavy_dns_microburst",

                packet=STLPktBuilder(
                    pkt=pkt
                ),

                mode=STLTXMultiBurst(
                    percentage=RATE_PERCENT,
                    pkts_per_burst=PKTS_PER_BURST,
                    ibg=IBG_US,
                    count=BURST_COUNT,
                )
            )
        ]


def register():
    return STLSuriMicroburst()
