#include #include #include #include #include #include #include #include #include #include #include #include "smtp.h" #include "util.h" enum { SRV_NOAUTH = 0, SRV_AUTH_PLAIN, SRV_AUTH_LOGIN, SRV_RCPT_REJECT, SRV_TLS, SRV_STARTTLS }; struct srv_result { int saw_ehlo; int saw_auth; int saw_auth_plain; int saw_auth_login; int saw_starttls; size_t data_len; }; static int failures; static void check(int cond, const char *what) { if (!cond) { printf("FAIL %s\n", what); failures++; } else { printf("ok %s\n", what); } } static int contains(const unsigned char *data, size_t len, const char *needle) { size_t n = strlen(needle); if (n > len) return 0; for (size_t i = 0; i + n <= len; i++) if (memcmp(data + i, needle, n) == 0) return 1; return 0; } static int count_occ(const unsigned char *data, size_t len, const char *needle) { size_t n = strlen(needle); int c = 0; if (n == 0 || n > len) return 0; for (size_t i = 0; i + n <= len; i++) if (memcmp(data + i, needle, n) == 0) c++; return c; } static ssize_t read_full(int fd, void *buf, size_t n) { unsigned char *p = buf; size_t got = 0; while (got < n) { ssize_t r = read(fd, p + got, n - got); if (r < 0) { if (errno == EINTR) continue; return -1; } if (r == 0) break; got += (size_t)r; } return (ssize_t)got; } static int write_full(int fd, const void *buf, size_t n) { const unsigned char *p = buf; size_t done = 0; while (done < n) { ssize_t w = write(fd, p + done, n - done); if (w < 0) { if (errno == EINTR) continue; return -1; } done += (size_t)w; } return 0; } struct srv_io { int fd; SSL *ssl; unsigned char buf[4096]; size_t len; size_t pos; }; static int srv_write_all(struct srv_io *io, const char *s) { size_t n = strlen(s); size_t o = 0; while (o < n) { ssize_t w = io->ssl ? SSL_write(io->ssl, s + o, (int)(n - o)) : write(io->fd, s + o, n - o); if (w <= 0) return -1; o += (size_t)w; } return 0; } static ssize_t srv_line(struct srv_io *io, char *out, size_t outsz) { size_t o = 0; for (;;) { if (io->pos >= io->len) { ssize_t n = io->ssl ? SSL_read(io->ssl, io->buf, (int)sizeof io->buf) : read(io->fd, io->buf, sizeof io->buf); if (n <= 0) return -1; io->len = (size_t)n; io->pos = 0; } unsigned char ch = io->buf[io->pos++]; if (ch == '\n') { while (o > 0 && out[o - 1] == '\r') o--; out[o] = '\0'; return (ssize_t)o; } if (o + 1 < outsz) out[o++] = (char)ch; } } static void srv_tls_accept(struct srv_io *io) { SSL_CTX *ctx = SSL_CTX_new(TLS_server_method()); if (!ctx || SSL_CTX_use_certificate_chain_file(ctx, "tests/tls_test_cert.pem") != 1 || SSL_CTX_use_PrivateKey_file(ctx, "tests/tls_test_key.pem", SSL_FILETYPE_PEM) != 1) _exit(1); SSL *ssl = SSL_new(ctx); if (!ssl || SSL_set_fd(ssl, io->fd) != 1 || SSL_accept(ssl) != 1) _exit(1); io->ssl = ssl; } static void srv_ehlo(struct srv_io *io, int mode, int after_tls) { srv_write_all(io, "250-test.local hello\r\n"); if (mode == SRV_AUTH_PLAIN) srv_write_all(io, "250-AUTH PLAIN LOGIN\r\n"); else if (mode == SRV_AUTH_LOGIN) srv_write_all(io, "250-AUTH LOGIN\r\n"); if (mode == SRV_STARTTLS && !after_tls) srv_write_all(io, "250-STARTTLS\r\n"); srv_write_all(io, "250 SIZE 10485760\r\n"); } static int plain_ok(const char *b64) { static const unsigned char want[] = { 0, 'm', 'a', 'i', 'l', 'e', 'r', 0, 's', '3', 'c', 'r', 'e', 't' }; unsigned char *raw = NULL; size_t len = 0; if (util_b64_decode(b64, strlen(b64), &raw, &len) != 0) { if (getenv("SMTP_CHECK_DUMP")) fprintf(stderr, "server: AUTH PLAIN decode failed [%s]\n", b64); return 0; } int ok = len == sizeof want && memcmp(raw, want, len) == 0; if (!ok && getenv("SMTP_CHECK_DUMP")) fprintf(stderr, "server: AUTH PLAIN mismatch len=%zu [%s]\n", len, b64); free(raw); return ok; } static void server_main(int wfd, int mode) { int sfd = socket(AF_INET, SOCK_STREAM | SOCK_CLOEXEC, 0); if (sfd < 0) _exit(1); struct sockaddr_in sa; memset(&sa, 0, sizeof sa); sa.sin_family = AF_INET; sa.sin_addr.s_addr = htonl(INADDR_LOOPBACK); sa.sin_port = 0; if (bind(sfd, (struct sockaddr *)&sa, sizeof sa) != 0 || listen(sfd, 1) != 0) _exit(1); socklen_t slen = sizeof sa; if (getsockname(sfd, (struct sockaddr *)&sa, &slen) != 0) _exit(1); int port = ntohs(sa.sin_port); if (write_full(wfd, &port, sizeof port) != 0) _exit(1); int cfd = accept(sfd, NULL, NULL); if (cfd < 0) _exit(1); close(sfd); struct srv_io io; memset(&io, 0, sizeof io); io.fd = cfd; struct srv_result res; memset(&res, 0, sizeof res); struct buf data; buf_init(&data); char *b64_user = util_b64((const unsigned char *)"mailer", 6); char *b64_pass = util_b64((const unsigned char *)"s3cret", 6); if (mode == SRV_TLS) srv_tls_accept(&io); srv_write_all(&io, "220 test.local ESMTP\r\n"); int in_data = 0; char line[2048]; for (;;) { ssize_t n = srv_line(&io, line, sizeof line); if (n < 0) break; if (in_data) { if (strcmp(line, ".") == 0) { res.data_len = data.len; in_data = 0; if (srv_write_all(&io, "250 2.0.0 queued\r\n") != 0) break; continue; } const char *payload = line[0] == '.' ? line + 1 : line; buf_append(&data, payload, strlen(payload)); buf_append(&data, "\r\n", 2); continue; } if (strncmp(line, "EHLO", 4) == 0) { res.saw_ehlo++; srv_ehlo(&io, mode, mode == SRV_TLS || io.ssl != NULL); } else if (strcmp(line, "STARTTLS") == 0) { res.saw_starttls = 1; srv_write_all(&io, "220 2.0.0 ready\r\n"); srv_tls_accept(&io); } else if (strncmp(line, "AUTH PLAIN", 10) == 0) { res.saw_auth = 1; res.saw_auth_plain = 1; const char *arg = line + 10; while (*arg == ' ') arg++; srv_write_all(&io, plain_ok(arg) ? "235 2.7.0 authenticated\r\n" : "535 5.7.8 bad credentials\r\n"); } else if (strcmp(line, "AUTH LOGIN") == 0) { res.saw_auth = 1; res.saw_auth_login = 1; srv_write_all(&io, "334 VXNlcm5hbWU6\r\n"); if (srv_line(&io, line, sizeof line) < 0) _exit(1); if (strcmp(line, b64_user) != 0) { srv_write_all(&io, "535 5.7.8 bad credentials\r\n"); continue; } srv_write_all(&io, "334 UGFzc3dvcmQ6\r\n"); if (srv_line(&io, line, sizeof line) < 0) _exit(1); srv_write_all(&io, strcmp(line, b64_pass) == 0 ? "235 2.7.0 authenticated\r\n" : "535 5.7.8 bad credentials\r\n"); } else if (strncmp(line, "MAIL FROM:", 10) == 0) { srv_write_all(&io, "250 2.1.0 ok\r\n"); } else if (strncmp(line, "RCPT TO:", 8) == 0) { srv_write_all(&io, mode == SRV_RCPT_REJECT ? "550 5.1.1 no such user\r\n" : "250 2.1.5 ok\r\n"); } else if (strcmp(line, "DATA") == 0) { if (mode == SRV_RCPT_REJECT) { srv_write_all(&io, "503 5.5.1 need RCPT\r\n"); } else { srv_write_all(&io, "354 end with .\r\n"); in_data = 1; } } else if (strcmp(line, "QUIT") == 0) { srv_write_all(&io, "221 2.0.0 bye\r\n"); break; } else { srv_write_all(&io, "500 5.5.1 unknown command\r\n"); } } (void)write_full(wfd, &res, sizeof res); if (data.len) (void)write_full(wfd, data.p, data.len); close(wfd); close(cfd); buf_free(&data); free(b64_user); free(b64_pass); _exit(0); } struct case_result { int rc; char err[512]; struct srv_result res; unsigned char data[131072]; size_t data_len; int server_ok; }; static void run_case(int mode, struct smtp_message m, struct case_result *out) { signal(SIGPIPE, SIG_IGN); int pfd[2]; memset(out, 0, sizeof *out); if (pipe(pfd) != 0) { printf("FAIL pipe\n"); failures++; return; } pid_t pid = fork(); if (pid == 0) { close(pfd[0]); server_main(pfd[1], mode); } close(pfd[1]); int port = 0; if (read_full(pfd[0], &port, sizeof port) != (ssize_t)sizeof port) port = 0; m.port = port; out->rc = smtp_send(&m, out->err, sizeof out->err); struct srv_result res; memset(&res, 0, sizeof res); if (read_full(pfd[0], &res, sizeof res) == (ssize_t)sizeof res) { out->res = res; size_t cap = res.data_len < sizeof out->data ? res.data_len : sizeof out->data; if (cap && read_full(pfd[0], out->data, cap) == (ssize_t)cap) out->data_len = cap; } close(pfd[0]); int status = 0; waitpid(pid, &status, 0); out->server_ok = WIFEXITED(status) && WEXITSTATUS(status) == 0; } static void check_delivered(const char *what, const struct case_result *c) { check(c->server_ok && c->rc == 0, what); if (c->server_ok && c->rc == 0) return; if (!c->server_ok) printf(" server aborted\n"); if (c->rc != 0) printf(" err: %s\n", c->err); } static struct smtp_message base_message(void) { struct smtp_message m; memset(&m, 0, sizeof m); m.host = "localhost"; m.security = "plain"; m.from = "invoice@example.test"; m.from_name = "Bokf AB"; m.to = "kund@example.test"; m.subject = "Faktura 1001"; m.body = "Hej!\nBifogad faktura.\n"; return m; } static const unsigned char pdf[] = "%PDF-1.4\n1 0 obj\n<< /Type /Catalog >>\nendobj\ntrailer\n" "<< /Root 1 0 R >>\n%%EOF\n"; static void flatten(const unsigned char *data, size_t len, unsigned char **out, size_t *out_len) { unsigned char *flat = malloc(len + 1); size_t o = 0; for (size_t i = 0; i < len; i++) if (data[i] != '\r' && data[i] != '\n') flat[o++] = data[i]; flat[o] = '\0'; *out = flat; *out_len = o; } static void check_crlf(const unsigned char *data, size_t len) { int crlf = 1; int max = 0; size_t start = 0; for (size_t i = 0; i < len; i++) { if (data[i] != '\n') continue; size_t l = i - start; if (l > 0 && data[i - 1] == '\r') l--; else crlf = 0; if (l > (size_t)max) max = (int)l; start = i + 1; } check(crlf, "message uses CRLF only"); check(max > 0 && max < 998, "line lengths under 998"); } static void test_multipart(void) { struct smtp_message m = base_message(); m.from_name = "Bokf \xc3\x84rende"; m.subject = "Faktura p\xc3\xa5 svenska"; m.attach_name = "faktura.pdf"; m.attach = pdf; m.attach_len = sizeof pdf - 1; struct case_result c; run_case(SRV_NOAUTH, m, &c); check_delivered("multipart: delivered", &c); check(c.res.saw_ehlo == 1, "multipart: one EHLO"); check(!c.res.saw_auth, "multipart: no AUTH when user empty"); const unsigned char *d = c.data; size_t n = c.data_len; if (getenv("SMTP_CHECK_DUMP")) printf("--- message ---\n%.*s\n---\n", (int)n, d); char *name_word = NULL; char *subj_word = NULL; { char *b = util_b64((const unsigned char *)m.from_name, strlen(m.from_name)); size_t l = strlen(b) + 13; name_word = malloc(l); snprintf(name_word, l, "=?UTF-8?B?%s?=", b); free(b); b = util_b64((const unsigned char *)m.subject, strlen(m.subject)); l = strlen(b) + 13; subj_word = malloc(l); snprintf(subj_word, l, "=?UTF-8?B?%s?=", b); free(b); } char from_hdr[256]; snprintf(from_hdr, sizeof from_hdr, "From: %s ", name_word); char subj_hdr[256]; snprintf(subj_hdr, sizeof subj_hdr, "Subject: %s", subj_word); check(contains(d, n, "MIME-Version: 1.0\r\n"), "multipart: MIME-Version"); check(contains(d, n, from_hdr), "multipart: From with B-encoded name"); check(contains(d, n, "To: \r\n"), "multipart: To"); check(contains(d, n, subj_hdr), "multipart: B-encoded subject"); check(contains(d, n, "Date: ") && contains(d, n, "+0000\r\n"), "multipart: UTC Date"); check(contains(d, n, "Message-ID: <"), "multipart: Message-ID"); check(contains(d, n, "Content-Type: multipart/mixed; boundary=\""), "multipart: multipart content type"); check(count_occ(d, n, "Content-Type: text/plain; charset=utf-8") == 1, "multipart: text part"); check(contains(d, n, "Content-Type: application/pdf; name=\"faktura.pdf\""), "multipart: attachment type and name"); check(contains(d, n, "Content-Disposition: attachment; filename=\"faktura.pdf\""), "multipart: attachment disposition"); check(count_occ(d, n, "Content-Transfer-Encoding: base64") == 2, "multipart: two base64 parts"); unsigned char *flat = NULL; size_t flat_len = 0; flatten(d, n, &flat, &flat_len); char *b64 = util_b64(pdf, sizeof pdf - 1); check(contains(flat, flat_len, b64), "multipart: attachment base64"); free(b64); b64 = util_b64((const unsigned char *)m.body, strlen(m.body)); check(contains(flat, flat_len, b64), "multipart: body base64"); free(b64); free(flat); check(count_occ(d, n, "--=_bokf_") == 3, "multipart: boundary delimiters"); check_crlf(d, n); free(name_word); free(subj_word); } static void test_text_only(void) { struct smtp_message m = base_message(); struct case_result c; run_case(SRV_NOAUTH, m, &c); check_delivered("text-only: delivered", &c); const unsigned char *d = c.data; size_t n = c.data_len; check(contains(d, n, "From: Bokf AB \r\n"), "text-only: plain display name"); check(contains(d, n, "Subject: Faktura 1001\r\n"), "text-only: plain subject"); check(contains(d, n, "Content-Type: text/plain; charset=utf-8\r\n"), "text-only: text content type"); check(!contains(d, n, "multipart/mixed"), "text-only: not multipart"); check_crlf(d, n); } static void test_rcpt_reject(void) { struct smtp_message m = base_message(); m.attach_name = "faktura.pdf"; m.attach = pdf; m.attach_len = sizeof pdf - 1; struct case_result c; run_case(SRV_RCPT_REJECT, m, &c); check(c.server_ok, "reject: server completed session"); check(c.rc == -1, "reject: smtp_send failed"); check(strstr(c.err, "550") != NULL, "reject: err has 550"); check(strstr(c.err, "no such user") != NULL, "reject: err has server text"); check(!contains(c.data, c.data_len, "faktura"), "reject: no DATA delivered"); } static void test_auth_plain(void) { struct smtp_message m = base_message(); m.user = "mailer"; m.password = "s3cret"; struct case_result c; run_case(SRV_AUTH_PLAIN, m, &c); check_delivered("auth plain: delivered", &c); check(c.res.saw_auth_plain == 1, "auth plain: PLAIN used"); check(c.res.saw_auth_login == 0, "auth plain: LOGIN not used"); check(!contains(c.data, c.data_len, "s3cret"), "auth plain: no secret in message"); } static void test_auth_login(void) { struct smtp_message m = base_message(); m.user = "mailer"; m.password = "s3cret"; struct case_result c; run_case(SRV_AUTH_LOGIN, m, &c); check_delivered("auth login: delivered", &c); check(c.res.saw_auth_login == 1, "auth login: LOGIN used"); check(c.res.saw_auth_plain == 0, "auth login: PLAIN not used"); } static void test_auth_reject(void) { struct smtp_message m = base_message(); m.user = "mailer"; m.password = "wrong"; struct case_result c; run_case(SRV_AUTH_PLAIN, m, &c); check(c.rc == -1, "auth reject: smtp_send failed"); check(strstr(c.err, "535") != NULL, "auth reject: err has 535"); check(strstr(c.err, "s3cret") == NULL, "auth reject: no secret in err"); } static void test_validation(void) { char err[256]; check(smtp_send(NULL, err, sizeof err) == -1, "validation: NULL message"); struct smtp_message m = base_message(); m.port = 25; m.security = "bogus"; check(smtp_send(&m, err, sizeof err) == -1, "validation: bad security"); check(strstr(err, "security") != NULL, "validation: bad security text"); m = base_message(); m.from = ""; check(smtp_send(&m, err, sizeof err) == -1, "validation: empty from"); m = base_message(); m.host = NULL; check(smtp_send(&m, err, sizeof err) == -1, "validation: NULL host"); } static int tls_available(void) { return access("tests/tls_test_cert.pem", R_OK) == 0 && access("tests/tls_test_key.pem", R_OK) == 0; } static void with_test_ca(void (*fn)(struct smtp_message *, struct case_result *)) { char *saved = getenv("SSL_CERT_FILE"); char *copy = saved ? strdup(saved) : NULL; setenv("SSL_CERT_FILE", "tests/tls_test_cert.pem", 1); struct smtp_message m = base_message(); struct case_result c; fn(&m, &c); check_delivered("TLS case delivered", &c); if (saved) setenv("SSL_CERT_FILE", copy, 1); else unsetenv("SSL_CERT_FILE"); free(copy); } static void tls_case(struct smtp_message *m, struct case_result *c) { m->security = "tls"; m->attach_name = "faktura.pdf"; m->attach = pdf; m->attach_len = sizeof pdf - 1; run_case(SRV_TLS, *m, c); } static void starttls_case(struct smtp_message *m, struct case_result *c) { m->security = "starttls"; m->attach_name = "faktura.pdf"; m->attach = pdf; m->attach_len = sizeof pdf - 1; run_case(SRV_STARTTLS, *m, c); check(c->res.saw_starttls == 1, "starttls: STARTTLS sent"); check(c->res.saw_ehlo == 2, "starttls: EHLO before and after"); } int main(void) { signal(SIGPIPE, SIG_IGN); alarm(120); test_multipart(); test_text_only(); test_rcpt_reject(); test_auth_plain(); test_auth_login(); test_auth_reject(); test_validation(); if (tls_available()) { with_test_ca(tls_case); with_test_ca(starttls_case); } else { printf("skip TLS cases: tests/tls_test_cert.pem missing\n"); } if (failures) { printf("%d failures\n", failures); return 1; } puts("all ok"); return 0; }