#include <patriciatree.h>
#include <string.h>

static struct patriciaTreeNode_t *patriciaTreeInsertHelper(struct patriciaTreeNode_t *root, char *str, uint32t val);
static struct patriciaTreeNode_t *patriciaTreeGetNode(struct patriciaTreeNode_t *root, char *str);

static struct patriciaTreeNode_t *root;
static struct patriciaTreeNode_t *newNode;

void patriciaTreeInit(void *mem) {
	newNode= mem;
}

struct patriciaTreeNode_t *patriciaTreeDeinit() {
	return root;
}

struct patriciaTreeNode_t *patriciaTreeGetEndMem() {
	return newNode;
}

void patriciaTreeInsert(char *str, uint32t val) {
	//no root
	if(unlikely(root == 0)) {
		root= newNode++;
		
		root->left= root;
		root->right= root;
		root->str= str;
		root->bit= 0;
		root->value= val;
	} else {
		patriciaTreeInsertHelper(root,str,val);
	}
}

struct patriciaTreeNode_t *patriciaTreeInsertHelper(struct patriciaTreeNode_t *root, char *str, uint32t val) {
	struct patriciaTreeNode_t *new, *node= root, *parent= root;
	uint32t bit= 0;
	
	do {
		if(unlikely(bit == node->bit)) {
			parent= node;
			if(getBit(str,bit) == 0)
				node= node->left;
			else
				node= node->right;
		}
		bit++;
	} while(likely(getBit(str,bit) == getBit(node->str,bit) || node->bit == bit));
	
	new= newNode++;
	
	new->str= str;
	new->bit= bit;
	new->value= val;
	
	if(getBit(node->str,bit) == 0) {
		new->left= node;
		new->right= new;
	} else {
		new->left= new;
		new->right= node;
	}
	
	if(getBit(str,parent->bit) == 0)
		parent->left= new;
	else
		parent->right= new;
	
	return new;
}

uint32t patriciaTreeGetValue(char *str) {
	struct patriciaTreeNode_t *node= patriciaTreeGetNode(root,str);
	
	if(likely(strcmp(node->str,str) == 0))
		return node->value;
	else
		return 0;
}

struct patriciaTreeNode_t *patriciaTreeGetNode(struct patriciaTreeNode_t *root, char *str) {
	struct patriciaTreeNode_t *node= root, *parent;
	uint32t maxBit= ((strlen(str) + 1) << 3) - 1;
	
	do {
		parent= node;
		if(getBit(str,node->bit) == 0) {
			node= node->left;
		} else {
			node= node->right;
		}
	} while(likely(parent->bit < node->bit && node->bit <= maxBit));
	
	return node;
}