master xplshn/aruu / shared / libutil / tls.c
  1/* see license file for copyright and license details */
  2#include "../tls.h"
  3#include "../paths.h"
  4#include "../util.h"
  5
  6#include <errno.h>
  7#include <signal.h>
  8#include <stdio.h>
  9#include <stdlib.h>
 10#include <string.h>
 11#include <unistd.h>
 12
 13#if FEATURE_USE_LIBTLS
 14#include <tls.h>
 15
 16struct TlsSocket {
 17  int         fd;
 18  int         is_tls;
 19  struct tls *ctx;
 20};
 21
 22#elif FEATURE_USE_BEARSSL
 23#include <bearssl.h>
 24
 25struct x509_noverify_ctx {
 26  const br_x509_class    *vtable;
 27  br_x509_minimal_context minimal;
 28};
 29
 30struct TlsSocket {
 31  int                      fd;
 32  int                      is_tls;
 33  br_ssl_client_context    sc;
 34  struct x509_noverify_ctx x509_noverify;
 35  unsigned char            iobuf[BR_SSL_BUFSIZE_BIDI];
 36  br_sslio_context         ioc;
 37};
 38
 39struct dn_accum {
 40  unsigned char *data;
 41  size_t         len;
 42  size_t         cap;
 43};
 44
 45static void
 46append_dn_callback(void *dn_ctx, const void *src, size_t len)
 47{
 48  struct dn_accum *accum = dn_ctx;
 49  if (accum->len + len > accum->cap) {
 50    accum->cap  = accum->len + len + 256;
 51    accum->data = erealloc(accum->data, accum->cap);
 52  }
 53  memcpy(accum->data + accum->len, src, len);
 54  accum->len += len;
 55}
 56
 57static void
 58x509_noverify_start_chain(const br_x509_class **ctx, const char *server_name)
 59{
 60  struct x509_noverify_ctx *t = (struct x509_noverify_ctx *)ctx;
 61  br_x509_minimal_vtable.start_chain((const br_x509_class **)&t->minimal, server_name);
 62}
 63
 64static void
 65x509_noverify_start_cert(const br_x509_class **ctx, uint32_t length)
 66{
 67  struct x509_noverify_ctx *t = (struct x509_noverify_ctx *)ctx;
 68  br_x509_minimal_vtable.start_cert((const br_x509_class **)&t->minimal, length);
 69}
 70
 71static void
 72x509_noverify_append(const br_x509_class **ctx, const unsigned char *buf, size_t len)
 73{
 74  struct x509_noverify_ctx *t = (struct x509_noverify_ctx *)ctx;
 75  br_x509_minimal_vtable.append((const br_x509_class **)&t->minimal, buf, len);
 76}
 77
 78static void
 79x509_noverify_end_cert(const br_x509_class **ctx)
 80{
 81  struct x509_noverify_ctx *t = (struct x509_noverify_ctx *)ctx;
 82  br_x509_minimal_vtable.end_cert((const br_x509_class **)&t->minimal);
 83}
 84
 85static unsigned
 86x509_noverify_end_chain(const br_x509_class **ctx)
 87{
 88  struct x509_noverify_ctx *t = (struct x509_noverify_ctx *)ctx;
 89  (void)br_x509_minimal_vtable.end_chain((const br_x509_class **)&t->minimal);
 90  /* always succeed and accept certificate when check_cert is 0 */
 91  return 0;
 92}
 93
 94static const br_x509_pkey *
 95x509_noverify_get_pkey(const br_x509_class *const *ctx, unsigned *usages)
 96{
 97  const struct x509_noverify_ctx *t = (const struct x509_noverify_ctx *)ctx;
 98  return br_x509_minimal_vtable.get_pkey((const br_x509_class *const *)&t->minimal, usages);
 99}
