#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"

#ifdef MAXEPOLLSIZE
#undef MAXEPOLLSIZE
#endif

/*
handle_message - 处理每个 socket 上的消息收发
*/
static int handle_message(var_t *var, int socket, int header_length)
{
    uint8_t query[MODBUS_TCP_MAX_ADU_LENGTH];
    int rc = 0;
    int i;
    int port_index = 0;
    reg_data_t *reg_data = NULL;

    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];

        for (i = 0; i < var->nodes_cfg_table->node_cnt; i++)
        {
            node_data_t *node = &var->node[i];
            if (slave_id == node->reg_data.term_addr)
            {
                if (node->reg_data.mb_mapping != NULL)
                    modbus_reply(var->ctx, query, rc, node->reg_data.mb_mapping);
                break;
            }
        }
    }
    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;

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