#include <stdio.h>
#include <stdarg.h>
#include <stdlib.h>
#include <string.h>
#include <unistd.h>
#include <signal.h>
#include <sys/socket.h>
#include <netinet/in.h>
#include <arpa/inet.h>
#include <sys/select.h>
#include <sys/wait.h>
#include <fcntl.h>
#include <errno.h>

#define STB_ARGPARSE_IMPLEMENTATION
#include "stb_argparse.h"
#include "fort.h"

#define MAX_SESSIONS 256
#define MAX_CLIENTS 64
#define BUFFER_SIZE 4096
#define LISTEN_PORT 9999

// 会话结构：代表一个子进程
typedef struct {
    int id;               // 会话ID (1-based)
    pid_t pid;            // 子进程PID
    int stdin_fd;         // 写入子进程stdin的管道写端
    int stdout_fd;        // 读取子进程stdout的管道读端
    int owner_fd;         // 当前拥有该会话的客户端socket，-1表示空闲
    int valid;
} Session;

static Session sessions[MAX_SESSIONS];
static int next_session_id = 1;

// 客户端连接结构
typedef struct {
    int sockfd;
    char inbuf[BUFFER_SIZE];
    size_t inlen;
    int attached_session; // 当前附加的会话ID，-1表示未附加
} Client;

static Client clients[MAX_CLIENTS];
static int listen_fd = -1;

// 初始化
static void init() {
    for (int i = 0; i < MAX_SESSIONS; i++) {
        sessions[i].valid = 0;
    }
    for (int i = 0; i < MAX_CLIENTS; i++) {
        clients[i].sockfd = -1;
        clients[i].attached_session = -1;
    }
    signal(SIGPIPE, SIG_IGN);
    signal(SIGCHLD, SIG_IGN); // 简单处理，用waitpid非阻塞回收
}

// 查找空闲会话槽
static int find_free_session() {
    for (int i = 1; i < MAX_SESSIONS; i++) {
        if (!sessions[i].valid) return i;
    }
    return -1;
}

// 创建新会话，执行程序
static int create_session(char *prog, char **args) {
    int idx = find_free_session();
    if (idx == -1) return -1;

    int stdin_pipe[2], stdout_pipe[2];
    if (pipe(stdin_pipe) < 0 || pipe(stdout_pipe) < 0) {
        perror("pipe");
        return -1;
    }

    pid_t pid = fork();
    if (pid < 0) {
        perror("fork");
        close(stdin_pipe[0]); close(stdin_pipe[1]);
        close(stdout_pipe[0]); close(stdout_pipe[1]);
        return -1;
    }
    if (pid == 0) {
        // 子进程
        close(stdin_pipe[1]);  // 关闭写端
        close(stdout_pipe[0]); // 关闭读端
        dup2(stdin_pipe[0], STDIN_FILENO);
        dup2(stdout_pipe[1], STDOUT_FILENO);
        dup2(stdout_pipe[1], STDERR_FILENO); // 也合并stderr到stdout，方便
        close(stdin_pipe[0]);
        close(stdout_pipe[1]);

        // 执行程序
        execvp(prog, args);
        perror("execvp");
        exit(1);
    }
    // 父进程
    close(stdin_pipe[0]);  // 关闭读端
    close(stdout_pipe[1]); // 关闭写端
    // 设置为非阻塞？可选
    fcntl(stdout_pipe[0], F_SETFL, O_NONBLOCK); // 非阻塞读取，防止死锁

    sessions[idx].id = idx;
    sessions[idx].pid = pid;
    sessions[idx].stdin_fd = stdin_pipe[1];
    sessions[idx].stdout_fd = stdout_pipe[0];
    sessions[idx].owner_fd = -1;
    sessions[idx].valid = 1;

    return idx;
}

// 关闭会话
static void close_session(int sid) {
    if (!sessions[sid].valid) return;
    // 终止子进程
    kill(sessions[sid].pid, SIGTERM);
    // 等待？为了简单，我们忽略，让SIGCHLD处理
    close(sessions[sid].stdin_fd);
    close(sessions[sid].stdout_fd);
    sessions[sid].valid = 0;
}

