#include <unistd.h>
#include <stdio.h>
#include <stdlib.h>
#include <stdbool.h>
#include <string.h>
#include <fcntl.h>
#include <linux/fb.h>
 #include <sys/ipc.h>
#include <sys/shm.h>
#include <sys/types.h> 
#include <sys/ioctl.h> 
#include <errno.h>
#include <lua.h>
#include <lauxlib.h>
#include <lualib.h>

#define IPCKEY 0x437800
#define MAX_TERM_CNT 100
#define REG_ADDR_MASK 0x1000 //寄存器地址范围 0xfff
#define SHM_SIZE REG_ADDR_MASK*2 //寄存器地址范围 0xfff

typedef struct
{
    unsigned short *m_data;
    bool connected;
    int shm_id;
} mmap_t;

mmap_t conn_list[MAX_TERM_CNT] = {0};

int l_mmap(int term_addr,int lenght)
{
    int shm_id;
    key_t key; 
    void *m_data;
    key = ftok("/tmp",term_addr);
    // key = IPCKEY | term_addr;
    shm_id=shmget(key,lenght,IPC_CREAT|IPC_EXCL|0600); 
    if(shm_id==-1 && errno == EEXIST) {
        shm_id=shmget(key,lenght,0); 
    }
    if (shm_id == -1)
    {
        char tmp[256];
        snprintf(tmp,sizeof(tmp),"shmget error=%d,key:%d term_addr:%d size:%d",errno,key,term_addr,SHM_SIZE);
        perror(tmp);
        return -1;
    }
    // printf("shm_id=%d\n", shm_id);
    m_data = (unsigned short *)shmat(shm_id, NULL, 0);
    if(m_data == NULL) return 0;
    conn_list[term_addr].shm_id = shm_id;
    conn_list[term_addr].m_data = m_data;  
    conn_list[term_addr].connected = true;   
    return 0;   
}
void l_unmmap(int term_addr)
{ 
    conn_list[term_addr].m_data = NULL;    
    conn_list[term_addr].connected = false;
}



static int l_connect(lua_State* L)
{
    int ntop, term_addr,lenght;

    lenght = SHM_SIZE;
    if (lua_checkstack(L, 2) == 0)
    return 0;

    // 检查栈
    ntop = lua_gettop(L);
    if (ntop < 1 ) {
        lua_pushnil(L);
        lua_pushstring(L, "invalid argument(s) for connect");
        return 2;
    }
    else if (ntop == 1 ) {
        if (lua_isinteger(L, 1) == 0)  {
            lua_pushnil(L);
            return 2;
        }
    }

    term_addr = (int) lua_tointeger(L, 1);
    if (ntop == 2) lenght = (int) lua_tointeger(L, 2);
    l_mmap(term_addr,lenght); 
    return 1;                           //告诉lua返回了一个变量
}

static int l_disconnect(lua_State* L)
{
    int ntop, term_addr;
    if (lua_checkstack(L, 2) == 0)
    return 0;

    // 检查栈
    ntop = lua_gettop(L);
    if (ntop < 1 ) {
        lua_pushnil(L);
        lua_pushstring(L, "invalid argument(s) for connect");
        return 2;
    }
    else if (ntop == 1 ) {
        if (lua_isinteger(L, 1) == 0)  {
            lua_pushnil(L);
            return 2;
        }
    }

    term_addr = (int) lua_tointeger(L, 1);
    l_unmmap(term_addr);
    return 1;                           //告诉lua返回了一个变量
}

static int l_shmget(lua_State* L)
{
    int ntop;
    unsigned int start_addr,bytes_cnt,term_addr;
    char *s = NULL; 
    if (lua_checkstack(L, 2) == 0)
    return 0;

    // 检查栈
    ntop = lua_gettop(L);
    if (ntop < 3 ||
        lua_isinteger(L, 1) == 0 ||
        lua_isinteger(L, 2) == 0 ||
        lua_isinteger(L, 3) == 0) {
        lua_pushnil(L);
        lua_pushstring(L, "invalid argument(s) ");
        return 2;
    }
    term_addr = (int) lua_tointeger(L, 1);
    start_addr = (int) lua_tointeger(L, 2);
    bytes_cnt = (int) lua_tointeger(L, 3);
    if(term_addr >= 100) {
        return 2;
    }
    if((start_addr*2 + bytes_cnt) >= (SHM_SIZE)) {
        return 2;
    }    
    if(!conn_list[term_addr].connected) {
        lua_pushnil(L);
        lua_pushstring(L, "term_addr not connected");
        return 2;
    }  
    s = (char *)&conn_list[term_addr].m_data[start_addr];  
    lua_pushlstring (L, s, bytes_cnt);
    return 1;                           //告诉lua返回了一个变量
}

