ntptun/src/Server.cpp
2026-07-24 21:05:01 +02:00

494 lines
17 KiB
C++

#include "ntptun/Server.hpp"
#include <algorithm>
#include <array>
#include <cerrno>
#include <cstddef>
#include <stdexcept>
#include <utility>
#include <vector>
#include "ntptun/Logging.hpp"
#include "ntptun/Platform.hpp"
namespace ntptun {
namespace {
constexpr std::size_t kBufSize = 65536;
constexpr std::size_t kMaxPendingRelays = 4096;
constexpr std::size_t kMaxRateEntries = 8192;
constexpr std::size_t kMaxRoutes = 100000;
constexpr std::size_t kMaxRelaySockets = 4096;
constexpr int kPollTimeoutMs = 200;
constexpr std::size_t kNoIndex = static_cast<std::size_t>(-1);
bool plausible_request(const NtpHeader& header) {
const std::uint8_t version = header.version();
return (version == 3 || version == 4) && header.mode() == kNtpModeClient &&
header.transmit_ts != 0;
}
} // namespace
Server::Server(ServerConfig config)
: cfg_(std::move(config)), sock_(UdpSocket::bind(cfg_.listen)) {
if (cfg_.common.transport == Transport::Tun) {
#ifdef NTPTUN_HAVE_TUN
tun_ = TunDevice::open(cfg_.common.tun);
#else
throw std::runtime_error("tun transport is not supported on this platform");
#endif
}
if (cfg_.upstream) {
UdpSocket up = UdpSocket::create(cfg_.upstream->family());
up.connect(*cfg_.upstream);
upstream_ = std::move(up);
}
}
void Server::run(const std::atomic<bool>& stop) {
const bool tun_xport = cfg_.common.transport == Transport::Tun;
#ifdef NTPTUN_HAVE_TUN
if (tun_xport) {
tun_.set_nonblocking(true);
}
#endif
sock_.set_nonblocking(true);
if (upstream_) {
upstream_->set_nonblocking(true);
}
const char* xport = tun_xport ? "tun" : "udp";
const std::string inner_desc =
tun_xport ? ("tun=" + cfg_.common.tun)
: ("target=" + cfg_.udp_target.to_string() + " base_port=" +
std::to_string(cfg_.udp_base_port));
log_info("server up: transport=", xport, " ", inner_desc,
" listen=", cfg_.listen.to_string(), upstream_ ? " relay=on" : " relay=off",
cfg_.common.push_mode ? " mode=push" : " mode=poll");
if (!tun_xport) {
// In UDP transport each client relays from source port
// udp_base_port + client_id, so the largest usable client_id is
// (65535 - udp_base_port). Clients above that cannot be served and are
// rejected per-request in relay_for().
const unsigned max_client_id = 65535u - cfg_.udp_base_port;
log_info("udp relay: source ports ", cfg_.udp_base_port, "..65535; ",
"usable client_id range 1..", max_client_id,
" (udp_base_port + client_id must be <= 65535)");
if (max_client_id < 65535u) {
log_warn("udp relay: client_id > ", max_client_id,
" cannot be served with udp_base_port=", cfg_.udp_base_port);
}
}
while (!stop.load()) {
std::vector<PollFd> fds;
std::size_t up_idx = kNoIndex;
fds.push_back(PollFd{sock_.fd(), kPollIn, 0});
const std::size_t udp_idx = 0;
#ifdef NTPTUN_HAVE_TUN
std::size_t tun_idx = kNoIndex;
if (tun_xport) {
tun_idx = fds.size();
fds.push_back(PollFd{tun_.fd(), kPollIn, 0});
}
#endif
if (upstream_) {
up_idx = fds.size();
fds.push_back(PollFd{upstream_->fd(), kPollIn, 0});
}
std::vector<std::uint16_t> relay_ids;
if (!tun_xport) {
relay_ids.reserve(relay_socks_.size());
for (auto& entry : relay_socks_) {
relay_ids.push_back(entry.first);
fds.push_back(PollFd{entry.second.fd(), kPollIn, 0});
}
}
const std::size_t relay_base = fds.size() - relay_ids.size();
const int rc = net_poll(fds.data(), fds.size(), kPollTimeoutMs);
if (rc < 0) {
#if !defined(_WIN32)
if (errno == EINTR) {
continue;
}
#endif
socket_throw("poll");
}
if (rc > 0) {
if (fds[udp_idx].revents & kPollIn) {
drain_udp();
}
#ifdef NTPTUN_HAVE_TUN
if (tun_idx != kNoIndex && (fds[tun_idx].revents & kPollIn)) {
drain_tun();
}
#endif
if (up_idx != kNoIndex && (fds[up_idx].revents & kPollIn)) {
drain_upstream();
}
for (std::size_t i = 0; i < relay_ids.size(); ++i) {
if (fds[relay_base + i].revents & kPollIn) {
const auto it = relay_socks_.find(relay_ids[i]);
if (it != relay_socks_.end()) {
drain_relay(relay_ids[i], it->second);
}
}
}
}
expire_relays();
}
log_info("server shutting down");
}
void Server::drain_udp() {
std::array<Byte, kBufSize> buf{};
Endpoint from;
for (;;) {
const ssize_t n = sock_.recv_from(buf.data(), buf.size(), from);
if (n < 0) {
break;
}
process_request(ByteSpan(buf.data(), static_cast<std::size_t>(n)), from);
}
}
void Server::process_request(ByteSpan packet, const Endpoint& from) {
const auto header = parse_ntp_header(packet);
if (!header) {
return; // not NTP -> drop
}
if (header->mode() != kNtpModeClient) {
return; // not a client request -> drop
}
// Attempt tunnel classification before touching any state.
std::vector<NtpExtension> exts;
if (parse_ntp_extensions(packet, exts)) {
const NtpExtension* tunnel_ext = nullptr;
int matches = 0;
for (const auto& ext : exts) {
if (ext.type == cfg_.common.ext_type) {
tunnel_ext = &ext;
++matches;
}
}
if (matches == 1) {
DecodedCarrier decoded;
const CarrierStatus status = decode_carrier_value(
tunnel_ext->value, cfg_.common.keys, Direction::ClientToServer,
header->transmit_ts, decoded);
if (status == CarrierStatus::Ok) {
// Reject replayed or too-old carriers before mutating any state.
ClientState& state = clients_[decoded.client_id];
if (!state.replay.accept(header->transmit_ts,
cfg_.common.replay_window)) {
log_debug("replay: dropping carrier for client ",
decoded.client_id, " ts=", header->transmit_ts);
return; // authenticated but replayed -> drop, do not relay
}
handle_tunnel(header->transmit_ts, decoded, from);
return;
}
}
}
// Not a valid carrier: fall through to the real-NTP relay path when the
// outer packet is independently a plausible NTP client request.
if (plausible_request(*header)) {
relay(packet, *header, from);
}
}
void Server::handle_tunnel(std::uint64_t request_transmit_ts,
const DecodedCarrier& carrier, const Endpoint& from) {
ClientState& state = clients_[carrier.client_id];
state.last_src = from; // remember the validated source
state.has_src = true;
state.last_request_ts = request_transmit_ts;
if (cfg_.common.transport == Transport::Udp) {
if (carrier.kind == kKindUdpDatagram && !carrier.inner_ip.empty()) {
UdpSocket* r = relay_for(carrier.client_id);
if (r && r->send(carrier.inner_ip.data(), carrier.inner_ip.size()) < 0) {
log_debug("relay send would block; payload dropped");
}
}
} else {
#ifdef NTPTUN_HAVE_TUN
if (carrier.kind == kKindIpDatagram) {
if (tun_.write_packet(carrier.inner_ip.data(), carrier.inner_ip.size()) < 0) {
log_debug("tun write would block; inner datagram dropped");
}
// Source-learn a route back to this client for downstream traffic.
const auto src = ip::source(carrier.inner_ip);
if (src) {
if (routes_.size() >= kMaxRoutes && routes_.find(*src) == routes_.end()) {
log_warn("route table full; not learning new route");
} else {
routes_[*src] = carrier.client_id;
}
}
}
#endif
}
// Answer with exactly one response, carrying a queued datagram if available.
std::uint8_t kind = kKindEmptyPoll;
ByteVector payload;
if (!state.downstream.empty()) {
kind = (cfg_.common.transport == Transport::Udp) ? kKindUdpDatagram
: kKindIpDatagram;
payload = std::move(state.downstream.front());
state.downstream.pop_front();
}
send_response(carrier.client_id, request_transmit_ts, kind, payload, from);
}
UdpSocket* Server::relay_for(std::uint16_t client_id) {
const auto existing = relay_socks_.find(client_id);
if (existing != relay_socks_.end()) {
relay_last_[client_id] = std::chrono::steady_clock::now();
return &existing->second;
}
const std::uint32_t port =
static_cast<std::uint32_t>(cfg_.udp_base_port) + client_id;
if (port > 65535u) {
log_error("relay: source port ", port, " out of range for client ", client_id);
return nullptr;
}
// Evict the least-recently-used relay socket when the table is full.
if (relay_socks_.size() >= kMaxRelaySockets && !relay_last_.empty()) {
const auto lru = std::min_element(
relay_last_.begin(), relay_last_.end(),
[](const auto& a, const auto& b) { return a.second < b.second; });
const std::uint16_t victim = lru->first;
relay_socks_.erase(victim);
relay_last_.erase(victim);
log_debug("relay: evicted LRU socket for client ", victim);
}
try {
UdpSocket s = UdpSocket::bind(
Endpoint::wildcard(cfg_.udp_target.family(), static_cast<std::uint16_t>(port)));
s.connect(cfg_.udp_target);
s.set_nonblocking(true);
const auto inserted = relay_socks_.emplace(client_id, std::move(s));
relay_last_[client_id] = std::chrono::steady_clock::now();
return &inserted.first->second;
} catch (const std::exception& e) {
log_error("relay: cannot open socket for client ", client_id, " on port ",
port, ": ", e.what());
return nullptr;
}
}
void Server::drain_relay(std::uint16_t client_id, UdpSocket& sock) {
std::array<Byte, kBufSize> buf{};
for (;;) {
const ssize_t n = sock.recv(buf.data(), buf.size());
if (n < 0) {
break;
}
relay_last_[client_id] = std::chrono::steady_clock::now();
if (static_cast<std::size_t>(n) > cfg_.common.inner_mtu) {
log_error("relay: dropping oversized reply of ", n, " bytes for client ",
client_id, " (max payload is ", cfg_.common.inner_mtu,
"); raise inner_mtu (and ntp_payload_cap >= inner_mtu+80), "
"e.g. inner_mtu >= 1200 for QUIC/Hysteria2");
continue;
}
queue_downstream(client_id, kKindUdpDatagram,
ByteSpan(buf.data(), static_cast<std::size_t>(n)));
}
}
void Server::queue_downstream(std::uint16_t client_id, std::uint8_t kind,
ByteSpan payload) {
ClientState& state = clients_[client_id];
if (cfg_.common.push_mode && state.has_src) {
// Push immediately to the client's last known address instead of waiting
// for a poll (breaks strict NTP request/response semantics).
send_response(client_id, state.last_request_ts, kind, payload, state.last_src);
return;
}
state.downstream.emplace_back(payload.begin(), payload.end());
while (state.downstream.size() > cfg_.max_queue_per_client) {
state.downstream.pop_front();
}
}
void Server::send_response(std::uint16_t client_id, std::uint64_t request_transmit_ts,
std::uint8_t kind, ByteSpan payload, const Endpoint& from) {
const auto key = cfg_.common.keys.resolve(client_id);
if (!key) {
return; // unreachable: decode already resolved this key
}
const std::uint64_t receive_ts = ntp_now();
const std::uint64_t transmit_ts = next_transmit_ts(); // monotonic, >= receive_ts
NtpHeader header;
header.set_li_vn_mode(0, 4, kNtpModeServer);
header.stratum = 2;
header.poll = 4;
header.precision = -20;
header.reference_ts = receive_ts;
header.origin_ts = request_transmit_ts; // NTP correlation
header.receive_ts = receive_ts;
header.transmit_ts = transmit_ts;
const ByteVector value = encode_carrier_value(
client_id, kind, payload, *key, Direction::ServerToClient, transmit_ts);
const ByteVector packet =
build_packet_with_extension(header, cfg_.common.ext_type, value);
if (packet.size() > cfg_.common.ntp_payload_cap) {
log_warn("response exceeds NTP payload cap; dropping");
return;
}
sock_.send_to(packet.data(), packet.size(), from);
}
void Server::drain_tun() {
#ifdef NTPTUN_HAVE_TUN
std::array<Byte, kBufSize> buf{};
for (;;) {
const ssize_t n = tun_.read_packet(buf.data(), buf.size());
if (n < 0) {
break;
}
if (n == 0) {
continue;
}
const ByteSpan pkt(buf.data(), static_cast<std::size_t>(n));
const auto ver = ip::version(pkt);
if (!ver || (*ver != 4 && *ver != 6)) {
continue;
}
const auto total = ip::total_length(pkt);
if (!total || *total != static_cast<std::size_t>(n)) {
continue;
}
if (static_cast<std::size_t>(n) > cfg_.common.inner_mtu) {
log_warn("tun: dropping downstream datagram of ", n,
" bytes (inner MTU is ", cfg_.common.inner_mtu, ")");
continue;
}
const auto dst = ip::destination(pkt);
if (!dst) {
continue;
}
const auto route = routes_.find(*dst);
if (route == routes_.end()) {
log_debug("no route for inner destination; dropping downstream datagram");
continue;
}
queue_downstream(route->second, kKindIpDatagram, pkt);
}
#endif
}
void Server::relay(ByteSpan packet, const NtpHeader& header, const Endpoint& from) {
if (!upstream_) {
return; // relay disabled -> silently drop non-tunnel requests
}
if (!relay_rate_ok(from)) {
log_debug("relay: rate limited ", from.to_string());
return;
}
if (upstream_->send(packet.data(), packet.size()) < 0) {
return;
}
PendingRelay pending;
pending.requester = from;
pending.origin_ts = header.transmit_ts;
pending.request_size = packet.size();
pending.deadline = std::chrono::steady_clock::now() +
std::chrono::milliseconds(cfg_.relay_timeout_ms);
relays_.push_back(std::move(pending));
if (relays_.size() > kMaxPendingRelays) {
relays_.pop_front();
}
}
void Server::drain_upstream() {
std::array<Byte, kBufSize> buf{};
for (;;) {
const ssize_t n = upstream_->recv(buf.data(), buf.size());
if (n < 0) {
break;
}
const auto header = parse_ntp_header(ByteSpan(buf.data(), static_cast<std::size_t>(n)));
if (!header || header->mode() != kNtpModeServer) {
continue; // only valid mode-4 responses are forwarded
}
// Correlate against a pending relay by Origin == forwarded Transmit.
const auto it = std::find_if(
relays_.begin(), relays_.end(),
[&](const PendingRelay& p) { return p.origin_ts == header->origin_ts; });
if (it == relays_.end()) {
continue;
}
const PendingRelay pending = *it;
relays_.erase(it);
// Amplification guard: never return more bytes than the request. If the
// upstream reply is larger, return only the validated 48-byte base.
std::size_t out_len = static_cast<std::size_t>(n);
if (out_len > pending.request_size) {
out_len = kNtpHeaderSize;
}
sock_.send_to(buf.data(), out_len, pending.requester);
}
}
bool Server::relay_rate_ok(const Endpoint& from) {
const auto now = std::chrono::steady_clock::now();
const ip::Address key = from.ip_key();
const double capacity = static_cast<double>(cfg_.relay_rate_per_sec);
if (rate_.size() > kMaxRateEntries && rate_.find(key) == rate_.end()) {
rate_.clear(); // coarse bound on memory for the rate table
}
RateBucket& bucket = rate_[key];
if (bucket.last.time_since_epoch().count() == 0) {
bucket.tokens = capacity;
bucket.last = now;
}
const double elapsed = std::chrono::duration<double>(now - bucket.last).count();
bucket.last = now;
bucket.tokens = std::min(capacity, bucket.tokens + elapsed * capacity);
if (bucket.tokens < 1.0) {
return false;
}
bucket.tokens -= 1.0;
return true;
}
void Server::expire_relays() {
const auto now = std::chrono::steady_clock::now();
while (!relays_.empty() && relays_.front().deadline <= now) {
relays_.pop_front();
}
}
std::uint64_t Server::next_transmit_ts() {
std::uint64_t ts = ntp_now();
if (ts <= last_transmit_ts_) {
ts = last_transmit_ts_ + 1;
}
last_transmit_ts_ = ts;
return ts;
}
} // namespace ntptun