// 查找会话通过ID
static Session* get_session(int sid) {
    if (sid < 1 || sid >= MAX_SESSIONS) return NULL;
    if (!sessions[sid].valid) return NULL;
    return &sessions[sid];
}

// 向客户端发送字符串
static void client_printf(int fd, const char *fmt, ...) {
    char buf[4096];
    va_list args;
    va_start(args, fmt);
    vsnprintf(buf, sizeof(buf), fmt, args);
    va_end(args);
    send(fd, buf, strlen(buf), 0);
}

// 处理客户端在命令模式下的输入
static void handle_command_mode(Client *client, char *line) {
    int sock = client->sockfd;
    char *cmd = strtok(line, " \t");
    if (!cmd) return;

    if (strcmp(cmd, "quit") == 0 || strcmp(cmd, "exit") == 0) {
        client_printf(sock, "Goodbye\n");
	shutdown(sock, SHUT_RDWR);
        close(sock);
        client->sockfd = -1;
        return;
    }
    else if (strcmp(cmd, "new") == 0) {
        char *prog = strtok(NULL, " \t");
        if (!prog) {
            client_printf(sock, "Usage: new <program> [args...]\n");
            return;
        }
        // 收集参数
        char *args[64];
        int argc = 0;
        args[argc++] = prog;
        char *arg;
        while ((arg = strtok(NULL, " \t")) != NULL && argc < 63) {
            args[argc++] = arg;
        }
        args[argc] = NULL;

        int sid = create_session(prog, args);
        if (sid == -1) {
            client_printf(sock, "Failed to create session\n");
        } else {
            client_printf(sock, "s%x\n", sid);
        }
    }
    else if (strcmp(cmd, "attach") == 0) {
        char *sid_str = strtok(NULL, " \t");
        if (!sid_str) {
            client_printf(sock, "Usage: attach <session-id>\n");
            return;
        }
        char *endptr;
        long sid = strtol(sid_str + 1, &endptr, 16);
        if (*endptr != '\0' || sid <= 0) {
            client_printf(sock, "Invalid session ID\n");
            return;
        }
        Session *sess = get_session((int)sid);
        if (!sess) {
            client_printf(sock, "Session not found\n");
            return;
        }
        if (sess->owner_fd != -1) {
            client_printf(sock, "Session already in use\n");
            return;
        }
        // 将会话所有权给该客户端
        sess->owner_fd = sock;
        client->attached_session = (int)sid;
        client_printf(sock, "Attached to session s%x. Type '~detach' to detach.\n", (int)sid);
    }
    else if (strcmp(cmd, "detach") == 0) {
        if (client->attached_session == -1) {
            client_printf(sock, "Not attached to any session\n");
        } else {
            Session *sess = get_session(client->attached_session);
            if (sess && sess->owner_fd == sock) {
                sess->owner_fd = -1;
            }
            client->attached_session = -1;
            client_printf(sock, "Detached\n");
        }
    }
    else if (strcmp(cmd, "kill") == 0) {
        char *sid_str = strtok(NULL, " \t");
        if (!sid_str) {
            client_printf(sock, "Usage: kill <session-id>\n");
            return;
        }
        char *endptr;
        long sid = strtol(sid_str + 1, &endptr, 16);
        if (*endptr != '\0' || sid <= 0) {
            client_printf(sock, "Invalid session ID\n");
            return;
        }
        if (!get_session((int)sid)) {
            client_printf(sock, "Session not found\n");
            return;
        }
        close_session((int)sid);
        client_printf(sock, "Session killed\n");
    }
    else if (strcmp(cmd, "list") == 0) {
        client_printf(sock, "Active sessions:\n");
        for (int i = 1; i < MAX_SESSIONS; i++) {
            if (sessions[i].valid) {
                client_printf(sock, "  s%x (pid %d)\n", i, sessions[i].pid);
            }
        }
    }
    else {
        client_printf(sock, "Unknown command: %s\n", cmd);
    }
}

