#include <stdio.h>
#include <stdlib.h>
#include <errno.h>
#include <string.h>
#include <sys/types.h>
#include <netinet/in.h>
#include <sys/socket.h>
#include <sys/wait.h>
#include <unistd.h>
#include <arpa/inet.h>
#include <fcntl.h>
#include <sys/epoll.h>
#include <sys/time.h>
#include <sys/resource.h>
#include <time.h>
#include <netinet/tcp.h>
#include <sys/socket.h>

#include "common.h"
#include <hiredis/hiredis.h>

#ifdef MAXEPOLLSIZE
#undef MAXEPOLLSIZE
#endif

redisContext *context;

static bool lnxall_modbus_reply(var_t *var, modbus_t *ctx, const uint8_t *req, int req_length, int header_length)
{
    modbus_mapping_t *mb_mapping = NULL;
    unsigned int offset;
    int functioncode;
    int slave_id;
    uint16_t address;
    bool writeback = false;
    char cmd[CMD_MAX_LENGTH];
    char topic[TOPIC_MAX_LEN] = {0};
    char ptaddr[8] = {0};
    char ptval[8] = {0};
    int readnum = 0;
    unsigned char *data;
    int isreg = 0;
    time_t now = time(NULL);
    //cJSON *point;
    cJSON *rt_data = cJSON_CreateObject();
    cJSON *pts = cJSON_CreateObject();

    slave_id = req[header_length - 1];
    address = (req[header_length + 1] << 8) + req[header_length + 2];
    functioncode = req[header_length];
    switch(functioncode)
    {
        case MODBUS_FC_READ_COILS:
            readnum = (req[header_length + 3] << 8) + req[header_length + 4];
            mb_mapping = modbus_mapping_new_start_address(address,readnum,0,0,0,0,0,0);
            data = mb_mapping->tab_bits;
            break;
        case MODBUS_FC_READ_DISCRETE_INPUTS:
            readnum = (req[header_length + 3] << 8) + req[header_length + 4];
            mb_mapping = modbus_mapping_new_start_address(0,0,address,readnum,0,0,0,0);
            data = mb_mapping->tab_input_bits;
            break;
        case MODBUS_FC_READ_HOLDING_REGISTERS:
            readnum = (req[header_length + 3] << 8) + req[header_length + 4];
            mb_mapping = modbus_mapping_new_start_address(0,0,0,0,address,readnum,0,0);
            data = (unsigned char *)mb_mapping->tab_registers;
            isreg = 1;
            break;
        case MODBUS_FC_READ_INPUT_REGISTERS:
            readnum = (req[header_length + 3] << 8) + req[header_length + 4];
            mb_mapping = modbus_mapping_new_start_address(0,0,0,0,0,0,address,readnum);
            data = (unsigned char *)mb_mapping->tab_input_registers;
            isreg = 1;
            break;
        case MODBUS_FC_WRITE_SINGLE_COIL:
            writeback = true;
            readnum = 1;
            mb_mapping = modbus_mapping_new_start_address(address,readnum,0,0,0,0,0,0);
            data = mb_mapping->tab_bits;
            //point = cJSON_CreateObject();
            int coildata = (req[header_length + 3] << 8) + req[header_length + 4];
            snprintf(ptaddr, 8, "%u", address);
            snprintf(ptval, 8, "%d", coildata ? 1 : 0);
            cJSON_AddStringToObject(pts, ptaddr, ptval);
            //cJSON_AddItemToArray(pts, point);
            break;
        case MODBUS_FC_WRITE_SINGLE_REGISTER:
            writeback = true;
            readnum = 1;
            mb_mapping = modbus_mapping_new_start_address(0,0,0,0,address,readnum,0,0);
            data = (unsigned char *)mb_mapping->tab_registers;
            isreg = 1;
            //point = cJSON_CreateObject();
            short regdata = (req[header_length + 3] << 8) + req[header_length + 4];
            snprintf(ptaddr, 8, "%u", address);
            snprintf(ptval, 8, "%d", regdata);
            cJSON_AddStringToObject(pts, ptaddr, ptval);
            //cJSON_AddItemToArray(pts, point);
            break;
        case MODBUS_FC_WRITE_MULTIPLE_COILS:
            writeback = true;
            readnum = (req[header_length + 3] << 8) + req[header_length + 4];
            mb_mapping = modbus_mapping_new_start_address(address,readnum,0,0,0,0,0,0);
            data = mb_mapping->tab_bits;
            int bitcount = (req[header_length + 3] << 8) + req[header_length + 4];
            //int bytecount = req[header_length + 5];
            for (int i = 0; i < bitcount; i++)
            {
                //point = cJSON_CreateObject();
                snprintf(ptaddr, 8, "%u", address + i);
                snprintf(ptval, 8, "%d", ((req[header_length + 6 + i/8]) & (1 << (i%8))) ? 1 : 0 );
                cJSON_AddStringToObject(pts, ptaddr, ptval);
                //cJSON_AddItemToArray(pts, point);
            }
            break;
        case MODBUS_FC_WRITE_MULTIPLE_REGISTERS:
            writeback = true;
            readnum = (req[header_length + 3] << 8) + req[header_length + 4];
            mb_mapping = modbus_mapping_new_start_address(0,0,0,0,address,readnum,0,0);
            data = (unsigned char *)mb_mapping->tab_registers;
            isreg = 1;
            int wordcount = (req[header_length + 3] << 8) + req[header_length + 4];
            //int bytecount = req[header_length + 5];
            for (int i = 0; i < wordcount; i++)
            {
                //point = cJSON_CreateObject();
                snprintf(ptaddr, 8, "%u", address + i);
                snprintf(ptval, 8, "%d", (req[header_length + 6 + 2 * i] << 8) + req[header_length + 7 + 2 * i]);
                cJSON_AddStringToObject(pts, ptaddr, ptval);
                //cJSON_AddItemToArray(pts, point);
            }
            break;
        case MODBUS_FC_MASK_WRITE_REGISTER:
            writeback = true;
            readnum = 1;
            mb_mapping = modbus_mapping_new_start_address(0,0,0,0,address,readnum,0,0);
            data = (unsigned char *)mb_mapping->tab_registers;
            isreg = 1;
            break;
        case MODBUS_FC_WRITE_AND_READ_REGISTERS:
            writeback = true;
            readnum = (req[header_length + 3] << 8) + req[header_length + 4];
            uint16_t address_write = (req[header_length + 5] << 8) + req[header_length + 6];
            int writenum = (req[header_length + 7] << 8) + req[header_length + 8];
            uint16_t startaddress = address < address_write ? address: address_write;
            int tnum = ((address+readnum) > (address_write+writenum)) ? (address+readnum) : (address_write+writenum);
            tnum -= startaddress;
            mb_mapping = modbus_mapping_new_start_address(0,0,0,0,startaddress,tnum,0,0);
            data = (unsigned char *)mb_mapping->tab_registers;
            address = startaddress;
            readnum = tnum;
            isreg = 1;
            break;
        case MODBUS_FC_REPORT_SLAVE_ID:
        case MODBUS_FC_READ_EXCEPTION_STATUS:
        default:
            readnum = 1;
            mb_mapping = modbus_mapping_new_start_address(0,0,0,0,address,readnum,0,0);
            data = (unsigned char *)mb_mapping->tab_registers;
            isreg = 1;
            break;
    }
    dy_syslog(LOG_DEBUG, "data %d readnum %d", address, readnum);
    if (writeback)
    {
        cJSON_AddStringToObject(rt_data, "proto", "modbus");
        cJSON_AddNumberToObject(rt_data, "caddr", slave_id);
        cJSON_AddStringToObject(rt_data, "identifier", "setems");
        cJSON_AddNumberToObject(rt_data, "time", time(NULL));
        cJSON_AddItemToObject(rt_data, "pts", pts);
        
        char *data_tmp = cJSON_Print(rt_data);
        snprintf(topic, TOPIC_MAX_LEN, "ipc/%s/lnxall/device/setems", var->sn_str);
        dy_syslog(LOG_ERR, "publish %s payload %s str %d", topic, data_tmp, strlen(data_tmp));
        ipc_session_publish(var->session, topic, data_tmp, strlen(data_tmp));
        free(data_tmp);
    }
    else 
    {
        for (int i = 0; i < readnum; )
        {
            int val = 0;
            float fval;
            snprintf(cmd, CMD_MAX_LENGTH, "zrangebyscore %d_modbus %u %u", slave_id, address + i, address + i);
            redisReply *reply = (redisReply*)redisCommand(context, cmd);
            if (NULL == reply)
            {
                dy_syslog(LOG_DEBUG, "%d execcmd %s error", __LINE__, cmd);
                i++;
                continue;
            }
            if (reply->elements < 1 || reply->element[0]->str == NULL)
            {
                dy_syslog(LOG_DEBUG, "execcmd error %s\n", cmd);
                freeReplyObject(reply);
                i++;
                continue;
            }
            else
            {
                char *tok = strtok(reply->element[0]->str, "|");
                int dtype = atoi(strtok(0, "|"));
                switch (dtype)
                {
                case REDIS_DATA_U8:
                case REDIS_DATA_U16:
                case REDIS_DATA_I16:
                    val = atoi(tok);
                    *((uint16_t *)(data + 2 * i)) = (uint16_t)val;
                    i++;
                    break;
                case REDIS_DATA_U32:
                case REDIS_DATA_I32:
                    val = atoi(tok);
                    *((int *)(data + 2 * i)) = val;
                    uint16_t tmp = *(uint16_t *)(data+2*i);
                    *((uint16_t*)(data+2*i)) = *((uint16_t*)(data + 2*i +2));
                    *((uint16_t*)(data + 2*i +2)) = tmp;
                    i+=2;
                    break;
                case REDIS_DATA_F32:
                    fval = atof(tok);
                    *((float *)(data + 2 * i)) = fval;
                    i+=2;
                    break;
                default:
                    i++;
                    break;
                }
            }
            freeReplyObject(reply);
        }
    }
    if(modbus_reply(ctx, req, req_length, mb_mapping) != -1)
    {
        //do nothing.
    }

    if (mb_mapping != NULL)
        modbus_mapping_free(mb_mapping);

    cJSON_Delete(rt_data);
    return true;
}

