#ifndef CORO3_H
#define CORO3_H

#include <stdbool.h>

void        coro_yield();
void        coro_spawn(void (*fn)(void*), void*);
bool        coro_poll();
void const* coro_gen(void (*fn)(void*), void*);
void*       coro_next(void const* generator);
void        coro_return(void*);

#endif
#ifdef	CORO3_IMPLEMENTATION

#include <stddef.h>
#include <stdlib.h>
#include <sys/mman.h>

#define auto_t __auto_type
#define typeof __Typeof
#define apush(vec, val) ({\
	auto_t x = &(val);\
	auto_t v = &(vec);\
	if (v->count == v->capacity) {\
		v->capacity = v->capacity? 2 * v->capacity: 1;\
		v->items = realloc(v->items, sizeof(*x) * v->capacity);\
		if (v->items == NULL)\
			abort();\
	}\
	v->items[v->count++] = *x;\
})

#ifndef CORO_STACK_SIZE
#define CORO_STACK_SIZE (4 << 20 /* 4MB */)
#endif

struct coroutine {
	void *rsp, *base;
	bool dying, paused;
	size_t id;
	void *oob;
};

struct _global_state {
	struct coroutine* items;
	size_t capacity, count, dying, current;
};

static struct _global_state _ = {0};
static struct coroutine* current(void);
static void   restore(void*);
static void   reap(void);
static void   die(void);
static void   save_rsp(void* rsp);
static void   jmp2next(void);

__attribute__((constructor)) static void coro_init() {
	apush(_, (struct coroutine){.id = 1});
}

__attribute__((destructor)) static void coro_fini() {
	free(_.items);
	_ = (struct _global_state) {0};
}

__attribute__((naked)) void coro_yield() {
	__asm__("push %rdi");
	__asm__("push %rbp");
	__asm__("push %rbx");
	__asm__("push %r12");
	__asm__("push %r13");
	__asm__("push %r14");
	__asm__("push %r15");
	__asm__("mov  %rsp, %rdi");
	__asm__("jmp *%0"::"r" (save_rsp));
}

void coro_spawn(void (*fn)(void*), void* arg) {
	static size_t id = 1; // 1 is for main.
	static int prot = PROT_READ | PROT_WRITE;
	static int flag = MAP_PRIVATE | MAP_ANONYMOUS | MAP_GROWSDOWN;
	void* base = mmap(NULL, CORO_STACK_SIZE, prot, flag, -1, 0);
	void** rsp = base;
	*(--rsp) = die;
	*(--rsp) = fn;
	*(--rsp) = arg; // %rdi
	*(--rsp) = 0;   // %rbp");
	*(--rsp) = 0;   // %rbx");
	*(--rsp) = 0;   // %r12");
	*(--rsp) = 0;   // %r13");
	*(--rsp) = 0;   // %r14");
	*(--rsp) = 0;   // %r15");	
	apush(_, ((struct coroutine){
		.rsp = rsp, .base = base,
		.id  = ++id
	}));
}

void const* coro_gen(void (*fn)(void*), void* arg) {
	coro_spawn(fn, arg);
	_.items[_.count - 1].paused = true;
	return (void const*) _.items[_.count - 1].id;
}

void* coro_next(void const* generator) {
	bool found = true;
	while (found) {
		found = false;
		// search for generator, unpause it and yield
		for (size_t i = 0; i < _.count; i++) {
			if (_.items[i].id != (size_t) generator)
				continue;
			found = true;
			_.items[i].oob = NULL;
			_.items[i].paused = false;
			coro_yield();
			break; 
		}
		// search for generator, if it's paused, it's finished
		// and we can return it's out of band pointer, otherwise
		// it's still doing work and we have to yield again
		for (size_t i = 0; found && i < _.count; i++) {
			if (_.items[i].id != (size_t) generator)
				continue;
			if (_.items[i].paused) // finished
				return _.items[i].oob;
			coro_yield(); 
		}
		// if we never found the generator, it's already dead,
		// and we have nothing to wait for, and can return a 
		// NULL value.
	}
	return NULL;
}

void coro_return(void* ptr) {
	current()->oob = ptr;
	current()->paused = true;
	coro_yield();
}

bool coro_poll() {
	return _.count > 1;
}

void die() {
	if (_.count == 0)
		abort();
	current()->dying = true;
	struct coroutine tmp = *current();
	*current() = _.items[_.count - 1];
	_.items[_.count - 1] = tmp;
	
	_.count -= 1;
	if (_.count) {
		_.dying += 1;
		_.current += 1;
		_.current %= _.count;
		restore(current()->rsp);
	} else {
		// munmap'ing our own stack sounds like an awful idea
		// just exit.
		exit(EXIT_SUCCESS);	
	}
}

__attribute__((naked)) void restore(void* rsp) {
	__asm__ __volatile__(""::"m" (rsp));
	__asm__             ("mov %rdi, %rsp"); (void) rsp; // probably illegal :)
	// we can now reap the dying stacks, as our stack isn't dying
	reap();
	__asm__("pop %r15");
	__asm__("pop %r14");
	__asm__("pop %r13");
	__asm__("pop %r12");
	__asm__("pop %rbx");
	__asm__("pop %rbp");
	__asm__("pop %rdi");
	__asm__("ret");
}

void reap() {
	for (size_t i = _.count; i < _.count + _.dying; i++) {
		munmap(_.items[i].base, CORO_STACK_SIZE);
		_.items[i] = (struct coroutine) {0};
	}
	_.dying = 0;
}

struct coroutine* current() {
	if (_.count == 0)
		abort();
	return &_.items[_.current];
}

void save_rsp(void* rsp) {
	if (_.count == 0)
		abort();
	current()->rsp = rsp;
	jmp2next();
}

void jmp2next(void) {
	if (_.count == 0)
		abort();
	size_t attempts = 0;
	do {
		if (attempts >= _.count)
			abort();
		attempts  += 1;
		_.current += 1;
		_.current %= _.count;
	} while(current()->paused);
	restore(current()->rsp);	
}

#undef	CORO3_IMPLEMENTATION
#endif