// 处理客户端在附加模式下的输入（直接转发给子进程）
static void handle_attached_mode(Client *client, char *line) {
    int sid = client->attached_session;
    Session *sess = get_session(sid);
    if (!sess || sess->owner_fd != client->sockfd) {
        // 会话已失效，退回命令模式
        client->attached_session = -1;
        client_printf(client->sockfd, "\nSession lost, returning to command mode.\n");
        return;
    }
    // 检查是否为detach命令
    if (strcmp(line, "~detach") == 0) {
        sess->owner_fd = -1;
        client->attached_session = -1;
        client_printf(client->sockfd, "\nDetached from session.\n");
        return;
    }
    // 否则将整行写入子进程stdin，加上换行符
    char *buf = malloc(strlen(line) + 2);
    sprintf(buf, "%s\n", line);
    write(sess->stdin_fd, buf, strlen(buf));
    free(buf);
}

// 处理客户端可读事件
static void handle_client_read(Client *client) {
    char buf[1024];
    ssize_t n = read(client->sockfd, buf, sizeof(buf)-1);
    if (n <= 0) {
        // 客户端断开
        if (n == 0) {
            printf("Client %d disconnected\n", client->sockfd);
        } else {
            perror("read");
        }
        // 释放该客户端拥有的会话
        if (client->attached_session != -1) {
            Session *sess = get_session(client->attached_session);
            if (sess && sess->owner_fd == client->sockfd) {
                sess->owner_fd = -1;
            }
        }
        close(client->sockfd);
        client->sockfd = -1;
        client->attached_session = -1;
        return;
    }

    // 追加到缓冲区
    if (client->inlen + n >= BUFFER_SIZE) {
        // 缓冲区满，关闭连接
        close(client->sockfd);
        client->sockfd = -1;
        return;
    }
    memcpy(client->inbuf + client->inlen, buf, n);
    client->inlen += n;

    // 按行处理
    char *line_start = client->inbuf;
    char *p;
    while ((p = memchr(line_start, '\n', client->inlen - (line_start - client->inbuf))) != NULL) {
        *p = '\0';
        // 处理一行
        if (client->attached_session == -1) {
            handle_command_mode(client, line_start);
        } else {
            handle_attached_mode(client, line_start);
        }
        line_start = p + 1;
        if (client->sockfd == -1) break; // 客户端可能已关闭
    }
    // 移动剩余数据
    if (line_start > client->inbuf) {
        size_t remaining = client->inlen - (line_start - client->inbuf);
        if (remaining > 0) {
            memmove(client->inbuf, line_start, remaining);
        }
        client->inlen = remaining;
    }
}

// 处理子进程输出（从stdout_fd读取，转发给所有者）
static void handle_session_output(Session *sess) {
    char buf[4096];
    ssize_t n = read(sess->stdout_fd, buf, sizeof(buf)-1);
    if (n <= 0) {
        if (n == 0) {
            // 子进程关闭了stdout，可能已退出
            // 关闭会话
            close_session(sess->id);
        } else if (errno != EAGAIN) {
            perror("read from child");
            close_session(sess->id);
        }
        return;
    }
    buf[n] = '\0';
    // 转发给所有者
    if (sess->owner_fd != -1) {
        send(sess->owner_fd, buf, n, 0);
    } else {
        // 没有所有者，丢弃输出
    }
}

// 创建监听socket
static int create_listen_socket(int port) {
    int fd = socket(AF_INET, SOCK_STREAM, 0);
    if (fd < 0) return -1;
    int opt = 1;
    setsockopt(fd, SOL_SOCKET, SO_REUSEADDR, &opt, sizeof(opt));
    struct sockaddr_in addr;
    memset(&addr, 0, sizeof(addr));
    addr.sin_family = AF_INET;
    addr.sin_addr.s_addr = INADDR_ANY;
    addr.sin_port = htons(port);
    if (bind(fd, (struct sockaddr*)&addr, sizeof(addr)) < 0) {
        close(fd);
        return -1;
    }
    if (listen(fd, 10) < 0) {
        close(fd);
        return -1;
    }
    return fd;
}

