#include "mudfish_dns_module.h"

#include <arpa/inet.h>
#include <arpa/nameser.h>
#include <stdlib.h>
#include <string.h>
#include <strings.h>

struct policy {
	uint8_t response[65535];
};

static int
question(const struct mudfish_dns_event *event, ns_rr *rr)
{
	ns_msg message;

	if (event->query.len > 65535 ||
	    ns_initparse(event->query.data, (int)event->query.len, &message) < 0)
		return (-1);
	return (ns_parserr(&message, ns_s_qd, 0, rr));
}

static int32_t
dns_query(void *context, const struct mudfish_dns_event *event,
    struct mudfish_dns_result *result)
{
	ns_rr rr;

	(void)context;
	if (question(event, &rr) < 0)
		return (-1);
	if (strcasecmp(ns_rr_name(rr), "allow.test") == 0)
		result->action = MUDFISH_DNS_ALLOW;
	else if (strcasecmp(ns_rr_name(rr), "block.test") == 0)
		result->action = MUDFISH_DNS_BLOCK;
	else if (strcasecmp(ns_rr_name(rr), "drop.test") == 0)
		result->action = MUDFISH_DNS_DROP;
	else if (strcasecmp(ns_rr_name(rr), "redirect.test") == 0) {
		result->action = MUDFISH_DNS_ANSWER;
		result->ttl = 60;
		if (ns_rr_type(rr) == ns_t_aaaa) {
			result->address_family = MUDFISH_DNS_IPV6;
			inet_pton(AF_INET6, "2001:db8::42", result->address);
		} else {
			result->address_family = MUDFISH_DNS_IPV4;
			inet_pton(AF_INET, "192.0.2.42", result->address);
		}
	}
	return (0);
}

static int32_t
dns_response(void *context, const struct mudfish_dns_event *event,
    struct mudfish_dns_result *result)
{
	struct policy *policy = context;
	ns_msg message;
	ns_rr rr;
	uint8_t *rdata;
	int section, i, changed = 0;

	if (question(event, &rr) < 0)
		return (-1);
	if (strcasecmp(ns_rr_name(rr), "rewrite.test") != 0)
		return (0);
	if (event->response.len > sizeof(policy->response))
		return (-1);
	memcpy(policy->response, event->response.data, event->response.len);
	if (ns_initparse(policy->response, (int)event->response.len, &message) < 0)
		return (-1);
	/* This example only edits unsigned answers. A policy that edits signed
	 * data must also rebuild/remove the affected signatures.
	 */
	for (section = ns_s_an; section <= ns_s_ar; section++) {
		for (i = 0; i < ns_msg_count(message, section); i++) {
			if (ns_parserr(&message, (ns_sect)section, i, &rr) < 0)
				return (-1);
			if (ns_rr_type(rr) == ns_t_rrsig || ns_rr_type(rr) == ns_t_sig ||
			    ns_rr_type(rr) == ns_t_tsig)
				return (0);
		}
	}
	for (i = 0; i < ns_msg_count(message, ns_s_an); i++) {
		if (ns_parserr(&message, ns_s_an, i, &rr) < 0)
			return (-1);
		if (ns_rr_class(rr) != ns_c_in || ns_rr_type(rr) != ns_t_a ||
		    ns_rr_rdlen(rr) != 4)
			continue;
		/* rdata belongs to our copy. TTL precedes RDLENGTH (2 bytes). */
		rdata = policy->response + (ns_rr_rdata(rr) - policy->response);
		inet_pton(AF_INET, "203.0.113.7", rdata);
		ns_put32(60, rdata - 6);
		changed = 1;
	}
	if (changed) {
		result->action = MUDFISH_DNS_REPLACE;
		result->response.data = policy->response;
		result->response.len = event->response.len;
	}
	return (0);
}

static void
destroy(void *context)
{

	free(context);
}

int32_t
mudfish_dns_module_init_v2(uint32_t abi_version, size_t module_size,
    struct mudfish_dns_module_v2 *module)
{

	if (abi_version != MUDFISH_DNS_MODULE_ABI_V2 || module_size != sizeof(*module))
		return (-1);
	module->context = calloc(1, sizeof(struct policy));
	if (module->context == NULL)
		return (-1);
	module->on_dns_query = dns_query;
	module->on_dns_response = dns_response;
	module->destroy = destroy;
	return (0);
}
