/*
  UDP redirect. This program redirect UDP connection from local address to send-to-ip address

  Copyright (C) 2013 Vladimir Oleynik <dzo@simtreas.ru>
  (Thanks by Ivan Tikhonov <kefeer@brokestream.com> for the original idea)

 Licensed under GPLv2 or later.


Usage:
udp_redir [-b incoming_addr[:incoming_port]] [-B outgoing_addr[:outgoing_port] [-B...]] [-F] [-p pidfile] send-to-ip:port

if incoming_port unspecified then used the same 'port', default pidfile is DEFAULT_PIDS_DIR/program_name.pid
if outgoing_port unspecified then used dynamic ports 1024...65535, -B options with port may be more one,
maximum clients is MAX_CLIENTS, -F - set foreground and verbose debug

  Example:
  - redirect MAX_CLIENTS programs connected to $proxy_address:$port as make connect
    to $external_address:$port from specified $local_address
udp_redir -b $proxy_address -B $local_address $external_address:$port
  - redirect DNS from localhost address:
udp_redir -b localhost 8.8.8.8:53

  - redirect one or two a program connected to $proxy_address:$proxy_port make connect to
    $external_address:$external_port from specified $local_addres1:$local_port1 or $local_addres2:$local_port2
udp_redir -b $proxy_address:$proxy_port -B $local_addres1:$local_port1 -B $local_addres2:$local_port2 $external_address:$external_port
udp_redir -b 192.168.1.1:53 -B 1.1.1.1:53 -B 1.1.1.2:53 8.8.8.8:53

*/


/* may be change for you */
#define MAX_CLIENTS 20
#define DEFAULT_PIDS_DIR "/var/run"

#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <getopt.h>
#include <errno.h>
#include <stdarg.h>
#include <unistd.h>
#include <sys/socket.h>
#include <netinet/in.h>
#include <netdb.h>
#include <arpa/inet.h>
#include <time.h>
#include <signal.h>

#if !defined(__GNUC__)
# ifndef __attribute__
#  define __attribute__(x)
# endif
#endif

static const char *program_name;

static void usage(int exit_code) __attribute__ ((__noreturn__));
static void usage(int exit_code)
{
	printf("UDP proxy. This program proxy UDP connection from local address to send-to-ip address\n");
	printf("Usage: %s [-b incoming_addr[:incoming_port]] [-B outgoing_addr[:outgoing_port] [-B...]] [-F] [-p pidfile] send-to-ip:port\n",
		program_name);
	printf("if incoming_port unspecified then used the same 'port', default pidfile is %s/%s.pid\n", DEFAULT_PIDS_DIR, program_name);
	printf("if outgoing_port unspecified then used a dynamic port 1024...65535, -B options with port may be more one\n");
	printf("maximum clients is %d, -F - set foreground and verbose debug\n", MAX_CLIENTS);
	exit(exit_code);
}

static void die(const char *s, ...) __attribute__ ((noreturn, format (printf, 1, 2)));
static void die(const char *s, ...)
{
	va_list p;
	char *f;

	va_start(p, s);
					/* ": \n\0" */
	f = alloca(strlen(program_name) + 4 + strlen(s));
	sprintf(f, "%s: %s\n", program_name, s);
	fflush(stdout);
	vfprintf(stderr, f, p);
	va_end(p);
	exit(1);
}

static int parse_host_colon_port(struct sockaddr_in *psa, char *name, int need_port)
{
	psa->sin_family = AF_INET;
	if (name != NULL && *name != '\0') {
		struct hostent *hp;
		char *colon = strchr(name, ':');

		if(colon == NULL && need_port)
			return -1;
		if(colon != NULL)
			*colon++ = '\0';
		if ((hp = gethostbyname(name)) == NULL)
			die("'%s' host unknown.", name);
		memcpy(&(psa->sin_addr.s_addr), hp->h_addr, hp->h_length);
		if(colon != NULL) {
			char *endptr;
			long port;

			if(*colon == '\0')
				die("need specified port after colon");
			errno = 0;
			port = strtol(colon, &endptr, 0);
			if (*endptr != '\0')
				die("can not parse port '%s'.", colon);
			if(errno || port < need_port || port > 65535)
				die("'%s' strange port.", colon);
			psa->sin_port = htons(port);
			*--colon = ':'; /* restore */
			return port != 0;
		} else {
			psa->sin_port = 0;
			return 0;
		}
	} else {
		if(need_port)
			return -1;
		psa->sin_addr.s_addr = 0;
		psa->sin_port = 0;
		return 0;
	}
}