// 回收僵尸进程
static void reap_children() {
    while (1) {
        int status;
        pid_t pid = waitpid(-1, &status, WNOHANG);
        if (pid <= 0) break;
        // 找到对应会话并清理
        for (int i = 1; i < MAX_SESSIONS; i++) {
            if (sessions[i].valid && sessions[i].pid == pid) {
                close_session(i); // 关闭管道，标记无效
                break;
            }
        }
    }
}

int main(int argc, char **argv) {
    argument_parser_t parser;
    argparse_init(&parser, argc, argv, "speed tcp server", "Session Persistent Execution Environment Daemon for bash scripting\n");
    int verbosity;
    int listen_port = LISTEN_PORT;

    argparse_arg_t arg1 = ARGPARSE_COUNT(
        'v', "--verbose", &verbosity, "verbosity level"
    );
    argparse_arg_t arg2 = ARGPARSE_OPTION(
        INT, 'p', "--port", &listen_port, "tcp port to listen (>1999 or 9999)"
    );
    argparse_add_argument(&parser, &arg1);
    argparse_add_argument(&parser, &arg2);
/*
    argparse_arg_t args[] = {
        ARGPARSE_COUNT('v', "--verbose", &verbosity, "verbosity level"),
        ARGPARSE_OPTION(INT, 'p', "--port", &listen_port, "tcp port to listen (>1999 or 9999)"),
    };
    argparse_add_arguments(&parser, args, 2);
*/
    argparse_parse_args(&parser);
    if (listen_port < 2000) {
	    listen_port = LISTEN_PORT;
    }

    init();

    listen_fd = create_listen_socket(listen_port);
    if (listen_fd < 0) {
        perror("listen socket");
        exit(1);
    }
    printf("speed server listening on port %d\n", listen_port);

    fd_set read_fds;
    int max_fd;

    while (1) {
        reap_children(); // 回收子进程

        FD_ZERO(&read_fds);
        FD_SET(listen_fd, &read_fds);
        max_fd = listen_fd;

        // 添加客户端
        for (int i = 0; i < MAX_CLIENTS; i++) {
            if (clients[i].sockfd != -1) {
                FD_SET(clients[i].sockfd, &read_fds);
                if (clients[i].sockfd > max_fd) max_fd = clients[i].sockfd;
            }
        }
        // 添加子进程stdout
        for (int i = 1; i < MAX_SESSIONS; i++) {
            if (sessions[i].valid) {
                FD_SET(sessions[i].stdout_fd, &read_fds);
                if (sessions[i].stdout_fd > max_fd) max_fd = sessions[i].stdout_fd;
            }
        }

        if (select(max_fd + 1, &read_fds, NULL, NULL, NULL) < 0) {
            if (errno == EINTR) continue;
            perror("select");
            break;
        }

        // 新连接
        if (FD_ISSET(listen_fd, &read_fds)) {
            struct sockaddr_in addr;
            socklen_t len = sizeof(addr);
            int client_fd = accept(listen_fd, (struct sockaddr*)&addr, &len);
            if (client_fd >= 0) {
                int i;
                for (i = 0; i < MAX_CLIENTS; i++) {
                    if (clients[i].sockfd == -1) {
                        clients[i].sockfd = client_fd;
                        clients[i].inlen = 0;
                        clients[i].attached_session = -1;
                        printf("New client from %s:%d, fd=%d\n",
                               inet_ntoa(addr.sin_addr), ntohs(addr.sin_port), client_fd);
                        break;
                    }
                }
                if (i == MAX_CLIENTS) {
                    close(client_fd);
                }
            }
        }

        // 处理客户端数据
        for (int i = 0; i < MAX_CLIENTS; i++) {
            if (clients[i].sockfd != -1 && FD_ISSET(clients[i].sockfd, &read_fds)) {
                handle_client_read(&clients[i]);
            }
        }

        // 处理子进程输出
        for (int i = 1; i < MAX_SESSIONS; i++) {
            if (sessions[i].valid && FD_ISSET(sessions[i].stdout_fd, &read_fds)) {
                handle_session_output(&sessions[i]);
            }
        }
    }

    close(listen_fd);
    return 0;
}
