libdpf/test/tests/security_test.cpp

1143 lines
39 KiB
C++
Raw Normal View History

#include <gtest/gtest.h>
#include <cstdint>
#include <cstdio>
#include <cstring>
#include <fstream>
#include <functional>
#include <memory>
#include <string>
#include <thread>
#include <sys/stat.h>
#include <unistd.h>
#include <sys/wait.h>
#include "dpf/compose.hpp"
#include "dpf/log.hpp"
#include "dpf/net/client_link.hpp"
#include "dpf/net/connect.hpp"
#include "dpf/net/identity.hpp"
#include "dpf/net/party_session.hpp"
#include "dpf/net/security.hpp"
#include "dpf/net/tls.hpp"
#include "dpf/online_session.hpp"
#include "dpf/party_runner.hpp"
#include "dpf/run_config.hpp"
#include <openssl/pem.h>
#include <openssl/x509v3.h>
namespace
{
using dpf::net::identity;
using dpf::net::public_key;
std::string temp_path(const std::string & name)
{
return "/tmp/libdpf_security_" + std::to_string(::getpid()) + "_" + name;
}
/// Two connected TCP sockets, each on its own io_context.
struct socket_pair
{
asio::io_context io_a, io_b;
asio::ip::tcp::socket a{io_a}, b{io_b};
socket_pair()
{
asio::ip::tcp::acceptor acc(io_a);
dpf::net::open_listener(acc, 0);
const auto port = acc.local_endpoint().port();
std::thread t([&] {
dpf::net::connect_until(b, "127.0.0.1", port, std::chrono::seconds(5));
});
dpf::net::accept_until(acc, a, std::chrono::seconds(5));
t.join();
a.set_option(asio::ip::tcp::no_delay(true));
b.set_option(asio::ip::tcp::no_delay(true));
}
};
/// Run `io` until `done` or `deadline`, restarting it when it runs dry.
void run_until(asio::io_context & io, const bool & done,
std::chrono::steady_clock::time_point deadline)
{
while (!done && std::chrono::steady_clock::now() < deadline)
{
if (io.stopped())
io.restart();
io.run_one_for(std::chrono::milliseconds(10));
}
}
/// Run `server` and `client` on two threads; throw both failures' messages.
void both(const std::function<void()> & server, const std::function<void()> & client)
{
std::string e1, e2;
std::thread ts([&] {
try
{
server();
}
catch (const std::exception & e)
{
e1 = e.what();
}
});
try
{
client();
}
catch (const std::exception & e)
{
e2 = e.what();
}
ts.join();
if (!e1.empty() || !e2.empty())
throw std::runtime_error("server: " + e1 + " | client: " + e2);
}
} // namespace
TEST(Identity, KeyFileRoundTripAndMode)
{
const auto path = temp_path("id.key");
::unlink(path.c_str());
const auto id = identity::generate();
id.save(path);
struct stat st{};
ASSERT_EQ(::stat(path.c_str(), &st), 0);
EXPECT_EQ(st.st_mode & 0777, 0600u);
EXPECT_THROW(id.save(path), std::runtime_error);
const auto back = identity::load(path);
EXPECT_EQ(back.key(), id.key());
EXPECT_EQ(id.key().base64().size(), 44u);
EXPECT_EQ(public_key::parse(id.key().base64()), id.key());
const auto pub = temp_path("id.pub");
{
std::ofstream out(pub);
out << "# peer 1\n" << id.key().base64() << "\n";
}
EXPECT_EQ(public_key::parse("file:" + pub), id.key());
EXPECT_THROW(public_key::parse("not-a-key"), std::invalid_argument);
EXPECT_THROW(public_key::parse(id.key().base64().substr(0, 40)), std::invalid_argument);
::unlink(path.c_str());
::unlink(pub.c_str());
}
TEST(Identity, DevelopmentIsFixedAndMarked)
{
const auto & a = identity::development();
const auto & b = identity::development();
EXPECT_TRUE(a.is_development());
EXPECT_EQ(a.key(), b.key());
EXPECT_FALSE(identity::generate().is_development());
EXPECT_NE(identity::generate().key(), identity::generate().key());
}
namespace
{
/// Party-link handshake between `ida` (accepting) and `idb` (connecting).
void peer_handshake(const identity & ida, const identity & idb,
const dpf::net::peer_security & pa, const dpf::net::peer_security & pb,
dpf::net::link_security & sa, dpf::net::link_security & sb)
{
socket_pair sp;
auto ca = dpf::net::make_peer_tls_context(ida);
auto cb = dpf::net::make_peer_tls_context(idb);
dpf::net::tls_stream ta(std::move(sp.a), *ca);
dpf::net::tls_stream tb(std::move(sp.b), *cb);
both(
[&] {
dpf::net::tls_handshake(sp.io_a, ta, true, std::chrono::seconds(5), "p0");
sa = dpf::net::tls_describe(ta);
try
{
dpf::net::check_peer(sa, pa, 1, "party 1");
}
catch (...)
{
std::error_code ec;
ta.lowest_layer().close(ec);
throw;
}
std::uint8_t in[5];
dpf::net::tls_read(sp.io_a, ta, in, 5, std::chrono::seconds(5), "read");
dpf::net::tls_write(sp.io_a, ta, in, 5, std::chrono::seconds(5), "echo");
},
[&] {
dpf::net::tls_handshake(sp.io_b, tb, false, std::chrono::seconds(5), "p1");
sb = dpf::net::tls_describe(tb);
dpf::net::check_peer(sb, pb, 0, "party 0");
const std::uint8_t out[5] = {'h', 'e', 'l', 'l', 'o'};
std::uint8_t back[5] = {};
dpf::net::tls_write(sp.io_b, tb, out, 5, std::chrono::seconds(5), "write");
dpf::net::tls_read(sp.io_b, tb, back, 5, std::chrono::seconds(5), "read");
if (std::memcmp(out, back, 5) != 0)
throw std::runtime_error("echo mismatch");
});
}
} // namespace
TEST(PeerTls, NoKeysEncryptsUnauthenticated)
{
const auto a = identity::generate();
const auto b = identity::generate();
dpf::net::peer_security pa, pb;
dpf::net::link_security sa, sb;
peer_handshake(a, b, pa, pb, sa, sb);
EXPECT_TRUE(sa.encrypted);
EXPECT_EQ(sa.protocol, "TLSv1.3");
EXPECT_EQ(sa.peer_auth, "none");
EXPECT_EQ(sb.peer_auth, "none");
ASSERT_TRUE(sa.peer_key.has_value());
EXPECT_EQ(*sa.peer_key, b.key());
EXPECT_EQ(*sb.peer_key, a.key());
}
TEST(PeerTls, KeyAuthenticatesOneDirection)
{
const auto a = identity::generate();
const auto b = identity::generate();
dpf::net::peer_security pa, pb;
pa.trusted[1] = b.key();
dpf::net::link_security sa, sb;
peer_handshake(a, b, pa, pb, sa, sb);
EXPECT_EQ(sa.peer_auth, "key");
EXPECT_EQ(sb.peer_auth, "none");
pb.trusted[0] = a.key();
peer_handshake(a, b, pa, pb, sa, sb);
EXPECT_EQ(sa.peer_auth, "key");
EXPECT_EQ(sb.peer_auth, "key");
}
TEST(PeerTls, WrongKeyNamesBothKeys)
{
const auto a = identity::generate();
const auto b = identity::generate();
const auto other = identity::generate();
dpf::net::peer_security pa, pb;
pa.trusted[1] = other.key();
dpf::net::link_security sa, sb;
try
{
peer_handshake(a, b, pa, pb, sa, sb);
FAIL() << "a mismatched key must fail the link";
}
catch (const std::exception & e)
{
const std::string what = e.what();
EXPECT_NE(what.find(b.key().base64()), std::string::npos) << what;
EXPECT_NE(what.find(other.key().base64()), std::string::npos) << what;
}
}
namespace
{
struct client_result
{
dpf::net::link_security client;
dpf::net::link_security server;
};
client_result client_handshake(const dpf::net::server_security & ss,
const dpf::net::client_security & cs, const std::string & host = "")
{
socket_pair sp;
bool dev = false;
auto sctx = dpf::net::make_server_tls_context(ss, dev);
auto cctx = dpf::net::make_client_tls_context(cs);
dpf::net::tls_stream ts(std::move(sp.a), *sctx);
dpf::net::tls_stream tc(std::move(sp.b), *cctx);
const std::string name = cs.server_name.empty() ? host : cs.server_name;
if (!cs.ca_file.empty())
dpf::net::tls_expect_host(tc, name);
client_result r;
both(
[&] {
dpf::net::tls_handshake(sp.io_a, ts, true, std::chrono::seconds(5), "server");
r.server = dpf::net::tls_describe(ts);
dpf::net::check_client(r.server, ss);
std::uint8_t in[1];
std::error_code ec;
asio::read(ts, asio::buffer(in, 1), ec);
},
[&] {
dpf::net::tls_handshake(sp.io_b, tc, false, std::chrono::seconds(5), "client");
r.client = dpf::net::tls_describe(tc);
std::error_code ec;
try
{
dpf::net::check_server(r.client, tc, cs, "localhost");
}
catch (...)
{
tc.lowest_layer().close(ec);
throw;
}
const std::uint8_t b = 1;
dpf::net::tls_write(sp.io_b, tc, &b, 1, std::chrono::seconds(5), "write");
});
return r;
}
/// A throwaway CA and a server certificate for `name`, as PEM files.
void make_ca_chain(const std::string & name, const std::string & ca_pem,
const std::string & cert_pem, const std::string & key_pem)
{
auto keygen = [] {
EVP_PKEY * k = nullptr;
EVP_PKEY_CTX * c = EVP_PKEY_CTX_new_id(EVP_PKEY_ED25519, nullptr);
EVP_PKEY_keygen_init(c);
EVP_PKEY_keygen(c, &k);
EVP_PKEY_CTX_free(c);
return k;
};
auto cert = [](EVP_PKEY * subject, EVP_PKEY * issuer_key, X509 * issuer,
const char * cn, bool ca, const std::string & dns) {
X509 * x = X509_new();
X509_set_version(x, 2);
ASN1_INTEGER_set(X509_get_serialNumber(x), ca ? 1 : 2);
X509_gmtime_adj(X509_getm_notBefore(x), -3600);
X509_gmtime_adj(X509_getm_notAfter(x), 86400);
X509_set_pubkey(x, subject);
X509_NAME * n = X509_get_subject_name(x);
X509_NAME_add_entry_by_txt(n, "CN", MBSTRING_ASC,
reinterpret_cast<const unsigned char *>(cn), -1, -1, 0);
X509_set_issuer_name(x, issuer ? X509_get_subject_name(issuer) : n);
X509V3_CTX v3;
X509V3_set_ctx_nodb(&v3);
X509V3_set_ctx(&v3, issuer ? issuer : x, x, nullptr, nullptr, 0);
auto add = [&](int nid, const char * value) {
X509_EXTENSION * e = X509V3_EXT_conf_nid(nullptr, &v3, nid, value);
X509_add_ext(x, e, -1);
X509_EXTENSION_free(e);
};
add(NID_basic_constraints, ca ? "critical,CA:TRUE" : "CA:FALSE");
if (ca)
add(NID_key_usage, "critical,keyCertSign,cRLSign");
else
add(NID_subject_alt_name, ("DNS:" + dns).c_str());
X509_sign(x, issuer_key, nullptr);
return x;
};
EVP_PKEY * cak = keygen();
EVP_PKEY * srvk = keygen();
X509 * cacert = cert(cak, cak, nullptr, "test CA", true, "");
X509 * srv = cert(srvk, cak, cacert, name.c_str(), false, name);
FILE * f = std::fopen(ca_pem.c_str(), "w");
PEM_write_X509(f, cacert);
std::fclose(f);
f = std::fopen(cert_pem.c_str(), "w");
PEM_write_X509(f, srv);
std::fclose(f);
f = std::fopen(key_pem.c_str(), "w");
PEM_write_PrivateKey(f, srvk, nullptr, nullptr, 0, nullptr, nullptr);
std::fclose(f);
X509_free(cacert);
X509_free(srv);
EVP_PKEY_free(cak);
EVP_PKEY_free(srvk);
}
} // namespace
namespace
{
/// Forwards bytes between a client and a server and keeps a copy of both
/// directions, so a test can look at what crossed the wire.
class recording_relay
{
public:
explicit recording_relay(unsigned short upstream) : upstream_(upstream)
{
dpf::net::open_listener(acc_, 0);
port_ = acc_.local_endpoint().port();
thread_ = std::thread([this] { run(); });
}
~recording_relay()
{
std::error_code ec;
acc_.close(ec);
{
std::lock_guard<std::mutex> lock(mu_);
for (int fd : fds_)
::shutdown(fd, SHUT_RDWR);
}
if (thread_.joinable())
thread_.join();
}
unsigned short port() const { return port_; }
std::string seen()
{
std::lock_guard<std::mutex> lock(mu_);
return seen_;
}
private:
void run()
{
try
{
asio::ip::tcp::socket in(io_), out(io_);
dpf::net::accept_until(acc_, in, std::chrono::seconds(10));
dpf::net::connect_until(out, "127.0.0.1", upstream_, std::chrono::seconds(10));
{
std::lock_guard<std::mutex> lock(mu_);
fds_ = {in.native_handle(), out.native_handle()};
}
std::thread back([&] { pump(out, in); });
pump(in, out);
back.join();
}
catch (...)
{
}
}
void pump(asio::ip::tcp::socket & from, asio::ip::tcp::socket & to)
{
std::uint8_t buf[4096];
std::error_code ec;
for (;;)
{
const std::size_t n = from.read_some(asio::buffer(buf), ec);
if (ec || n == 0)
break;
{
std::lock_guard<std::mutex> lock(mu_);
seen_.append(reinterpret_cast<const char *>(buf), n);
}
asio::write(to, asio::buffer(buf, n), ec);
if (ec)
break;
}
std::error_code e;
to.shutdown(asio::socket_base::shutdown_send, e);
}
unsigned short upstream_;
asio::io_context io_;
asio::ip::tcp::acceptor acc_{io_};
unsigned short port_ = 0;
std::mutex mu_;
std::string seen_;
std::vector<int> fds_;
std::thread thread_;
};
const std::string k_marker = "SECRET-SHARE-MARKER-0123456789";
/// Party 0 accepts (through a relay), party 1 connects; party 1 writes the
/// marker on lane 1 and party 0 reads it. Returns what crossed the wire.
template <bool Tls, bool Parallel>
std::string relay_exchange()
{
asio::io_context io0, io1;
asio::ip::tcp::acceptor acc(io0);
dpf::net::open_listener(acc, 0);
recording_relay relay(acc.local_endpoint().port());
const std::size_t lanes = Parallel ? 1 : 2;
const std::size_t lane = Parallel ? 0 : 1;
(void)lanes;
std::string got;
both(
[&] {
asio::ip::tcp::socket s(io0);
dpf::net::accept_until(acc, s, std::chrono::seconds(5));
s.set_option(asio::ip::tcp::no_delay(true));
std::unique_ptr<dpf::net::async_stream_array> arr;
auto id = identity::generate();
auto ctx = dpf::net::make_peer_tls_context(id);
if constexpr (Tls)
{
dpf::net::tls_stream t(std::move(s), *ctx);
dpf::net::tls_handshake(io0, t, true, std::chrono::seconds(5), "p0");
if constexpr (Parallel)
{
std::vector<dpf::net::tls_stream> v;
v.push_back(std::move(t));
arr = std::make_unique<dpf::net::async_tls_parallel_stream_array>(
io0, std::move(v));
}
else
arr = std::make_unique<dpf::net::async_tls_mux_stream_array>(io0,
std::move(t), 0, 1, lanes);
}
else
arr = std::make_unique<dpf::net::async_mux_stream_array>(io0,
std::move(s), 0, 1, lanes);
std::string in(k_marker.size() * 4, '\0');
bool done = false;
std::error_code rec;
arr->async_read(lane, in.data(), in.size(), [&](const std::error_code & ec) {
rec = ec;
done = true;
});
const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(5);
run_until(io0, done, deadline);
if (!done || rec)
throw std::runtime_error("read failed: " + rec.message() + " (done "
+ std::to_string(done) + ", relay saw "
+ std::to_string(relay.seen().size()) + " bytes)");
got = in;
// The peer closes after writing: the next read ends cleanly.
done = false;
char extra = 0;
arr->async_read(lane, &extra, 1, [&](const std::error_code & ec) {
rec = ec;
done = true;
});
run_until(io0, done, deadline);
if (rec != asio::error::eof)
throw std::runtime_error("expected eof after peer close, got "
+ rec.message());
arr.reset();
for (int i = 0; i < 50; ++i)
io0.poll();
},
[&] {
asio::ip::tcp::socket s(io1);
dpf::net::connect_until(s, "127.0.0.1", relay.port(), std::chrono::seconds(5));
s.set_option(asio::ip::tcp::no_delay(true));
std::unique_ptr<dpf::net::async_stream_array> arr;
auto id = identity::generate();
auto ctx = dpf::net::make_peer_tls_context(id);
if constexpr (Tls)
{
dpf::net::tls_stream t(std::move(s), *ctx);
dpf::net::tls_handshake(io1, t, false, std::chrono::seconds(5), "p1");
if constexpr (Parallel)
{
std::vector<dpf::net::tls_stream> v;
v.push_back(std::move(t));
arr = std::make_unique<dpf::net::async_tls_parallel_stream_array>(
io1, std::move(v));
}
else
arr = std::make_unique<dpf::net::async_tls_mux_stream_array>(io1,
std::move(t), 1, 0, lanes);
}
else
arr = std::make_unique<dpf::net::async_mux_stream_array>(io1,
std::move(s), 1, 0, lanes);
const std::string out = k_marker + k_marker + k_marker + k_marker;
bool done = false;
arr->async_write(lane, out.data(), out.size(),
[&](const std::error_code &) { done = true; });
const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(5);
run_until(io1, done, deadline);
arr.reset(); // graceful: drains, then closes
for (int i = 0; i < 50; ++i)
io1.poll();
});
if (got.find(k_marker) != 0)
throw std::runtime_error("payload did not arrive intact");
return relay.seen();
}
} // namespace
TEST(TlsBackends, PlainThroughRelay)
{
const auto plain = relay_exchange<false, false>();
EXPECT_NE(plain.find(k_marker), std::string::npos);
}
TEST(TlsBackends, TlsMuxDirect)
{
socket_pair sp;
auto ca = dpf::net::make_peer_tls_context(identity::generate());
auto cb = dpf::net::make_peer_tls_context(identity::generate());
dpf::net::tls_stream ta(std::move(sp.a), *ca);
dpf::net::tls_stream tb(std::move(sp.b), *cb);
std::string got(8, '\0');
both(
[&] {
dpf::net::tls_handshake(sp.io_a, ta, true, std::chrono::seconds(5), "a");
dpf::net::async_tls_mux_stream_array m(sp.io_a, std::move(ta), 0, 1, 2);
bool done = false;
std::error_code rec;
m.async_read(1, got.data(), 8, [&](const std::error_code & ec) {
rec = ec;
done = true;
});
const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(3);
run_until(sp.io_a, done, deadline);
if (!done)
throw std::runtime_error("direct read timed out; stats in="
+ std::to_string(m.stats().bytes_in));
// The kernel counts TLS records and the handshake too.
const auto st = m.stats();
if (st.socket_bytes_in <= st.bytes_in + 22)
throw std::runtime_error("socket bytes " + std::to_string(st.socket_bytes_in)
+ " vs frame bytes " + std::to_string(st.bytes_in));
},
[&] {
dpf::net::tls_handshake(sp.io_b, tb, false, std::chrono::seconds(5), "b");
dpf::net::async_tls_mux_stream_array m(sp.io_b, std::move(tb), 1, 0, 2);
bool done = false;
m.async_write(1, "abcdefgh", 8, [&](const std::error_code &) { done = true; });
const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(3);
run_until(sp.io_b, done, deadline);
for (int i = 0; i < 100; ++i)
sp.io_b.run_one_for(std::chrono::milliseconds(5));
if (!done)
throw std::runtime_error("direct write timed out; stats out="
+ std::to_string(m.stats().bytes_out));
});
EXPECT_EQ(got, "abcdefgh");
}
TEST(TlsBackends, MuxCiphertextOnTheWire)
{
const auto plain = relay_exchange<false, false>();
EXPECT_NE(plain.find(k_marker), std::string::npos);
const auto tls = relay_exchange<true, false>();
EXPECT_EQ(tls.find(k_marker), std::string::npos);
EXPECT_GT(tls.size(), k_marker.size() * 4);
}
TEST(TlsBackends, ParallelCiphertextOnTheWire)
{
const auto tls = relay_exchange<true, true>();
EXPECT_EQ(tls.find(k_marker), std::string::npos);
EXPECT_GT(tls.size(), k_marker.size() * 4);
}
TEST(ClientTls, DevelopmentCertificateIsTheDefault)
{
dpf::net::server_security ss;
dpf::net::client_security cs;
const auto r = client_handshake(ss, cs);
EXPECT_EQ(r.client.peer_auth, "development");
EXPECT_EQ(r.server.peer_auth, "none");
EXPECT_EQ(r.client.protocol, "TLSv1.3");
}
TEST(ClientTls, UnknownServerKeyIsRefusedUntilPinned)
{
const auto id = std::make_shared<identity>(identity::generate());
dpf::net::server_security ss;
ss.self = id;
dpf::net::client_security cs;
try
{
(void)client_handshake(ss, cs);
FAIL() << "a client with no pins must refuse a non-development key";
}
catch (const std::exception & e)
{
const std::string what = e.what();
EXPECT_NE(what.find("does not trust"), std::string::npos) << what;
EXPECT_NE(what.find(id->key().base64()), std::string::npos) << what;
}
cs.pins.push_back(id->key());
EXPECT_EQ(client_handshake(ss, cs).client.peer_auth, "key");
dpf::net::client_security off;
off.verify = false;
EXPECT_EQ(client_handshake(ss, off).client.peer_auth, "none");
// Once a client pins anything, the development certificate is refused.
dpf::net::server_security dev;
EXPECT_THROW((void)client_handshake(dev, cs), std::runtime_error);
}
TEST(ClientTls, CaChainAndHostName)
{
const auto ca = temp_path("ca.pem");
const auto crt = temp_path("srv.pem");
const auto key = temp_path("srv.key");
make_ca_chain("localhost", ca, crt, key);
dpf::net::server_security ss;
ss.cert_file = crt;
ss.key_file = key;
dpf::net::client_security cs;
cs.ca_file = ca;
cs.server_name = "localhost";
EXPECT_EQ(client_handshake(ss, cs).client.peer_auth, "ca");
cs.server_name = "elsewhere.example";
EXPECT_THROW((void)client_handshake(ss, cs), std::runtime_error);
::unlink(ca.c_str());
::unlink(crt.c_str());
::unlink(key.c_str());
}
TEST(ClientTls, ServerChecksPinnedClientKeys)
{
const auto cid = std::make_shared<identity>(identity::generate());
dpf::net::server_security ss;
ss.client_pins.push_back(cid->key());
dpf::net::client_security cs;
cs.self = cid;
const auto r = client_handshake(ss, cs);
EXPECT_EQ(r.server.peer_auth, "key");
EXPECT_EQ(r.client.peer_auth, "development");
}
// ---------------------------------------------------------------------------
// Party sessions
// ---------------------------------------------------------------------------
namespace
{
struct session_pair
{
std::string error[2];
dpf::net::link_security sec[2];
std::uint64_t got[2] = {0, 0};
};
/// Two party_sessions join in-process with `opts[me]`, swap one word, and
/// report each side's view of the link (or its error).
session_pair join_two(const dpf::net::session_options (&opts)[2],
const std::string & role_prefix = "")
{
auto ports = dpf::net::make_mesh_ports(2);
session_pair out;
auto run = [&](unsigned me) {
dpf::log::role_scope role(role_prefix.empty() ? dpf::log::role()
: role_prefix + std::to_string(me));
try
{
asio::io_context io;
dpf::net::party_session s(io, me, 2, opts[me]);
s.join("127.0.0.1", ports);
out.sec[me] = s.edge_security(1 - me);
std::uint64_t mine = 100 + me;
bool wrote = false, read = false;
s.peer(1 - me).async_write(0, &mine, 8, [&](const std::error_code &) {
wrote = true;
});
s.peer(1 - me).async_read(0, &out.got[me], 8, [&](const std::error_code &) {
read = true;
});
const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(5);
while (!(wrote && read) && std::chrono::steady_clock::now() < deadline)
{
if (io.stopped())
io.restart();
io.run_one_for(std::chrono::milliseconds(10));
}
}
catch (const std::exception & e)
{
out.error[me] = e.what();
}
};
std::thread t0([&] { run(0); });
std::thread t1([&] { run(1); });
t0.join();
t1.join();
return out;
}
dpf::net::session_options with_limits(dpf::net::session_options o)
{
o.limits.accept = std::chrono::milliseconds(3000);
o.limits.connect = std::chrono::milliseconds(3000);
o.limits.handshake = std::chrono::milliseconds(3000);
o.limits.join = std::chrono::milliseconds(3000);
return o;
}
} // namespace
TEST(Session, EncryptedByDefaultWithoutKeys)
{
dpf::net::session_options o[2] = {with_limits({}), with_limits({})};
const auto r = join_two(o);
ASSERT_TRUE(r.error[0].empty() && r.error[1].empty()) << r.error[0] << r.error[1];
for (int me = 0; me < 2; ++me)
{
EXPECT_TRUE(r.sec[me].encrypted);
EXPECT_EQ(r.sec[me].protocol, "TLSv1.3");
EXPECT_EQ(r.sec[me].peer_auth, "none");
EXPECT_FALSE(r.sec[me].peer_verified_us);
EXPECT_EQ(r.got[me], 101u - me);
}
}
TEST(Session, KeysAuthenticatePerDirection)
{
auto k0 = std::make_shared<identity>(identity::generate());
auto k1 = std::make_shared<identity>(identity::generate());
dpf::net::session_options o[2] = {with_limits({}), with_limits({})};
o[0].security.self = k0;
o[1].security.self = k1;
o[0].security.trusted[1] = k1->key();
for (auto kind : {dpf::net::transport::mux, dpf::net::transport::parallel})
{
o[0].kind = o[1].kind = kind;
o[0].n_lanes = o[1].n_lanes = 2;
auto r = join_two(o);
ASSERT_TRUE(r.error[0].empty() && r.error[1].empty()) << r.error[0] << r.error[1];
EXPECT_EQ(r.sec[0].peer_auth, "key");
EXPECT_EQ(r.sec[1].peer_auth, "none");
EXPECT_FALSE(r.sec[0].peer_verified_us);
EXPECT_TRUE(r.sec[1].peer_verified_us);
ASSERT_TRUE(r.sec[1].peer_key.has_value());
EXPECT_EQ(*r.sec[1].peer_key, k0->key());
}
o[1].security.trusted[0] = k0->key();
auto r = join_two(o);
ASSERT_TRUE(r.error[0].empty() && r.error[1].empty()) << r.error[0] << r.error[1];
EXPECT_EQ(r.sec[0].peer_auth, "key");
EXPECT_EQ(r.sec[1].peer_auth, "key");
EXPECT_TRUE(r.sec[0].peer_verified_us);
EXPECT_TRUE(r.sec[1].peer_verified_us);
}
TEST(Session, WrongKeyFailsBothEndsQuickly)
{
auto k1 = std::make_shared<identity>(identity::generate());
dpf::net::session_options o[2] = {with_limits({}), with_limits({})};
o[1].security.self = k1;
o[0].security.trusted[1] = identity::generate().key();
const auto t0 = std::chrono::steady_clock::now();
const auto r = join_two(o);
EXPECT_LT(std::chrono::steady_clock::now() - t0, std::chrono::seconds(2));
EXPECT_NE(r.error[0].find(k1->key().base64()), std::string::npos) << r.error[0];
EXPECT_FALSE(r.error[1].empty());
}
TEST(Session, MixedEncryptionSettingsAreNamed)
{
dpf::net::session_options o[2] = {with_limits({}), with_limits({})};
o[1].security.encrypt = false;
const auto r = join_two(o);
const std::string both = r.error[0] + " | " + r.error[1];
EXPECT_NE(both.find("encryption"), std::string::npos) << both;
}
TEST(Session, SctpIsRefusedWhileEncrypted)
{
asio::io_context io;
dpf::net::session_options o;
o.kind = dpf::net::transport::sctp;
try
{
dpf::net::party_session s(io, 0, 2, o);
FAIL() << "an encrypted session must refuse SCTP";
}
catch (const std::invalid_argument & e)
{
EXPECT_NE(std::string(e.what()).find("cannot be encrypted"), std::string::npos);
}
if (dpf::net::sctp_available())
{
o.security.encrypt = false;
EXPECT_NO_THROW(dpf::net::party_session(io, 0, 2, o));
}
}
TEST(Session, LackOfKeysIsLogged)
{
const auto path = temp_path("lack.log");
::unlink(path.c_str());
dpf::log::settings ls;
ls.sinks = "file:" + path;
dpf::log::configure(ls);
dpf::net::session_options o[2] = {with_limits({}), with_limits({})};
const auto r = join_two(o, "lack-test-p");
dpf::log::configure(dpf::log::settings{});
ASSERT_TRUE(r.error[0].empty() && r.error[1].empty()) << r.error[0] << r.error[1];
std::ifstream in(path);
const std::string text((std::istreambuf_iterator<char>(in)),
std::istreambuf_iterator<char>());
EXPECT_NE(text.find("ev=security.no_identity"), std::string::npos) << text;
EXPECT_NE(text.find("ev=security.unauthenticated"), std::string::npos) << text;
EXPECT_NE(text.find("ephemeral=1"), std::string::npos) << text;
EXPECT_NE(text.find("auth=none"), std::string::npos) << text;
EXPECT_NE(text.find("encryption=TLSv1.3/"), std::string::npos) << text;
::unlink(path.c_str());
}
TEST(Session, DealerKeyAuthenticatesTheDealer)
{
auto dk = std::make_shared<identity>(identity::generate());
dpf::net::session_options dopt = with_limits({});
dopt.security.self = dk;
dpf::net::session_options popt = with_limits({});
popt.security.trusted[dpf::net::dealer_id] = dk->key();
std::atomic<unsigned short> port{0};
std::string derr;
dpf::net::link_security dealer_view;
std::thread dealer([&] {
try
{
asio::io_context io;
dpf::net::dealer_session d(io, 1, dopt);
port.store(d.listen());
d.accept_parties();
dealer_view = d.party_security(0);
}
catch (const std::exception & e)
{
derr = e.what();
}
});
while (port.load() == 0)
std::this_thread::yield();
asio::io_context io;
dpf::net::party_session p(io, 0, 2, popt);
p.connect_dealer("127.0.0.1", port.load());
dealer.join();
ASSERT_TRUE(derr.empty()) << derr;
EXPECT_EQ(p.dealer_security().peer_auth, "key");
EXPECT_EQ(dealer_view.peer_auth, "none");
EXPECT_TRUE(dealer_view.peer_verified_us);
}
// ---------------------------------------------------------------------------
// Client links
// ---------------------------------------------------------------------------
namespace
{
/// A client connects to a `client_listener` and sends one word to it.
std::string client_roundtrip(const dpf::net::server_security & ss,
const dpf::net::client_security & cs, dpf::net::link_security * client_view = nullptr)
{
std::atomic<unsigned short> port{0};
std::string serr, cerr;
std::uint64_t got = 0;
std::thread server([&] {
try
{
asio::io_context io;
dpf::net::deadlines lim;
lim.accept = std::chrono::milliseconds(3000);
dpf::net::client_listener l(io, ss, 1, {}, lim);
port.store(l.listen());
auto c = l.accept();
bool done = false;
c.link->async_read(0, &got, 8, [&](const std::error_code &) { done = true; });
const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(3);
run_until(io, done, deadline);
}
catch (const std::exception & e)
{
serr = e.what();
}
});
while (port.load() == 0)
std::this_thread::yield();
try
{
asio::io_context io;
auto c = dpf::net::connect_server(io, "127.0.0.1", port.load(), cs);
if (client_view != nullptr)
*client_view = c.security;
std::uint64_t v = 77;
bool done = false;
c.link->async_write(0, &v, 8, [&](const std::error_code &) { done = true; });
const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(3);
run_until(io, done, deadline);
c.link.reset();
for (int i = 0; i < 20; ++i)
io.poll();
}
catch (const std::exception & e)
{
cerr = e.what();
}
server.join();
if (!cerr.empty())
return "client: " + cerr;
if (!serr.empty())
return "server: " + serr;
return got == 77 ? "" : "server got the wrong word";
}
} // namespace
TEST(ClientLink, DevelopmentDefaultWorksAndSaysSo)
{
dpf::net::link_security view;
EXPECT_EQ(client_roundtrip({}, {}, &view), "");
EXPECT_EQ(view.peer_auth, "development");
}
TEST(ClientLink, ServerIdentityNeedsAPinOrVerifyOff)
{
dpf::net::server_security ss;
ss.self = std::make_shared<identity>(identity::generate());
EXPECT_NE(client_roundtrip(ss, {}).find("does not trust"), std::string::npos);
dpf::net::client_security pinned;
pinned.pins.push_back(ss.self->key());
dpf::net::link_security view;
EXPECT_EQ(client_roundtrip(ss, pinned, &view), "");
EXPECT_EQ(view.peer_auth, "key");
dpf::net::client_security off;
off.verify = false;
EXPECT_EQ(client_roundtrip(ss, off, &view), "");
EXPECT_EQ(view.peer_auth, "none");
}
// ---------------------------------------------------------------------------
// Configuration
// ---------------------------------------------------------------------------
TEST(Config, SecurityFromAFileWithRelativePaths)
{
const std::string dir = temp_path("cfgdir");
(void)::system(("rm -rf '" + dir + "' && mkdir -p '" + dir + "'").c_str());
const auto me = identity::generate();
const auto other = identity::generate();
me.save(dir + "/p0.key");
{
std::ofstream pub(dir + "/p1.pub");
pub << other.key().base64() << "\n";
std::ofstream cfg(dir + "/p0.conf");
cfg << "# party 0\n"
<< "identity = p0.key\n"
<< "peer.1 = file:p1.pub # the other party\n"
<< "transport = mux\n"
<< "client_pin = " << other.key().base64() << "\n";
}
dpf::app::run_config c;
const std::string arg = "--config=" + dir + "/p0.conf";
const char * argv[] = {"x", arg.c_str()};
EXPECT_TRUE(c.apply_args(2, const_cast<char **>(argv)).empty());
ASSERT_TRUE(c.security.self);
EXPECT_EQ(c.security.self->key(), me.key());
ASSERT_NE(c.security.trusted_key(1), nullptr);
EXPECT_EQ(*c.security.trusted_key(1), other.key());
EXPECT_EQ(c.kind, dpf::net::transport::mux);
ASSERT_EQ(c.client.pins.size(), 1u);
const auto summary = c.summary();
EXPECT_NE(summary.find("encryption=on"), std::string::npos) << summary;
EXPECT_NE(summary.find("identity=" + me.key().base64()), std::string::npos);
EXPECT_NE(summary.find("trusted=1"), std::string::npos);
{
std::ofstream bad(dir + "/bad.conf");
bad << "transport = mux\npeer.1 = not-a-key\n";
}
try
{
c.set("config", dir + "/bad.conf");
FAIL() << "a bad key must be rejected";
}
catch (const std::invalid_argument & e)
{
EXPECT_NE(std::string(e.what()).find("bad.conf:2"), std::string::npos) << e.what();
}
EXPECT_THROW(c.set("encryption", "maybe"), std::invalid_argument);
(void)::system(("rm -rf '" + dir + "'").c_str());
}
// ---------------------------------------------------------------------------
// Separate processes from key files and config files
// ---------------------------------------------------------------------------
namespace
{
dpf::protocol::plan open_plan(std::size_t party, dpf::protocol::node & x,
dpf::protocol::node & out)
{
dpf::protocol::composer c(party);
x = c.input(dpf::protocol::domain::a, 8);
out = c.exchange(x);
return c.schedule();
}
} // namespace
TEST(Deployment, TwoProcessesAuthenticateEachOther)
{
const std::string dir = temp_path("deploy");
(void)::system(("rm -rf '" + dir + "' && mkdir -p '" + dir + "'").c_str());
const identity ids[2] = {identity::generate(), identity::generate()};
unsigned short ports[2];
for (int i = 0; i < 2; ++i)
{
ids[i].save(dir + "/p" + std::to_string(i) + ".key");
asio::io_context io;
asio::ip::tcp::acceptor a(io);
dpf::net::open_listener(a, 0);
ports[i] = a.local_endpoint().port();
}
for (int i = 0; i < 2; ++i)
{
std::ofstream cfg(dir + "/p" + std::to_string(i) + ".conf");
cfg << "identity = p" << i << ".key\n"
<< "peer." << (1 - i) << " = " << ids[1 - i].key().base64() << "\n"
<< "transport = mux\n";
}
const std::string peers = "--peers=127.0.0.1:" + std::to_string(ports[0])
+ ",127.0.0.1:" + std::to_string(ports[1]);
pid_t kids[2];
for (int me = 0; me < 2; ++me)
{
const pid_t pid = ::fork();
ASSERT_GE(pid, 0);
if (pid == 0)
{
int code = 1;
try
{
const std::string party = "--party=" + std::to_string(me);
const std::string cfg = "--config=" + dir + "/p" + std::to_string(me)
+ ".conf";
const std::string log = "--log=file:" + dir + "/p" + std::to_string(me)
+ ".log";
const char * argv[] = {"node", party.c_str(), peers.c_str(), cfg.c_str(),
log.c_str()};
auto args = dpf::app::parse_node_args(5, const_cast<char **>(argv),
dpf::app::run_config{});
dpf::log::settings ls;
ls.sinks = args.cfg.log_sinks;
dpf::log::configure(ls);
dpf::protocol::node x, out;
auto plan = open_plan(static_cast<std::size_t>(me), x, out);
dpf::app::party_values v(plan.nodes().size());
v[x.id].assign(8, 0);
const std::uint64_t in = me == 0 ? 40 : 2;
std::memcpy(v[x.id].data(), &in, 8);
(void)dpf::app::run_node(args, plan, v, {});
std::uint64_t o = 0;
std::memcpy(&o, v[out.id].data(), 8);
code = o == 42 ? 0 : 2;
}
catch (const std::exception & e)
{
std::fprintf(stderr, "party %d: %s\n", me, e.what());
code = 3;
}
::_exit(code);
}
kids[me] = pid;
}
for (int me = 0; me < 2; ++me)
{
int status = 0;
ASSERT_EQ(::waitpid(kids[me], &status, 0), kids[me]);
ASSERT_TRUE(WIFEXITED(status));
EXPECT_EQ(WEXITSTATUS(status), 0) << "party " << me;
}
for (int me = 0; me < 2; ++me)
{
std::ifstream in(dir + "/p" + std::to_string(me) + ".log");
const std::string text((std::istreambuf_iterator<char>(in)),
std::istreambuf_iterator<char>());
EXPECT_NE(text.find("ev=link.up"), std::string::npos) << text;
EXPECT_NE(text.find("auth=key"), std::string::npos) << text;
EXPECT_NE(text.find("peer_verified_us=1"), std::string::npos) << text;
EXPECT_EQ(text.find("ev=security.unauthenticated"), std::string::npos) << text;
}
(void)::system(("rm -rf '" + dir + "'").c_str());
}