/*
handle_message - 处理每个 socket 上的消息收发
*/
static int handle_message(var_t *var, int socket, int header_length)
{
    uint8_t query[MODBUS_TCP_MAX_ADU_LENGTH];
    char cmd[CMD_MAX_LENGTH] = {0};
    int rc = 0;
    int i;
    int port_index = 0;
    reg_data_t *reg_data = NULL;
    dy_syslog(LOG_DEBUG, "handle_message, socket is %d, header_length is %d", socket, header_length);
    modbus_set_socket(var->ctx, socket);
    rc = modbus_receive(var->ctx, query);
    if (rc > 0)
    {
        uint16_t reg_addr = MODBUS_GET_INT16_FROM_INT8(query, header_length + 1);
        int slave_id = query[header_length - 1];
        lnxall_modbus_reply(var, var->ctx, query, rc, header_length);
    }
    else if (rc == -1)
    {
        printf("rc:%d err:%s\n", rc, modbus_strerror(errno));
        return -2;
    }

    return 0;
}

/*
setnonblocking - 设置句柄为非阻塞方式
*/
static int setnonblocking(int sockfd)
{
    if (fcntl(sockfd, F_SETFL, fcntl(sockfd, F_GETFD, 0) | O_NONBLOCK) == -1)
    {
        return -1;
    }
    return 0;
}