//TODO: 增加format参数
static int l_shmget_single(lua_State* L)
{
    int ntop;
    short value;
    unsigned int start_addr,term_addr;
    char *s = NULL; 
    if (lua_checkstack(L, 2) == 0)
    return 0;

    // 检查栈
    ntop = lua_gettop(L);
    if (ntop < 2 ||
        lua_isinteger(L, 1) == 0 ||
        lua_isinteger(L, 2) == 0)  {
        lua_pushnil(L);
        lua_pushstring(L, "invalid argument(s) ");
        return 2;
    }
    term_addr = (int) lua_tointeger(L, 1);
    start_addr = (int) lua_tointeger(L, 2);
    if(term_addr >= 100) {
        return 2;
    }
    if((start_addr*2 + 2) >= (SHM_SIZE)) {
        return 2;
    }    
    if(!conn_list[term_addr].connected) {
        lua_pushnil(L);
        lua_pushstring(L, "term_addr not connected");
        return 2;
    }  
    s = (char *)&conn_list[term_addr].m_data[start_addr]; 
    value = (s[0] << 8) + s[1];
    lua_pushinteger(L, value);
    return 1;                           //告诉lua返回了一个变量
}

static int l_shmset(lua_State* L)
{
    int ntop;
    unsigned int start_addr,data_len,term_addr;
    char *data;
    char *s = NULL; 
    if (lua_checkstack(L, 2) == 0)
    return 0;

    // 检查栈
    ntop = lua_gettop(L);
    if (ntop < 3 ||
        lua_isinteger(L, 1) == 0 ||
        lua_isinteger(L, 2) == 0 ||
        lua_isstring(L, 3) == 0 ||
        lua_isinteger(L, 4) == 0) {
        lua_pushnil(L);
        lua_pushstring(L, "invalid argument(s) ");
        return 2;
    }
    term_addr = (int) lua_tointeger(L, 1);
    start_addr = (int) lua_tointeger(L, 2);
    data = lua_tostring (L, 3);
    if(data == NULL)
    {
        lua_pushnil(L);
        lua_pushstring(L, "not valid data");
    }
    data_len = (int) lua_tointeger(L, 4);
    if(term_addr >= 100) {
        lua_pushnil(L);
        lua_pushstring(L, "invalid argument(s) ");
        return 2;
    }
    if((start_addr*2 + data_len) >= (SHM_SIZE)) {
                lua_pushnil(L);
        lua_pushstring(L, "length is exceeded");
        return 2;
    }    
    if(!conn_list[term_addr].connected) {
        lua_pushnil(L);
        lua_pushstring(L, "term_addr not connected");
        return 2;
    }  
    s = (char *)&conn_list[term_addr].m_data[start_addr];  
    printf("shm_id=%p,%p,%d\n", s,data,data_len);
    memcpy(s,data,data_len);
    return 0;                           //告诉lua返回了一个变量
}

//映射表，"shmget"为lua中的函数名，l_shmget为真正C中的函数地址
static const struct luaL_Reg shm[] = {
    {"shmget", l_shmget},
    {"shmsingle", l_shmget_single},
    {"shmset", l_shmset},
    {"connect", l_connect},
    {"disconnect", l_disconnect},
    {NULL, NULL},
};
 
//模块注册函数
int luaopen_shmem(lua_State* lua)
{
    //注册本模块中所有的功能函数，libshmem为模块名，shm数组存储所有函数的映射关系
    luaL_register(lua, "shmem", shm);
    return 1;
}
