Merge pull request #114834 from Faless/net/enet/dtls_memory_management

ENet: Fix DTLS server leaking peers
This commit is contained in:
Thaddeus Crews 2026-01-26 13:14:16 -06:00
commit 013fffa10d
No known key found for this signature in database
GPG key ID: 8C6E5FEB5FC03CCC
3 changed files with 111 additions and 18 deletions

View file

@ -182,6 +182,9 @@ int ENetPacketPeer::get_packet_flags() const {
void ENetPacketPeer::_on_disconnect() {
if (peer) {
#ifdef GODOT_ENET
enet_peer_socket_destroy(peer);
#endif
peer->data = nullptr;
}
peer = nullptr;
@ -256,6 +259,9 @@ void ENetPacketPeer::_bind_methods() {
ENetPacketPeer::ENetPacketPeer(ENetPeer *p_peer) {
peer = p_peer;
peer->data = this;
#ifdef GODOT_ENET
enet_peer_socket_bind(peer);
#endif
}
ENetPacketPeer::~ENetPacketPeer() {

View file

@ -49,5 +49,7 @@ ENET_API void enet_address_set_ip(ENetAddress * address, const uint8_t * ip, siz
ENET_API int enet_host_dtls_server_setup (ENetHost *, void *);
ENET_API int enet_host_dtls_client_setup (ENetHost *, const char *, void *);
ENET_API void enet_host_refuse_new_connections (ENetHost *, int);
ENET_API void enet_peer_socket_bind (ENetPeer *);
ENET_API void enet_peer_socket_destroy (ENetPeer *);
#endif // __ENET_GODOT_EXT_H__

View file

@ -54,6 +54,8 @@ public:
virtual void close() = 0;
virtual void set_refuse_new_connections(bool p_enable) {} /* Only used by dtls server */
virtual bool can_upgrade() { return false; } /* Only true in ENetUDP */
virtual void disconnect_peer(ENetPeer *p_peer) {}
virtual void bind_peer(ENetPeer *p_peer) {}
virtual ~ENetGodotSocket() {}
};
@ -171,6 +173,8 @@ class ENetDTLSClient : public ENetGodotSocket {
Ref<TLSOptions> tls_options;
String for_hostname;
IPAddress local_address;
uint16_t local_port = 0;
bool bound = false;
public:
ENetDTLSClient(ENetUDP *p_base, String p_for_hostname, Ref<TLSOptions> p_options) {
@ -179,10 +183,9 @@ public:
udp.instantiate();
dtls = Ref<PacketPeerDTLS>(PacketPeerDTLS::create());
if (p_base->bound) {
uint16_t port;
p_base->get_socket_address(&local_address, &port);
p_base->get_socket_address(&local_address, &local_port);
p_base->close();
bind(local_address, port);
bind(local_address, local_port);
}
}
@ -191,7 +194,9 @@ public:
}
Error bind(IPAddress p_ip, uint16_t p_port) {
bound = true;
local_address = p_ip;
local_port = p_port;
return udp->bind(p_port, p_ip);
}
@ -208,7 +213,11 @@ public:
if (!connected) {
udp->connect_to_host(p_ip, p_port);
if (dtls->connect_to_peer(udp, for_hostname, tls_options)) {
return FAILED;
close();
if (bound) {
bind(local_address, local_port);
}
return FAILED; // May still retry...
}
connected = true;
}
@ -216,7 +225,11 @@ public:
if (dtls->get_status() == PacketPeerDTLS::STATUS_HANDSHAKING) {
return ERR_BUSY;
} else if (dtls->get_status() != PacketPeerDTLS::STATUS_CONNECTED) {
return FAILED;
close();
if (bound) {
bind(local_address, local_port);
}
return FAILED; // May still retry...
}
r_sent = p_len;
return dtls->put_packet(p_buffer, p_len);
@ -228,6 +241,10 @@ public:
return ERR_BUSY;
}
if (dtls->get_status() != PacketPeerDTLS::STATUS_CONNECTED) {
close();
if (bound) {
bind(local_address, local_port);
}
return FAILED;
}
int pc = dtls->get_available_packet_count();
@ -255,6 +272,7 @@ public:
void close() {
dtls->disconnect_from_peer();
udp->close();
connected = false;
}
};
@ -262,12 +280,24 @@ public:
class ENetDTLSServer : public ENetGodotSocket {
Ref<DTLSServer> server;
Ref<UDPServer> udp_server;
HashMap<String, Ref<PacketPeerDTLS>> peers;
int last_service = 0;
enum {
ENET_DTLS_TIMEOUT_MS = 8000,
};
struct Peer {
Ref<PacketPeerDTLS> conn;
bool bound = false;
uint64_t timeout = 0;
};
HashMap<String, Peer> peers;
IPAddress local_address;
size_t max_clients = 0;
public:
ENetDTLSServer(ENetUDP *p_base, Ref<TLSOptions> p_options) {
ENetDTLSServer(ENetUDP *p_base, Ref<TLSOptions> p_options, size_t p_max_clients) {
max_clients = p_max_clients;
udp_server.instantiate();
if (p_base->bound) {
uint16_t port;
@ -303,14 +333,14 @@ public:
Error sendto(const uint8_t *p_buffer, int p_len, int &r_sent, IPAddress p_ip, uint16_t p_port) {
String key = String(p_ip) + ":" + itos(p_port);
if (unlikely(!peers.has(key))) {
Peer *peer = peers.getptr(key);
if (unlikely(peer == nullptr)) {
// The peer might have been disconnected due to a DTLS error.
// We need to wait for it to time out, just mark the packet as sent.
r_sent = p_len;
return OK;
}
Ref<PacketPeerDTLS> peer = peers[key];
Error err = peer->put_packet(p_buffer, p_len);
Error err = peer->conn->put_packet(p_buffer, p_len);
if (err == OK) {
r_sent = p_len;
} else if (err == ERR_BUSY) {
@ -333,9 +363,12 @@ public:
int peer_port = udp->get_packet_port();
Ref<PacketPeerDTLS> peer = server->take_connection(udp);
PacketPeerDTLS::Status status = peer->get_status();
if (status == PacketPeerDTLS::STATUS_HANDSHAKING || status == PacketPeerDTLS::STATUS_CONNECTED) {
if (peers.size() < max_clients && status == PacketPeerDTLS::STATUS_HANDSHAKING || status == PacketPeerDTLS::STATUS_CONNECTED) {
String key = String(peer_ip) + ":" + itos(peer_port);
peers[key] = peer;
Peer p;
p.conn = peer;
p.timeout = OS::get_singleton()->get_ticks_msec() + ENET_DTLS_TIMEOUT_MS;
peers[key] = p;
}
}
@ -343,8 +376,12 @@ public:
Error err = ERR_BUSY;
// TODO this needs to be fair!
for (KeyValue<String, Ref<PacketPeerDTLS>> &E : peers) {
Ref<PacketPeerDTLS> peer = E.value;
for (KeyValue<String, Peer> &E : peers) {
if (unlikely(!E.value.bound && OS::get_singleton()->get_ticks_msec() >= E.value.timeout)) {
remove.push_back(E.key);
continue;
}
Ref<PacketPeerDTLS> peer = E.value.conn;
peer->poll();
if (peer->get_status() == PacketPeerDTLS::STATUS_HANDSHAKING) {
@ -388,9 +425,35 @@ public:
return -1;
}
void bind_peer(ENetPeer *p_peer) {
if (p_peer == nullptr) {
return;
}
IPAddress ip;
ip.set_ipv6((uint8_t *)&(p_peer->address.host));
int port = p_peer->address.port;
const String key = String(ip) + ":" + itos(port);
Peer *p = peers.getptr(key);
if (p == nullptr) {
return;
}
p->bound = true;
}
void disconnect_peer(ENetPeer *p_peer) {
if (p_peer == nullptr) {
return;
}
IPAddress ip;
ip.set_ipv6((uint8_t *)&(p_peer->address.host));
int port = p_peer->address.port;
const String key = String(ip) + ":" + itos(port);
peers.erase(key);
};
void close() {
for (KeyValue<String, Ref<PacketPeerDTLS>> &E : peers) {
E.value->disconnect_from_peer();
for (KeyValue<String, Peer> &E : peers) {
E.value.conn->disconnect_from_peer();
}
peers.clear();
udp_server->stop();
@ -448,13 +511,35 @@ ENetSocket enet_socket_create(ENetSocketType type) {
return socket;
}
void enet_peer_socket_bind(ENetPeer *p_peer) {
if (p_peer == nullptr || p_peer->host == nullptr) {
return;
}
ENetGodotSocket *sock = (ENetGodotSocket *)p_peer->host->socket;
if (sock == nullptr) {
return;
}
sock->bind_peer(p_peer);
}
void enet_peer_socket_destroy(ENetPeer *p_peer) {
if (p_peer == nullptr || p_peer->host == nullptr) {
return;
}
ENetGodotSocket *sock = (ENetGodotSocket *)p_peer->host->socket;
if (sock == nullptr) {
return;
}
sock->disconnect_peer(p_peer);
}
int enet_host_dtls_server_setup(ENetHost *host, void *p_options) {
ERR_FAIL_COND_V_MSG(!DTLSServer::is_available(), -1, "DTLS server is not available in this build.");
ENetGodotSocket *sock = (ENetGodotSocket *)host->socket;
if (!sock->can_upgrade()) {
return -1;
}
host->socket = memnew(ENetDTLSServer(static_cast<ENetUDP *>(sock), Ref<TLSOptions>(static_cast<TLSOptions *>(p_options))));
host->socket = memnew(ENetDTLSServer(static_cast<ENetUDP *>(sock), Ref<TLSOptions>(static_cast<TLSOptions *>(p_options)), host->peerCount));
memdelete(sock);
return 0;
}