mirror of
https://gitlab.com/shadow_contributor/ntptun.git
synced 2026-10-02 19:36:45 +00:00
494 lines
17 KiB
C++
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
|