summaryrefslogtreecommitdiff
path: root/src/bokfd.c
diff options
context:
space:
mode:
Diffstat (limited to 'src/bokfd.c')
-rw-r--r--src/bokfd.c267
1 files changed, 234 insertions, 33 deletions
diff --git a/src/bokfd.c b/src/bokfd.c
index 70b654a..944e910 100644
--- a/src/bokfd.c
+++ b/src/bokfd.c
@@ -1,6 +1,8 @@
#include <errno.h>
#include <fcntl.h>
#include <netdb.h>
+#include <openssl/err.h>
+#include <openssl/ssl.h>
#include <poll.h>
#include <signal.h>
#include <stdio.h>
@@ -10,6 +12,7 @@
#include <sys/stat.h>
#include <sys/un.h>
#include <termios.h>
+#include <time.h>
#include <unistd.h>
#include "auth.h"
@@ -26,6 +29,8 @@
struct conn {
int fd;
+ SSL *ssl;
+ int handshaking;
struct buf in;
struct buf out;
size_t out_sent;
@@ -34,6 +39,10 @@ struct conn {
static volatile sig_atomic_t g_stop = 0;
static size_t g_line_limit = 1024 * 1024;
+static SSL_CTX *g_tls_ctx = NULL;
+static time_t g_tls_cert_mtime = 0;
+static time_t g_tls_key_mtime = 0;
+static time_t g_tls_next_try = 0;
static void on_signal(int sig)
{
@@ -208,7 +217,92 @@ static int tcp_listen(const char *addrport)
return fd;
}
-static void accept_conns(int lfd, struct conn *conns, size_t *nconns)
+static void log_tls_error(const char *what)
+{
+ unsigned long e = ERR_get_error();
+ char buf[256] = "unknown error";
+ if (e)
+ ERR_error_string_n(e, buf, sizeof buf);
+ log_error("%s: %s", what, buf);
+ ERR_clear_error();
+}
+
+static int tls_load_certs(void)
+{
+ SSL_CTX *ctx = SSL_CTX_new(TLS_server_method());
+ if (!ctx) {
+ log_tls_error("TLS context");
+ return -1;
+ }
+ SSL_CTX_set_min_proto_version(ctx, TLS1_2_VERSION);
+ SSL_CTX_set_options(ctx, SSL_OP_NO_COMPRESSION |
+ SSL_OP_CIPHER_SERVER_PREFERENCE |
+ SSL_OP_NO_RENEGOTIATION);
+ if (SSL_CTX_use_certificate_chain_file(ctx, g_cfg.tls_cert) != 1) {
+ log_tls_error(g_cfg.tls_cert);
+ SSL_CTX_free(ctx);
+ return -1;
+ }
+ if (SSL_CTX_use_PrivateKey_file(ctx, g_cfg.tls_key, SSL_FILETYPE_PEM) !=
+ 1) {
+ log_tls_error(g_cfg.tls_key);
+ SSL_CTX_free(ctx);
+ return -1;
+ }
+ if (SSL_CTX_check_private_key(ctx) != 1) {
+ log_error("TLS key does not match certificate %s", g_cfg.tls_cert);
+ SSL_CTX_free(ctx);
+ return -1;
+ }
+ SSL_CTX *old = g_tls_ctx;
+ g_tls_ctx = ctx;
+ if (old)
+ SSL_CTX_free(old);
+ struct stat st;
+ if (stat(g_cfg.tls_cert, &st) == 0)
+ g_tls_cert_mtime = st.st_mtime;
+ if (stat(g_cfg.tls_key, &st) == 0)
+ g_tls_key_mtime = st.st_mtime;
+ return 0;
+}
+
+static void tls_reload_check(void)
+{
+ struct stat st;
+ time_t cert = 0, key = 0;
+ if (stat(g_cfg.tls_cert, &st) == 0)
+ cert = st.st_mtime;
+ if (stat(g_cfg.tls_key, &st) == 0)
+ key = st.st_mtime;
+ if (cert == 0 || key == 0 ||
+ (cert == g_tls_cert_mtime && key == g_tls_key_mtime))
+ return;
+ time_t now = time(NULL);
+ if (now < g_tls_next_try)
+ return;
+ g_tls_next_try = now + 60;
+ if (tls_load_certs() == 0)
+ log_info("TLS certificate reloaded");
+ else
+ log_error("TLS reload failed; keeping previous certificate");
+}
+
+static void conn_handshake(struct conn *c)
+{
+ int r = SSL_accept(c->ssl);
+ if (r == 1) {
+ c->handshaking = 0;
+ return;
+ }
+ int e = SSL_get_error(c->ssl, r);
+ if (e == SSL_ERROR_WANT_READ || e == SSL_ERROR_WANT_WRITE)
+ return;
+ log_warn("TLS handshake failed from fd %d", c->fd);
+ ERR_clear_error();
+ c->closing = 1;
+}
+
+static void accept_conns(int lfd, struct conn *conns, size_t *nconns, int tls)
{
for (;;) {
int fd = accept4(lfd, NULL, NULL, SOCK_NONBLOCK | SOCK_CLOEXEC);
@@ -223,22 +317,54 @@ static void accept_conns(int lfd, struct conn *conns, size_t *nconns)
c->fd = fd;
buf_init(&c->in);
buf_init(&c->out);
+ if (tls) {
+ c->ssl = SSL_new(g_tls_ctx);
+ if (!c->ssl) {
+ log_tls_error("TLS connection");
+ close(fd);
+ buf_free(&c->in);
+ buf_free(&c->out);
+ (*nconns)--;
+ continue;
+ }
+ SSL_set_fd(c->ssl, fd);
+ SSL_set_accept_state(c->ssl);
+ c->handshaking = 1;
+ }
}
}
static void conn_read(struct conn *c, sqlite3 *db)
{
- char tmp[READ_CHUNK];
- ssize_t r = read(c->fd, tmp, sizeof tmp);
- if (r == 0) {
- c->closing = 1;
- return;
+ if (c->ssl && c->handshaking) {
+ conn_handshake(c);
+ if (c->closing || c->handshaking)
+ return;
}
- if (r < 0) {
- if (errno == EAGAIN || errno == EWOULDBLOCK || errno == EINTR)
+ char tmp[READ_CHUNK];
+ ssize_t r;
+ if (c->ssl) {
+ r = SSL_read(c->ssl, tmp, sizeof tmp);
+ if (r <= 0) {
+ int e = SSL_get_error(c->ssl, (int)r);
+ if (e == SSL_ERROR_WANT_READ || e == SSL_ERROR_WANT_WRITE)
+ return;
+ ERR_clear_error();
+ c->closing = 1;
return;
- c->closing = 1;
- return;
+ }
+ } else {
+ r = read(c->fd, tmp, sizeof tmp);
+ if (r == 0) {
+ c->closing = 1;
+ return;
+ }
+ if (r < 0) {
+ if (errno == EAGAIN || errno == EWOULDBLOCK || errno == EINTR)
+ return;
+ c->closing = 1;
+ return;
+ }
}
buf_append(&c->in, tmp, (size_t)r);
for (;;) {
@@ -262,18 +388,39 @@ static void conn_read(struct conn *c, sqlite3 *db)
static void conn_flush(struct conn *c)
{
+ if (c->ssl && c->handshaking) {
+ conn_handshake(c);
+ if (c->closing || c->handshaking)
+ return;
+ }
while (c->out_sent < c->out.len) {
- ssize_t w = write(c->fd, c->out.p + c->out_sent, c->out.len - c->out_sent);
- if (w > 0) {
- c->out_sent += (size_t)w;
- continue;
- }
- if (w < 0 && (errno == EAGAIN || errno == EWOULDBLOCK))
+ ssize_t w;
+ if (c->ssl) {
+ w = SSL_write(c->ssl, c->out.p + c->out_sent,
+ (int)(c->out.len - c->out_sent));
+ if (w <= 0) {
+ int e = SSL_get_error(c->ssl, (int)w);
+ if (e == SSL_ERROR_WANT_READ || e == SSL_ERROR_WANT_WRITE)
+ return;
+ ERR_clear_error();
+ c->closing = 1;
+ return;
+ }
+ } else {
+ w = write(c->fd, c->out.p + c->out_sent,
+ c->out.len - c->out_sent);
+ if (w > 0) {
+ c->out_sent += (size_t)w;
+ continue;
+ }
+ if (w < 0 && (errno == EAGAIN || errno == EWOULDBLOCK))
+ return;
+ if (w < 0 && errno == EINTR)
+ continue;
+ c->closing = 1;
return;
- if (w < 0 && errno == EINTR)
- continue;
- c->closing = 1;
- return;
+ }
+ c->out_sent += (size_t)w;
}
c->out.len = 0;
c->out_sent = 0;
@@ -292,7 +439,10 @@ static void usage(FILE *f)
" --socket PATH unix socket (default /run/bokfd/bokfd.sock)\n"
" --backup-dir DIR snapshot destination\n"
" --export-dir DIR SIE export destination\n"
- " --log-level L error|warn|info|debug\n");
+ " --log-level L error|warn|info|debug\n"
+ "\n"
+ "config/env: tls (host:port, BOKFD_TLS), tls_cert (BOKFD_TLS_CERT),\n"
+ " tls_key (BOKFD_TLS_KEY); TLS is off unless tls is set\n");
}
static const char *opt_value(const char *arg, const char *name, int *i,
@@ -404,6 +554,32 @@ int main(int argc, char **argv)
g_cfg.tcp_addr ? g_cfg.tcp_addr : "127.0.0.1:8787");
}
+ int tlsfd = -1;
+ if (g_cfg.tls_enabled) {
+ if (tls_load_certs() != 0) {
+ log_error("cannot load TLS certificate %s / key %s",
+ g_cfg.tls_cert, g_cfg.tls_key);
+ close(lfd);
+ if (tfd >= 0)
+ close(tfd);
+ sqlite3_close(db);
+ config_free();
+ return 1;
+ }
+ tlsfd = tcp_listen(g_cfg.tls_addr ? g_cfg.tls_addr : "127.0.0.1:8788");
+ if (tlsfd < 0) {
+ log_error("cannot listen on tls %s",
+ g_cfg.tls_addr ? g_cfg.tls_addr : "127.0.0.1:8788");
+ close(lfd);
+ if (tfd >= 0)
+ close(tfd);
+ sqlite3_close(db);
+ config_free();
+ return 1;
+ }
+ log_info("bokfd listening on tls %s", g_cfg.tls_addr);
+ }
+
signal(SIGINT, on_signal);
signal(SIGTERM, on_signal);
signal(SIGPIPE, SIG_IGN);
@@ -418,29 +594,40 @@ int main(int argc, char **argv)
struct conn conns[MAX_CONNS];
size_t nconns = 0;
- struct pollfd pfds[MAX_CONNS + 2];
- struct conn *map[MAX_CONNS + 2];
+ struct pollfd pfds[MAX_CONNS + 3];
+ struct conn *map[MAX_CONNS + 3];
+ int lfds[3];
+ int ltls[3];
while (!g_stop) {
int n = 0;
int nlisteners = 0;
- pfds[n].fd = lfd;
- pfds[n].events = POLLIN;
- map[n] = NULL;
- n++;
+ lfds[nlisteners] = lfd;
+ ltls[nlisteners] = 0;
nlisteners++;
if (tfd >= 0) {
- pfds[n].fd = tfd;
+ lfds[nlisteners] = tfd;
+ ltls[nlisteners] = 0;
+ nlisteners++;
+ }
+ if (tlsfd >= 0) {
+ lfds[nlisteners] = tlsfd;
+ ltls[nlisteners] = 1;
+ nlisteners++;
+ }
+ for (int i = 0; i < nlisteners; i++) {
+ pfds[n].fd = lfds[i];
pfds[n].events = POLLIN;
map[n] = NULL;
n++;
- nlisteners++;
}
for (size_t i = 0; i < nconns; i++) {
pfds[n].fd = conns[i].fd;
- pfds[n].events = POLLIN | (conns[i].out.len > conns[i].out_sent
- ? POLLOUT
- : 0);
+ pfds[n].events = POLLIN |
+ ((conns[i].out.len > conns[i].out_sent ||
+ conns[i].handshaking)
+ ? POLLOUT
+ : 0);
map[n] = &conns[i];
n++;
}
@@ -453,7 +640,7 @@ int main(int argc, char **argv)
}
for (int i = 0; i < nlisteners; i++)
if (pfds[i].revents & POLLIN)
- accept_conns(pfds[i].fd, conns, &nconns);
+ accept_conns(lfds[i], conns, &nconns, ltls[i]);
for (int i = nlisteners; i < n; i++) {
struct conn *c = map[i];
if (c->fd < 0)
@@ -463,12 +650,18 @@ int main(int argc, char **argv)
if (pfds[i].revents & POLLOUT)
conn_flush(c);
if (c->closing && c->out.len == c->out_sent) {
+ if (c->ssl) {
+ SSL_free(c->ssl);
+ c->ssl = NULL;
+ }
close(c->fd);
buf_free(&c->in);
buf_free(&c->out);
c->fd = -1;
}
}
+ if (tlsfd >= 0)
+ tls_reload_check();
size_t keep = 0;
for (size_t i = 0; i < nconns; i++) {
if (conns[i].fd >= 0)
@@ -479,6 +672,8 @@ int main(int argc, char **argv)
log_info("shutting down");
for (size_t i = 0; i < nconns; i++) {
+ if (conns[i].ssl)
+ SSL_free(conns[i].ssl);
close(conns[i].fd);
buf_free(&conns[i].in);
buf_free(&conns[i].out);
@@ -486,6 +681,12 @@ int main(int argc, char **argv)
close(lfd);
if (tfd >= 0)
close(tfd);
+ if (tlsfd >= 0)
+ close(tlsfd);
+ if (g_tls_ctx) {
+ SSL_CTX_free(g_tls_ctx);
+ g_tls_ctx = NULL;
+ }
unlink(g_cfg.socket_path);
sessions_free_all();
sqlite3_close(db);