static int create_socket(void)
{
	int sock = socket(AF_INET, SOCK_DGRAM, IPPROTO_UDP);
	if(sock < 0)
		die("can not create socket");
	return sock;
}


static char *show_ip_port(const struct sockaddr_in *p_sa)
{
	char ipstr[16];
	static char ip_port_buf[16+1+6];
	int p;

	inet_ntop(AF_INET, &(p_sa->sin_addr), ipstr, sizeof ipstr);
	p = ntohs(p_sa->sin_port);
	sprintf(ip_port_buf, "%s:%d", ipstr, p);
	return ip_port_buf;
}

static char *xstrdup (const char *s)
{
	char *t;

	if (s == NULL)
		return NULL;
	t = strdup (s);
	if (t == NULL)
		die("Memory exhausted while save '%s'", s);
	return t;
}

static void add_out_a(char ***p_out_a, int *pn, const char *a)
{
	char **dim = *p_out_a;
	int n = *pn;
	size_t size;

	if(n == MAX_CLIENTS)
		die("Too many outgoing address specified, recompile with bigger MAX_CLIENTS");
	size = (n + 1)  * sizeof(char *);
	if(dim == NULL)
		dim = malloc(size);
	else
		dim = realloc(dim, size);
	if (dim == NULL)
		die("Memory exhausted while save outgoing address '%s'", a);
	dim[n] = xstrdup(a);
	*p_out_a = dim;
	*pn = n + 1;
}


static void set_program_name(const char *a0)
{
	const char *s;

	program_name = a0;
	for (s = program_name; *s; )
		if (*(s++) == '/') program_name = s;
}

static char *pidfile;

typedef void (*sh_t)(int s);

static void xsignal(unsigned int sigs_bits, sh_t h)
{
	unsigned int b = 2;
	int signo = 1;
	struct sigaction act, old;

	memset(&act, 0, sizeof(act));
	act.sa_handler = h;
	act.sa_flags |= SA_RESTART;
	while(b) {
		if((sigs_bits & b) != 0) {
			sigaddset(&act.sa_mask, signo);
			sigaction(signo, &act, &old);
		}
		signo++;
		b <<= 1;
	}
}

static void sighandler(int signal)
{
#if defined(__GNUC__)
	(void)&signal;
#endif
	if(pidfile != NULL)
		unlink(pidfile);
	_exit(1);
}

static void xdaemon(void)
{
	FILE *fpid;

	if(pidfile == NULL) {
		pidfile = malloc(strlen(DEFAULT_PIDS_DIR) + strlen(program_name) + 6);  /* +"/.pid\0" */
		if(pidfile == NULL)
			die("Memory exhausted while make pidfile '%s/%s.pid'", DEFAULT_PIDS_DIR, program_name);
		sprintf(pidfile, "%s/%s.pid", DEFAULT_PIDS_DIR, program_name);
	}
	if ((fpid = fopen(pidfile , "w")) == NULL)
		die("can not create pidfile '%s': %s", pidfile, strerror(errno));
	/* remove pidfile if signaled */
	xsignal((1<<SIGINT)|(1<<SIGTERM)|(1<<SIGKILL)|(1<<SIGQUIT), sighandler);
	if(daemon(0, 0) < 0)
	  die("unable to stay daemon: %s", strerror(errno));
	fprintf(fpid, "%u\n", getpid());
	fclose(fpid);
}

static char buf[65535];