/* Set TCP keep alive option to detect dead peers. The interval option
 * is only used for Linux as we are using Linux-specific APIs to set
 * the probe send time, interval, and count. */
int set_keep_alive(int fd)
{
    int val = 1;
    int interval = 10;

    //开启keepalive机制
    if (setsockopt(fd, SOL_SOCKET, SO_KEEPALIVE, &val, sizeof(val)) == -1)
    {
        dy_syslog(LOG_ERR, "setsockopt SO_KEEPALIVE: %s", strerror(errno));
        return -1;
    }

    /* Default settings are more or less garbage, with the keepalive time
     * set to 7200 by default on Linux. Modify settings to make the feature
     * actually useful. */

    /* Send first probe after interval. */
    val = interval;
    if (setsockopt(fd, IPPROTO_TCP, TCP_KEEPIDLE, &val, sizeof(val)) < 0)
    {
        dy_syslog(LOG_ERR, "setsockopt TCP_KEEPIDLE: %s\n", strerror(errno));
        return -1;
    }

    /* Send next probes after the specified interval. Note that we set the
     * delay as interval / 3, as we send three probes before detecting
     * an error (see the next setsockopt call). */
    val = interval / 3;
    if (val == 0)
        val = 1;
    if (setsockopt(fd, IPPROTO_TCP, TCP_KEEPINTVL, &val, sizeof(val)) < 0)
    {
        dy_syslog(LOG_ERR, "setsockopt TCP_KEEPINTVL: %s\n", strerror(errno));
        return -1;
    }

    /* Consider the socket in error state after three we send three ACK
     * probes without getting a reply. */
    val = 3;
    if (setsockopt(fd, IPPROTO_TCP, TCP_KEEPCNT, &val, sizeof(val)) < 0)
    {
        dy_syslog(LOG_ERR, "setsockopt TCP_KEEPCNT: %s\n", strerror(errno));
        return -1;
    }

    dy_syslog(LOG_DEBUG, "set keep alive success");

    return 0;
}

