#include <sys/select.h>
#include <sys/time.h>
#include <sys/types.h>
#include <unistd.h>
#include <sys/syscall.h>
#include <stdlib.h>
#include <stdio.h>
#include <string.h>
#include <sys/stat.h>
#include <sys/types.h>
#include <errno.h>
#include <netinet/ip.h>
#include <netinet/in.h>
#include <net/ethernet.h>
#include "net_sniffer_common.h"


gw_port_e ports[] = {TCPDUMP};

// 建立与内部broker之间的MQTT连接
static int net_sniffer_mqtt_client_init(net_sniffer_var_t *var)
{
    char clientId[MAX_CLIENT_ID_LEN] = {0};

    snprintf(clientId, MAX_CLIENT_ID_LEN, "%s", var->proc_name);
    var->session = ipc_session_new(clientId, (void*)var, IPC_DEFAULT);
    if (var->session == NULL)
        return -1;
    ipc_session_start(var->session);
}

int net_sniffer_data_report(net_sniffer_var_t *var, char *sn, char *data, int len, unsigned char port)
{
    pp_real_time_data_t real_data = {0};
    real_data.port = port;
    real_data.len = len;
    real_data.instruction_code = data[1];
    real_data.data = malloc(real_data.len);

    real_data.mi = 0;
    strcpy(real_data.sn, sn);
    // strncpy(real_data.src_identifier, fields->src_identifier, sizeof(real_data.src_identifier));

    memcpy(real_data.data, data, real_data.len);

    send_to_proto_parser(var->session, &real_data);
    flush_device_status_by_sn(var->ds, ACTION_RCV, real_data.sn);

    free(real_data.data);

    return 0;
}

static int net_sniffer_handle_rcv_msg(net_sniffer_var_t *var, const char* ip_src,uint16_t th_sport,const char* ip_dst,uint16_t th_dport,const char *buf, int len)
{

    int mi = time(NULL);
    connect_config_t *connect_cfg = NULL;
    list_for_each_entry(connect_cfg, &var->connect_list, list)
    {
        if(strcmp(ip_src,connect_cfg->addr) == 0 && connect_cfg-> port == th_sport)
        {
            dy_syslog_hex(LOG_INFO,(void *)buf,len,"[%s:%d]->[%s:%d]",ip_src,th_sport,ip_dst,th_dport);
            net_sniffer_data_report(var, connect_cfg->sn, (void *)buf, len, TCPDUMP);
            break;
        }
    }
    // flush_device_status_by_sn(var->ds, ACTION_SND, connect_cfg->sn);
    return 0;
}

// static int connect_state_change(void *obj, socket_state_e state)
// {
//     connect_config_t *connect_cfg = (connect_config_t*)obj;
//     net_sniffer_var_t *var = connect_cfg->var;

//     dy_syslog(LOG_DEBUG, "connect_state_change state %d", state);

//     return 0;
// }

/** Get the index of an adapter by its network address
 *
 * @param netaddr network address of the adapter (e.g. 192.168.1.0)
 */
static int get_adapter_index_from_addr(net_sniffer_var_t *var,struct in_addr *netaddr)
{
    pcap_if_t *alldevs;
    pcap_if_t *d;
    char error_buffer[PCAP_ERRBUF_SIZE];
    int index = 0;



    /* Retrieve the interfaces list */
    if (pcap_findalldevs(&alldevs, error_buffer) == -1)
    {
        dy_syslog(LOG_ERR,"Error in pcap_findalldevs: %s\n", error_buffer);
        return -1;
    }
    /* Scan the list printing every entry */
    for (d = alldevs; d != NULL; d = d->next, index++)
    {
        pcap_addr_t *a;
        /* Scan the addresses and check the sub net*/
        for (a = d->addresses; a != NULL; a = a->next)
        {
            int len = 0;
            if (a->addr->sa_family == AF_INET)
            {
                uint32_t a_addr = ((struct sockaddr_in *)a->addr)->sin_addr.s_addr;
                uint32_t a_netmask = ((struct sockaddr_in *)a->netmask)->sin_addr.s_addr;
                uint32_t a_netaddr = a_addr & a_netmask;
                uint32_t addr = (*netaddr).s_addr & a_netmask;
                if (a_netaddr == addr) //matched and find the adapter device LAN
                {
                    dy_syslog(LOG_ERR,"netdev:%s ip:%s", d->name,inet_ntoa(((struct sockaddr_in *)a->addr)->sin_addr));
                    strcpy(var->adapter,d->name);
                    // pcap_freealldevs(alldevs);
                    return 0;
                }
            }
        }
    }
    dy_syslog(LOG_ERR, "Network address not found.\n");
    // pcap_freealldevs(alldevs);
    return -1;
}