100
101static const br_x509_class x509_noverify_vtable = {
102    sizeof(struct x509_noverify_ctx),
103    x509_noverify_start_chain,
104    x509_noverify_start_cert,
105    x509_noverify_append,
106    x509_noverify_end_cert,
107    x509_noverify_end_chain,
108    x509_noverify_get_pkey
109};
110
111static int
112sock_read(void *ctx, unsigned char *buf, size_t len)
113{
114  int fd = *(int *)ctx;
115  for (;;) {
116    ssize_t rlen = read(fd, buf, len);
117    if (rlen <= 0) {
118      if (rlen < 0 && errno == EINTR)
119        continue;
120      return -1;
121    }
122    return (int)rlen;
123  }
124}
125
126static int
127sock_write(void *ctx, const unsigned char *buf, size_t len)
128{
129  int fd = *(int *)ctx;
130  for (;;) {
131    ssize_t wlen = write(fd, buf, len);
132    if (wlen <= 0) {
133      if (wlen < 0 && errno == EINTR)
134        continue;
135      return -1;
136    }
137    return (int)wlen;
138  }
139}
140
141static int
142b64_decode_char(char c)
143{
144  if (c >= 'A' && c <= 'Z')
145    return c - 'A';
146  if (c >= 'a' && c <= 'z')
147    return c - 'a' + 26;
148  if (c >= '0' && c <= '9')
149    return c - '0' + 52;
150  if (c == '+')
151    return 62;
152  if (c == '/')
153    return 63;
154  return -1;
155}
156
157static size_t
158b64_decode(const char *src, size_t src_len, unsigned char *dst)
159{
160  size_t i, j = 0;
161  int    val = 0, valb = -8;
162  for (i = 0; i < src_len; i++) {
163    int c = b64_decode_char(src[i]);
164    if (c >= 0) {
165      val = (val << 6) | c;
166      valb += 6;
167      if (valb >= 0) {
168        dst[j++] = (val >> valb) & 0xFF;
169        valb -= 8;
170      }
171    }
172  }
173  return j;
174}
175
176static int
177decode_cert_der(const unsigned char *der, size_t der_len, br_x509_trust_anchor *ta)
178{
179  br_x509_decoder_context dc;
180  struct dn_accum         accum = {NULL, 0, 0};
181  const br_x509_pkey     *pk;
182
183  br_x509_decoder_init(&dc, append_dn_callback, &accum);
184  br_x509_decoder_push(&dc, der, der_len);
185  if (br_x509_decoder_last_error(&dc) != 0) {
186    free(accum.data);
187    return 0;
188  }
189  pk = br_x509_decoder_get_pkey(&dc);
190  if (!pk) {
191    free(accum.data);
192    return 0;
193  }
194
195  ta->dn.data       = accum.data;
196  ta->dn.len        = accum.len;
197  ta->flags         = br_x509_decoder_isCA(&dc) ? BR_X509_TA_CA : 0;
198  ta->pkey.key_type = pk->key_type;
199
200  if (pk->key_type == BR_KEYTYPE_RSA) {
201    ta->pkey.key.rsa.nlen = pk->key.rsa.nlen;
202    ta->pkey.key.rsa.n    = emalloc(pk->key.rsa.nlen);
203    memcpy(ta->pkey.key.rsa.n, pk->key.rsa.n, pk->key.rsa.nlen);
204    ta->pkey.key.rsa.elen = pk->key.rsa.elen;
205    ta->pkey.key.rsa.e    = emalloc(pk->key.rsa.elen);
206    memcpy(ta->pkey.key.rsa.e, pk->key.rsa.e, pk->key.rsa.elen);
207  } else if (pk->key_type == BR_KEYTYPE_EC) {
208    ta->pkey.key.ec.curve = pk->key.ec.curve;
209    ta->pkey.key.ec.qlen  = pk->key.ec.qlen;
210    ta->pkey.key.ec.q     = emalloc(pk->key.ec.qlen);
211    memcpy(ta->pkey.key.ec.q, pk->key.ec.q, pk->key.ec.qlen);
212  } else {
213    free(accum.data);
214    return 0;
215  }
216  return 1;
217}
218
219static br_x509_trust_anchor *tas     = NULL;
220static size_t                tas_num = 0;
221
222static void
223load_ca_certs(void)
224{
225  static const char *ca_paths[] = {
226      ARUU_PATH_ETC "/ssl/certs/ca-certificates.crt",
227      ARUU_PATH_ETC "/ssl/cert.pem",
228      ARUU_PATH_ETC "/pki/tls/certs/ca-bundle.crt",
229  };
230  FILE  *fp = NULL;
231  size_t i;
232  char   line[512];
233  struct {
234    char  *data;
235    size_t len;
236    size_t cap;
237  } pem       = {NULL, 0, 0};
238  int in_cert = 0;
239
240  for (i = 0; i < LEN(ca_paths); i++) {
241    if ((fp = fopen(ca_paths[i], "r")))
242      break;
243  }
244  if (!fp) {
245    weprintf(
246        "no CA certificates found, TLS verification will "
247        "fail\n"
248    );
249    return;
250  }
251
252  while (fgets(line, sizeof(line), fp)) {
253    if (strncmp(line, "-----BEGIN CERTIFICATE-----", 27) == 0) {
254      in_cert = 1;
255      pem.len = 0;
256    } else if (strncmp(line, "-----END CERTIFICATE-----", 25) == 0) {
257      if (in_cert) {
258        unsigned char       *der     = emalloc(pem.len);
259        size_t               der_len = b64_decode(pem.data, pem.len, der);
260        br_x509_trust_anchor ta;
261        if (decode_cert_der(der, der_len, &ta)) {
262          tas            = ereallocarray(tas, tas_num + 1, sizeof(*tas));
263          tas[tas_num++] = ta;
264        }
265        free(der);
266        in_cert = 0;
267      }
268    } else if (in_cert) {
269      size_t llen = strlen(line);
270      while (llen > 0 && (line[llen - 1] == '\r' || line[llen - 1] == '\n'))
271        llen--;
272      if (pem.len + llen > pem.cap) {
273        pem.cap  = pem.len + llen + 1024;
274        pem.data = erealloc(pem.data, pem.cap);
275      }
276      memcpy(pem.data + pem.len, line, llen);
277      pem.len += llen;
278    }
279  }
280  free(pem.data);
281  fclose(fp);
282}
283#else
284struct TlsSocket {
285  int fd;
286  int is_tls;
287};
288#endif
289
290struct TlsSocket *
291tlss_connect(int fd, const char *host, int check_cert, int is_tls)
292{
293  struct TlsSocket *s;
294
295  s         = emalloc(sizeof(*s));
296  s->fd     = fd;
297  s->is_tls = is_tls;
298
299  if (!is_tls)
300    return s;
301#if !FEATURE_USE_LIBTLS && !FEATURE_USE_BEARSSL && !FEATURE_USE_OPENSSL
302  (void)host;
303  (void)check_cert;
304#endif
305
306#if FEATURE_USE_LIBTLS
307  {
308    struct tls_config *cfg;
309
310    s->ctx = tls_client();
311    if (!s->ctx) {
312      weprintf("tls_client failed\n");
313      free(s);
314      return NULL;
315    }
316    cfg = tls_config_new();
317    if (!cfg) {
318      weprintf("tls_config_new failed\n");
319      tls_free(s->ctx);
320      free(s);
321      return NULL;
322    }
323    if (!check_cert) {
324      tls_config_insecure_noverifycert(cfg);
325      tls_config_insecure_noverifyname(cfg);
326    }
327    if (tls_configure(s->ctx, cfg) < 0) {
328      weprintf("tls_configure: %s\n", tls_error(s->ctx));
329      tls_config_free(cfg);
330      tls_free(s->ctx);
331      free(s);
332      return NULL;
333    }
334    tls_config_free(cfg);
335
336    if (tls_connect_socket(s->ctx, fd, host) < 0) {
337      weprintf("tls_connect_socket: %s\n", tls_error(s->ctx));
338      tls_free(s->ctx);
339      free(s);
340      return NULL;
341    }
342  }
343#elif FEATURE_USE_BEARSSL
344  {
345    signal(SIGPIPE, SIG_IGN);
346
347    if (check_cert) {
348      static int ca_loaded = 0;
349      if (!ca_loaded) {
350        load_ca_certs();
351        ca_loaded = 1;
352      }
353      br_ssl_client_init_full(&s->sc, &s->x509_noverify.minimal, tas, tas_num);
354    } else {
355      br_ssl_client_init_full(&s->sc, &s->x509_noverify.minimal, NULL, 0);
356      s->x509_noverify.vtable = &x509_noverify_vtable;
357      br_ssl_engine_set_x509(&s->sc.eng, &s->x509_noverify.vtable);
358    }
359
360    br_ssl_engine_set_buffer(&s->sc.eng, s->iobuf, sizeof(s->iobuf), 1);
361    br_ssl_client_reset(&s->sc, host, 0);
362    br_sslio_init(&s->ioc, &s->sc.eng, sock_read, &s->fd, sock_write, &s->fd);
363
364    if (br_sslio_flush(&s->ioc) < 0) {
365      weprintf("TLS handshake failed with %s\n", host);
366      free(s);
367      return NULL;
368    }
369  }
370#else
371  weprintf(
372      "TLS not supported, compile with FEATURE_USE_LIBTLS or "
373      "FEATURE_USE_BEARSSL\n"
374  );
375  free(s);
376  return NULL;
377#endif
378
379  return s;
380}
381
382ssize_t
383tlss_read(struct TlsSocket *s, void *buf, size_t len)
384{
385  if (!s->is_tls) {
386    for (;;) {
387      ssize_t r = read(s->fd, buf, len);
388      if (r < 0 && errno == EINTR)
389        continue;
390      return r;
391    }
392  }
393#if FEATURE_USE_LIBTLS
394  for (;;) {
395    ssize_t r = tls_read(s->ctx, buf, len);
396    if (r == TLS_WANT_POLLIN || r == TLS_WANT_POLLOUT)
397      continue;
398    return r;
399  }
400#elif FEATURE_USE_BEARSSL
401  return br_sslio_read(&s->ioc, buf, len);
402#else
403  return -1;
404#endif
405}
406
407ssize_t
408tlss_write(struct TlsSocket *s, const void *buf, size_t len)
409{
410  if (!s->is_tls) {
411    for (;;) {
412      ssize_t r = write(s->fd, buf, len);
413      if (r < 0 && errno == EINTR)
414        continue;
415      return r;
416    }
417  }
418#if FEATURE_USE_LIBTLS
419  for (;;) {
420    ssize_t r = tls_write(s->ctx, buf, len);
421    if (r == TLS_WANT_POLLIN || r == TLS_WANT_POLLOUT)
422      continue;
423    return r;
424  }
425#elif FEATURE_USE_BEARSSL
426  {
427    int r = br_sslio_write_all(&s->ioc, buf, len);
428    if (r < 0)
429      return -1;
430    if (br_sslio_flush(&s->ioc) < 0)
431      return -1;
432    return len;
433  }
434#else
435  return -1;
436#endif
437}
438
439void
440tlss_close(struct TlsSocket *s, int close_fd)
441{
442  if (s->is_tls) {
443#if FEATURE_USE_LIBTLS
444    tls_close(s->ctx);
445    tls_free(s->ctx);
446#elif FEATURE_USE_BEARSSL
447    br_sslio_close(&s->ioc);
448#endif
449  }
450  if (close_fd)
451    close(s->fd);
452  free(s);
453}