void *__modbus_tcp_slave_start(void *param)
{
    var_t *var = (var_t *)param;
    int server_socket = -1;
    struct epoll_event ev;
    struct epoll_event *events = NULL;
    int epfd, curfds;
    int new_fd, nfds, n, ret;
    int header_length;
    socklen_t len;
    struct sockaddr_in my_addr, their_addr;

    var->ctx = modbus_new_tcp(NULL, 502);

    header_length = modbus_get_header_length(var->ctx);

    server_socket = modbus_tcp_listen(var->ctx, NB_CONNECTION);
    if (server_socket == -1)
    {
        fprintf(stderr, "Unable to listen TCP connection\n");
        modbus_free(var->ctx);
        return (void *)-1;
    }

    printf("after listen\n");

    events = malloc((NB_CONNECTION + 1) * sizeof(struct epoll_event));
    if (events == NULL)
    {
        dy_syslog(LOG_INFO, "malloc failed");
        return (void *)-1;
    }
    /* 创建 epoll 句柄，把监听 socket 加入到 epoll 集合里 */
    epfd = epoll_create(NB_CONNECTION + 1);
    len = sizeof(struct sockaddr_in);
    ev.events = EPOLLIN;
    ev.data.fd = server_socket;
    if (epoll_ctl(epfd, EPOLL_CTL_ADD, server_socket, &ev) < 0)
    {
        dy_syslog(LOG_INFO, "epoll set insertion error: fd=%d", server_socket);
        free(events);
        close(server_socket);
        close(epfd);
        return (void *)-1;
    }
    else
    {
        dy_syslog(LOG_INFO, "listen socket add into epoll");
    }
    curfds = 1;

    while (1)
    {
        /* 等待有事件发生 */
        nfds = epoll_wait(epfd, events, NB_CONNECTION + 1, 10 * 1000);
        if (nfds == -1)
        {
            dy_syslog(LOG_INFO, "epoll_wait");
            break;
        }

        /* 处理所有事件 */
        for (n = 0; n < nfds; ++n)
        {
            if (events[n].data.fd == server_socket)
            {
                dy_syslog(LOG_DEBUG, "listener receive event");
                new_fd = accept(server_socket, (struct sockaddr *)&their_addr, &len);
                if (new_fd < 0)
                {
                    dy_syslog(LOG_INFO, "accept");
                    continue;
                }
                else
                {
                    dy_syslog(LOG_INFO, "connection request from %s:%d, allocated socket is:%d",
                              inet_ntoa(their_addr.sin_addr), ntohs(their_addr.sin_port), new_fd);
                }
                setnonblocking(new_fd);
                set_keep_alive(new_fd);
                ev.events = EPOLLIN | EPOLLRDHUP | EPOLLERR | EPOLLHUP;
                ev.data.fd = new_fd;
                if (epoll_ctl(epfd, EPOLL_CTL_ADD, new_fd, &ev) < 0)
                {
                    dy_syslog(LOG_INFO, "add socket '%d' into epoll failed%s",
                              new_fd, strerror(errno));
                    free(events);
                    close(server_socket);
                    close(epfd);
                    return (void *)-1;
                }
                curfds++;
            }
            else
            {
                if (events[n].events & (EPOLLERR | EPOLLRDHUP | EPOLLHUP))
                {
                    /* error on the fd, close the socket */
                    dy_syslog(LOG_INFO, "socket:%d leave", events[n].data.fd);
                    ret = -2;
                }
                else
                {
                    ret = handle_message(var, events[n].data.fd, header_length);
                }

                if (ret == -2)
                {
                    if (epoll_ctl(epfd, EPOLL_CTL_DEL, events[n].data.fd,
                                  &ev) < 0)
                    {
                        dy_syslog(LOG_INFO, "remove socket '%d' from epoll failed%s",
                                  events[n].data.fd, strerror(errno));
                    }
                    curfds--;
                    close(events[n].data.fd);
                }
            }
        }
    }
    return (void *)0;
}

int __init_modbus_tcp_slave(var_t *var)
{
    pthread_t mt;

    context = redisConnect(REDIS_SERVER_IP, REDIS_SERVER_PORT);
    if (context->err)
    {
        redisFree(context);
        printf("%d connect redis server failure: %s\n", __LINE__, context->errstr);
        return false;
    }
    dy_syslog(LOG_DEBUG, "connect redis server success.\n");

    pthread_create(&mt, NULL, __modbus_tcp_slave_start, (void *)var);
}