void packet_handler(
    u_char *args,
    const struct pcap_pkthdr *header,
    const u_char *packet
)
{
        /* First, lets make sure we have an IP packet */
    struct ether_header *eth_header;
    eth_header = (struct ether_header *) packet;
    if (ntohs(eth_header->ether_type) != ETHERTYPE_IP) {
        dy_syslog(LOG_ERR,"Not an IP packet. Skipping...\n\n");
        return;
    }


    /* Pointers to start point of various headers */
    const u_char *ip_header;
    const u_char *tcp_header;
    const u_char *payload;
    struct ip *ip_frame;
    /* Header lengths in bytes */
    int ethernet_header_length = 14; /* Doesn't change */
    int ip_header_length;
    int tcp_header_length;
    int payload_length;
    
    char *ip_src,*ip_dst;
    uint16_t sport,dport;
    /* Find start of IP header */
    ip_header = packet + ethernet_header_length;

    ip_frame = (struct ip *)ip_header;
    ip_src = strdup(inet_ntoa(ip_frame->ip_src));
    ip_dst = strdup(inet_ntoa(ip_frame->ip_dst));
    /* The second-half of the first byte in ip_header
       contains the IP header length (IHL). */
    ip_header_length = ((*ip_header) & 0x0F);
    /* The IHL is number of 32-bit segments. Multiply
       by four to get a byte count for pointer arithmetic */
    ip_header_length = ip_header_length * 4;

    /* Now that we know where the IP header is, we can 
       inspect the IP header for a protocol number to 
       make sure it is TCP before going any further. 
       Protocol is always the 10th byte of the IP header */
    u_char protocol = *(ip_header + 9);
    if (protocol != IPPROTO_TCP) {
        // dy_syslog(LOG_ERR,"Not a TCP packet. Skipping...\n\n");
        free(ip_src);
        free(ip_dst);
        return;
    }

    /* Add the ethernet and ip header length to the start of the packet
       to find the beginning of the TCP header */
    tcp_header = packet + ethernet_header_length + ip_header_length;
    struct tcphdr *tp = (struct tcphdr *)tcp_header;
    sport = ntohs(tp->th_sport);
    dport = ntohs(tp->th_dport);
    /* TCP header length is stored in the first half 
       of the 12th byte in the TCP header. Because we only want
       the value of the top half of the byte, we have to shift it
       down to the bottom half otherwise it is using the most 
       significant bits instead of the least significant bits */
    tcp_header_length = ((*(tcp_header + 12)) & 0xF0) >> 4;
    /* The TCP header length stored in those 4 bits represents
       how many 32-bit words there are in the header, just like
       the IP header length. We multiply by four again to get a
       byte count. */
    tcp_header_length = tcp_header_length * 4;

    /* Add up all the header sizes to find the payload offset */
    int total_headers_size = ethernet_header_length+ip_header_length+tcp_header_length;
    payload_length = header->caplen -
        (ethernet_header_length + ip_header_length + tcp_header_length);
    payload = packet + total_headers_size;


    net_sniffer_handle_rcv_msg((void *)args,ip_src,sport,ip_dst,dport,(void *)payload,payload_length);
    free(ip_src);
    free(ip_dst);
    return;
}

static void net_sniffer_loop(net_sniffer_var_t *var)
{
    if (var->pcap && strlen(var->adapter) > 0)
    {
        pcap_loop(var->pcap,0,packet_handler,(void *)var);
    }
}

static void set_cature_params(net_sniffer_var_t *var,const char *filter_exp)
{
    char error_buffer[PCAP_ERRBUF_SIZE];
    struct bpf_program filter;
    bpf_u_int32 subnet_mask, ip;
    pcap_t *handle = pcap_open_live(var->adapter, BUFSIZ, 1, 1000, error_buffer);
    if (handle == NULL)
    {
        printf("Could not open %s - %s\n", var->adapter, error_buffer);
        return ;
    }
    var->pcap = handle;
    dy_syslog(LOG_INFO,"filter exp: %s\n", filter_exp);
    if (pcap_compile(var->pcap, &filter, filter_exp, 0, ip) == -1) {
        dy_syslog(LOG_ERR,"Bad filter - %s\n", pcap_geterr(var->pcap));
        return ;
    }
    if (pcap_setfilter(var->pcap, &filter) == -1) {
        dy_syslog(LOG_ERR,"Error setting filter - %s\n", pcap_geterr(var->pcap));
        return ;
    }
    return ;
}