int main(int argc, char *argv[])
{
	int p, n, f, nbuf;
	char *bind_incoming_addr = NULL;
	char **bind_outgoing_addrs = NULL;
	int n_outgoing_addrs = 0;
	int foreground = 0;
	fd_set socks;
	int sock_in, dim_sock_out[MAX_CLIENTS];
	struct sockaddr_in clients[MAX_CLIENTS], dest, sa, new_out;
	int current_clients = 0;
	int max_clients;
	time_t last[MAX_CLIENTS];

	set_program_name(argv[0]);
	while ((p = getopt(argc, argv, "Fb:B:p:h")) != -1) {
		switch (p) {
		case 'F':
			foreground++;
			break;
		case 'b':
			if(bind_incoming_addr != NULL)
				die("Multiple -b option unsupport");
			bind_incoming_addr = xstrdup(optarg);
			break;
		case 'B':
			add_out_a(&bind_outgoing_addrs, &n_outgoing_addrs, optarg);
			break;
		case 'p':
			pidfile = optarg;
			break;
		case 'h':
		case '?':
			usage(0);
		default:
			usage(1);
		}
	}
	argv += optind;
	argc -= optind;

	if (argc != 1)
		usage(1);

	if(parse_host_colon_port(&dest, argv[0], 1) < 0)
		die("Need specified send-to-ip:port");

	sock_in = create_socket();
	if(parse_host_colon_port(&sa, bind_incoming_addr, 0) == 0)
		sa.sin_port = dest.sin_port;    /* set proxy_port=ext_port */
	if(bind(sock_in, (struct sockaddr *)&sa, sizeof(sa)) == -1)
		die("can not bind to incoming address '%s': %s",
			bind_incoming_addr == NULL ? "0.0.0.0:0" : bind_incoming_addr, strerror(errno));

	f = n_outgoing_addrs;
	if(f == 0)      /* outgoing unspecified, set automate find out interface with dynamic ports */
		add_out_a(&bind_outgoing_addrs, &f, "0.0.0.0:0");
	for(n = 0; n < f; n++) {
		p = parse_host_colon_port(&new_out, bind_outgoing_addrs[n], 0);
		if(n_outgoing_addrs > 1 && p == 0)
			die("Multiple outgoing address without specified ports unsupport");
		if(bind((dim_sock_out[n] = create_socket()), (struct sockaddr *)&new_out, sizeof(new_out)) == -1)
			die("can not bind to outgoing address '%s': %s", bind_outgoing_addrs[n], strerror(errno));
	}
	max_clients = new_out.sin_port == 0 ? MAX_CLIENTS : n_outgoing_addrs;
	/* already binded */
	n_outgoing_addrs = f;

	if(!foreground)
		xdaemon();

	while(1) {
		socklen_t sz_sa = sizeof(sa);

		FD_ZERO(&socks);
		f = sock_in;
		FD_SET(f, &socks);
		for(n = 0; n < current_clients; n++) {
			FD_SET(dim_sock_out[n], &socks);
			/* find a max descriptor */
			if(dim_sock_out[n] > f)
				f = dim_sock_out[n];
		}

		f = select(f + 1, &socks, NULL, NULL, NULL);
		if(f > 0) {
			if(FD_ISSET(sock_in, &socks)) {
				f = sock_in;
			} else {
loop:
				for(n = 0; n < current_clients; n++) {
					f = dim_sock_out[n];
					if(FD_ISSET(f, &socks))
						break;
				}
				if(n == current_clients)
					continue;       /* processed */
			}
			FD_CLR(f, &socks);

			nbuf = recvfrom(f, buf, sizeof(buf), 0, (struct sockaddr *)&sa, &sz_sa);
			if(foreground) {
				const char *s = show_ip_port(&sa);
				if(f == sock_in)
					printf("read from %s (local) %d bytes\n", s, nbuf);
				else
					printf("read from %s (remote #%d) %d bytes\n", s, n, nbuf);
			}
			if(nbuf <= 0)
				goto loop;

			if(f == sock_in) {
				/* read from local address, find number clients */
				for(n = 0; n < current_clients; n++) {
					if(memcmp(clients + n, &sa, sizeof(sa)) == 0) {
						/* found client */
						f = dim_sock_out[n];
						break;
					}
				}
				if(n == current_clients) {
					/* new client */
					if(current_clients < max_clients) {
						/* have "free" sockets */
						if(n < n_outgoing_addrs) {
							/* first bind sockets already */
							f = dim_sock_out[n];
						} else {
							bind((f = create_socket()), (struct sockaddr *)&new_out, sizeof(new_out));
							dim_sock_out[n] = f;
						}
						current_clients++;
					} else {
						/* forgot a old inactive client */
						n = 0;
						time_t old = last[0];
						for(f = 1; f < max_clients; f++) {
							if(last[f] < old) {
								old = last[f];
								n = f;
							}
						}
						if(foreground)
							printf("detect number client over max_clients=%d, drop old inactive #%d client\n",
								max_clients, n);
						f = dim_sock_out[n];
					}
					clients[n] = sa;
				}
				last[n] = max_clients == 1 ? 1 : time(NULL);
				sa = dest;
			} else {
				/* read from remote address, set a client's sock */
				f = sock_in;
				sa = clients[n];
			}
			nbuf = sendto(f, buf, nbuf, 0, (struct sockaddr *)&sa, sizeof(sa));
			if(foreground)
				printf("send to %s (%s #%d) %d bytes\n", show_ip_port(&sa),
					(f == sock_in ? "local" : "remote"), n, nbuf);
			goto loop;
		} else if(foreground)
			printf("select returned %d\n", f);
	}
	return 0;
}