int net_sniffer_load_tcp_nodes(net_sniffer_var_t *var)
{
    //获取网卡
    int i = 0;
    connect_config_t *connect_cfg = NULL;

    for (i = 0; i < var->nodes_cfg_table->node_cnt; i++)
    {
        if (strlen(var->nodes_cfg_table->node[i].app_key) == 0 && var->nodes_cfg_table->node[i].port == TCPDUMP && strlen(var->nodes_cfg_table->node[i].tcp_ip_addr) > 5)
        {
            connect_cfg = calloc(1, sizeof(connect_config_t));
            connect_cfg->var = var;
            strcpy(connect_cfg->sn, var->nodes_cfg_table->node[i].sn);
            connect_cfg->addr = strdup(var->nodes_cfg_table->node[i].tcp_ip_addr);
            connect_cfg->port = var->nodes_cfg_table->node[i].tcp_port;
            list_add_tail(&connect_cfg->list, &var->connect_list);
        }
    }
}

int net_sniffer_pcap_init(net_sniffer_var_t *var)
{
    connect_config_t *connect_cfg = NULL;
    list_for_each_entry(connect_cfg, &var->connect_list, list)
    {
        struct in_addr  addr;
        if (inet_aton(connect_cfg->addr,&addr) > 0)
        {
            if(get_adapter_index_from_addr(var,&addr) == 0) break; //Only supported that capture one netcard
        }
    }

    if(strlen(var->adapter) <= 0)
    {
        return -1;
    }
    char *filter_exp = calloc(1024,1) ;
    strcat(filter_exp,"(");
    list_for_each_entry(connect_cfg, &var->connect_list, list)
    {
        char buff[100] = {0};
        int ret = snprintf(buff,sizeof(buff), " (src port %d) ",connect_cfg->port);
        if(ret < 0 )
            dy_syslog(LOG_ERR,"host ip too long %s:%d\n", connect_cfg->addr,connect_cfg->port);
        else
        {
            if (strlen(filter_exp) > 1)//有至少存在一条过滤
            {
                strcat(filter_exp,"or");
            }
            strcat(filter_exp,buff);
        }
    }
    if (strlen(filter_exp) > 1)//有至少存在一条过滤
    {
        strcat(filter_exp,")");
        //目标格式 tcpdump '(((src port 80) and (src host 172.16.38.53)) or ((src port 99) and (src host 172.16.38.39)))'
        set_cature_params(var,filter_exp);
    }
    free(filter_exp);
}

static int net_sniffer_init(net_sniffer_var_t *var)
{
    INIT_LIST_HEAD(&var->connect_list);

    get_board_sn(var->sn_str);
    get_process_name(var->proc_name);

    if (load_nodes_cfg(&var->nodes_cfg_table, NODES_CFG_PATH) == -1)
    {
        dy_syslog(LOG_ERR, "load nodes cfg fail");
    }
    if (load_templates_cfg(&var->template_table, TEMPLATES_CFG_PATH) == -1)
    {
        dy_syslog(LOG_ERR, "load template cfg fail");
    }

    dev_status_init(&var->ds, var->nodes_cfg_table, var->template_table, ports, ARRAY_SIZE(ports));
    node_sta_recovery(var->ds, PCAP_NODE_STATUS_BAK_FILE);

    kv_array_init(&var->identifier_backup, 32);

    net_sniffer_mqtt_client_init(var);
    net_sniffer_load_tcp_nodes(var);
    net_sniffer_pcap_init(var);
    dy_syslog(LOG_INFO, "init done, board SN:%s", var->sn_str);

    return 0;
}

int main(int argc, char *argv[])
{
    net_sniffer_var_t var = {0};

    openlog("net_snifder", LOG_PID, LOG_DAEMON);

    memset(&var, 0, sizeof(net_sniffer_var_t));

    net_sniffer_init(&var);
    while (1)  net_sniffer_loop(&var);

    return 0;
}
