[socket] Share the server side as TcpListener (#20037)

Co-authored-by: J. Nick Koston <nick@home-assistant.io>
This commit is contained in:
Bascht74
2026-10-02 13:07:00 -05:00
committed by GitHub
co-authored by J. Nick Koston
parent cbe89c940d
commit b6c79fdc9a
13 changed files with 262 additions and 70 deletions
+20
View File
@@ -5,6 +5,7 @@ from ipaddress import IPv4Address, IPv4Network
import logging
import esphome.codegen as cg
from esphome.components.const import CONF_ROLE
from esphome.config_helpers import filter_source_files_from_defines
import esphome.config_validation as cv
from esphome.core import CORE, ID
@@ -162,6 +163,7 @@ def add_ipv4_allow(
"""
if not networks:
return
cg.add_define("USE_SOCKET_IPV4_ALLOW")
entries = [
cg.StructInitializer(
Ipv4AllowEntry,
@@ -186,6 +188,23 @@ def require_tcp_client_link() -> None:
cg.add_define("USE_SOCKET_TCP_CLIENT_LINK")
def require_tcp_listener() -> None:
"""Compile the TCP listener; call from a server role's to_code."""
require_tcp_client_link()
cg.add_define("USE_SOCKET_TCP_LISTENER")
def consume_role_sockets(component: str) -> Callable[[ConfigType], ConfigType]:
"""Socket accounting for a role keyed client or server schema."""
def validator(config: ConfigType) -> ConfigType:
if config[CONF_ROLE] == "server":
consume_sockets(1, component, SocketType.TCP_LISTEN)(config)
return consume_sockets(1, component)(config)
return validator
CONFIG_SCHEMA = cv.Schema(
{
cv.SplitDefault(
@@ -239,5 +258,6 @@ FILTER_SOURCE_FILES = filter_source_files_from_defines(
"lwip_sockets_impl.cpp": "USE_SOCKET_IMPL_LWIP_SOCKETS",
"ipv4_resolve.cpp": "USE_SOCKET_IPV4_RESOLVE",
"tcp_client_link.cpp": "USE_SOCKET_TCP_CLIENT_LINK",
"tcp_listener.cpp": "USE_SOCKET_TCP_LISTENER",
}
)
+10 -3
View File
@@ -39,15 +39,22 @@ class Ipv4Allow {
return true;
}
for (size_t i = 0; i != this->count_; i++) {
Ipv4AllowEntry entry;
progmem_memcpy(&entry, &this->entries_[i], sizeof(entry));
if ((addr & entry.mask) == entry.addr) {
Ipv4AllowEntry e = this->entry(i);
if ((addr & e.mask) == e.addr) {
return true;
}
}
return false;
}
size_t size() const { return this->count_; }
/// A copy of entry i, read from flash.
Ipv4AllowEntry entry(size_t i) const {
Ipv4AllowEntry e;
progmem_memcpy(&e, &this->entries_[i], sizeof(e));
return e;
}
private:
const Ipv4AllowEntry *entries_{nullptr};
size_t count_{0};
@@ -0,0 +1,90 @@
#include "tcp_listener.h"
#if defined(USE_SOCKET_IMPL_LWIP_TCP) || defined(USE_SOCKET_IMPL_LWIP_SOCKETS) || defined(USE_SOCKET_IMPL_BSD_SOCKETS)
#include "esphome/core/application.h"
#include "esphome/core/log.h"
#include <cerrno>
#include <span>
namespace esphome::socket {
// One client at a time; a second connection waits in the stack until the first drops.
static constexpr int LISTEN_BACKLOG = 1;
#ifdef USE_SOCKET_IPV4_ALLOW
static constexpr uint32_t REJECT_LOG_INTERVAL_MS = 5000;
#endif
void TcpListener::try_listen_(TcpClientLink &link) {
this->listen_ = socket_ip_loop_monitored(SOCK_STREAM, IPPROTO_TCP);
int err = errno;
if (this->listen_ != nullptr) {
int yes = 1;
this->listen_->setsockopt(SOL_SOCKET, SO_REUSEADDR, &yes, sizeof(yes));
struct sockaddr_storage local;
socklen_t local_len = set_sockaddr_any(reinterpret_cast<struct sockaddr *>(&local), sizeof(local), link.port());
// A blocking listener would stall loop() inside accept(), so its
// setblocking result is part of the success condition.
if (this->listen_->setblocking(false) == 0 && local_len != 0 &&
this->listen_->bind(reinterpret_cast<struct sockaddr *>(&local), local_len) == 0 &&
this->listen_->listen(LISTEN_BACKLOG) == 0) {
ESP_LOGI(this->tag_, "Listening on %u", link.port());
return;
}
// Captured before reset(); the close inside can overwrite errno.
err = errno;
this->listen_.reset();
}
ESP_LOGW(this->tag_, "Listen on %u failed: %d", link.port(), err);
link.note_attempt();
}
void TcpListener::accept_(TcpClientLink &link) {
struct sockaddr_storage peer {};
socklen_t peer_len = sizeof(peer);
auto client = this->listen_->accept_loop_monitored(reinterpret_cast<struct sockaddr *>(&peer), &peer_len);
if (client == nullptr) {
// A reset during the handshake or a signal only affects that connection.
if (errno == EAGAIN || errno == EWOULDBLOCK || errno == ECONNABORTED || errno == EINTR) {
return;
}
// Rebuild the listener after the backoff instead of spinning on it.
int err = errno;
this->listen_.reset();
ESP_LOGW(this->tag_, "Accept failed: %d", err);
link.note_attempt();
return;
}
const auto *sa = reinterpret_cast<const struct sockaddr *>(&peer);
char text[SOCKADDR_STR_LEN];
format_sockaddr_to(sa, peer_len, std::span<char, SOCKADDR_STR_LEN>(text));
#ifdef USE_SOCKET_IPV4_ALLOW
if (!this->allow_.allows(sa)) {
uint32_t now = App.get_loop_component_start_time();
if (this->last_reject_log_ms_ == 0 || now - this->last_reject_log_ms_ >= REJECT_LOG_INTERVAL_MS) {
this->last_reject_log_ms_ = now;
ESP_LOGW(this->tag_, "Rejected %s", text);
}
return;
}
#endif
link.adopt(std::move(client));
ESP_LOGI(this->tag_, "Client connected from %s", text);
}
void TcpListener::dump_config() const {
#ifdef USE_SOCKET_IPV4_ALLOW
for (size_t i = 0; i < this->allow_.size(); i++) {
Ipv4AllowEntry e = this->allow_.entry(i);
// Network order is dotted order, and the contiguous mask's popcount is the prefix.
const auto *b = reinterpret_cast<const uint8_t *>(&e.addr);
ESP_LOGCONFIG(this->tag_, " Allowed IP: %u.%u.%u.%u/%u", b[0], b[1], b[2], b[3],
static_cast<unsigned>(__builtin_popcount(e.mask)));
}
#endif
}
} // namespace esphome::socket
#endif
+60
View File
@@ -0,0 +1,60 @@
#pragma once
#include "headers.h"
#if defined(USE_SOCKET_IMPL_LWIP_TCP) || defined(USE_SOCKET_IMPL_LWIP_SOCKETS) || defined(USE_SOCKET_IMPL_BSD_SOCKETS)
#ifdef USE_SOCKET_IPV4_ALLOW
#include "ipv4_allow.h"
#endif
#include "socket.h"
#include "tcp_client_link.h"
#include <cstdint>
#include <memory>
namespace esphome::socket {
/// The server side of a bridged TCP link: owns the listen socket and the
/// allow list, accepts one peer at a time and adopts it into a TcpClientLink,
/// sharing that link's retry clock and connect port.
class TcpListener {
public:
#ifdef USE_SOCKET_IPV4_ALLOW
void set_allow(const Ipv4AllowEntry *entries, size_t count) { this->allow_.set(entries, count); }
#endif
/// Call from setup(); tag names the log lines.
void begin(const char *tag) { this->tag_ = tag; }
/// Server state machine; call every loop. may_accept lets the caller hold
/// accepts until its own disconnect edge has run.
void poll(TcpClientLink &link, bool may_accept) {
if (this->listen_ == nullptr) {
if (!link.in_backoff()) {
this->try_listen_(link);
}
return;
}
if (may_accept && !link.connected() && this->listen_->ready()) {
this->accept_(link);
}
}
void close() { this->listen_.reset(); }
/// One config line per allowed network.
void dump_config() const;
protected:
void try_listen_(TcpClientLink &link);
void accept_(TcpClientLink &link);
std::unique_ptr<ListenSocket> listen_;
const char *tag_{nullptr};
#ifdef USE_SOCKET_IPV4_ALLOW
uint32_t last_reject_log_ms_{0};
Ipv4Allow allow_;
#endif
};
} // namespace esphome::socket
#endif
+13 -10
View File
@@ -19,15 +19,10 @@ MULTI_CONF = True
uart_tcp_ns = cg.esphome_ns.namespace("uart_tcp")
UartTcp = uart_tcp_ns.class_("UartTcp", cg.Component, uart.UARTDevice)
CONF_ALLOWED_IPS = "allowed_ips"
CONF_CONNECTED = "connected"
def _consume_sockets(config: ConfigType) -> ConfigType:
if config[CONF_ROLE] == "server":
socket.consume_sockets(1, "uart_tcp", socket.SocketType.TCP_LISTEN)(config)
return socket.consume_sockets(1, "uart_tcp")(config)
BASE_SCHEMA = cv.Schema(
{
cv.GenerateID(): cv.declare_id(UartTcp),
@@ -47,22 +42,30 @@ CONFIG_SCHEMA = cv.All(
cv.typed_schema(
{
"client": BASE_SCHEMA.extend({cv.Required(CONF_HOST): cv.string}),
"server": BASE_SCHEMA,
"server": BASE_SCHEMA.extend(
{cv.Optional(CONF_ALLOWED_IPS): socket.IPV4_ALLOW_SCHEMA}
),
},
key=CONF_ROLE,
default_type="client",
lower=True,
),
_consume_sockets,
socket.consume_role_sockets("uart_tcp"),
)
async def to_code(config: ConfigType) -> None:
socket.require_tcp_client_link()
var = cg.new_Pvariable(config[CONF_ID])
await cg.register_component(var, config)
await uart.register_uart_device(var, config)
cg.add(var.set_server(config[CONF_ROLE] == "server"))
if config[CONF_ROLE] == "server":
socket.require_tcp_listener()
cg.add(var.set_server(True))
socket.add_ipv4_allow(
var.set_allow, config.get(CONF_ALLOWED_IPS), config[CONF_ID]
)
else:
socket.require_tcp_client_link()
cg.add(var.set_port(config[CONF_PORT]))
cg.add(var.set_reconnect_interval(config[CONF_RECONNECT_INTERVAL]))
if (host := config.get(CONF_HOST)) is not None:
+14 -52
View File
@@ -10,13 +10,14 @@ namespace esphome::uart_tcp {
static const char *const TAG = "uart_tcp";
// One client at a time; a second connection waits in the stack until the first drops.
static constexpr int LISTEN_BACKLOG = 1;
// Bytes per 16 ms loop pass at 10 bits per byte: baud / 10 / 62.5.
static constexpr uint32_t BAUD_PACE_DIVISOR = 625;
void UartTcp::setup() {
this->link_.begin(TAG);
#ifdef USE_SOCKET_TCP_LISTENER
this->listener_.begin(TAG);
#endif
if (this->connected_sensor_ != nullptr) {
this->connected_sensor_->publish_state(false);
}
@@ -30,12 +31,17 @@ void UartTcp::dump_config() {
this->server_ ? LOG_STR_LITERAL("Listen") : LOG_STR_LITERAL("Host"),
this->server_ ? LOG_STR_LITERAL("*") : this->link_.host(), this->link_.port(),
this->link_.reconnect_interval());
#ifdef USE_SOCKET_TCP_LISTENER
this->listener_.dump_config();
#endif
LOG_BINARY_SENSOR(" ", "Connected", this->connected_sensor_);
}
void UartTcp::on_shutdown() {
this->link_.close();
this->listen_.reset();
#ifdef USE_SOCKET_TCP_LISTENER
this->listener_.close();
#endif
}
void UartTcp::sync_link_() {
@@ -50,49 +56,6 @@ void UartTcp::sync_link_() {
}
}
void UartTcp::try_listen_() {
this->listen_ = socket::socket_ip_loop_monitored(SOCK_STREAM, IPPROTO_TCP);
int err = errno;
if (this->listen_ != nullptr) {
int yes = 1;
this->listen_->setsockopt(SOL_SOCKET, SO_REUSEADDR, &yes, sizeof(yes));
struct sockaddr_storage local;
socklen_t local_len =
socket::set_sockaddr_any(reinterpret_cast<struct sockaddr *>(&local), sizeof(local), this->link_.port());
// A blocking listener would stall loop() inside accept(), so its
// setblocking result is part of the success condition.
if (this->listen_->setblocking(false) == 0 && local_len != 0 &&
this->listen_->bind(reinterpret_cast<struct sockaddr *>(&local), local_len) == 0 &&
this->listen_->listen(LISTEN_BACKLOG) == 0) {
ESP_LOGI(TAG, "Listening on %u", this->link_.port());
return;
}
// Captured before reset(); the close inside can overwrite errno.
err = errno;
this->listen_.reset();
}
ESP_LOGW(TAG, "Listen on %u failed: %d", this->link_.port(), err);
this->link_.note_attempt();
}
void UartTcp::accept_client_() {
auto client = this->listen_->accept_loop_monitored(nullptr, nullptr);
if (client == nullptr) {
// A reset during the handshake or a signal only affects that connection.
if (errno == EAGAIN || errno == EWOULDBLOCK || errno == ECONNABORTED || errno == EINTR) {
return;
}
// Rebuild the listener after the backoff instead of spinning on it.
int err = errno;
this->listen_.reset();
ESP_LOGW(TAG, "Accept failed: %d", err);
this->link_.note_attempt();
return;
}
this->link_.adopt(std::move(client));
ESP_LOGI(TAG, "Client connected");
}
void UartTcp::read_socket_() {
// A hardware write blocks until the driver takes every byte. Leave what does
// not fit in the socket, so TCP flow control throttles the peer.
@@ -141,18 +104,17 @@ void UartTcp::read_uart_() {
}
void UartTcp::loop() {
#ifdef USE_SOCKET_TCP_LISTENER
if (this->server_) {
if (this->listen_ == nullptr && !this->link_.in_backoff()) {
this->try_listen_();
}
// link_was_up_ holds the accept until the previous drop's edge has run,
// so the sensor and the stale UART discard always see the disconnect.
if (this->listen_ != nullptr && !this->link_.connected() && !this->link_was_up_ && this->listen_->ready()) {
this->accept_client_();
}
this->listener_.poll(this->link_, !this->link_was_up_);
} else {
this->link_.poll();
}
#else
this->link_.poll();
#endif
if (this->link_.connected() != this->link_was_up_) {
this->sync_link_();
}
+12 -4
View File
@@ -2,6 +2,9 @@
#include "esphome/components/binary_sensor/binary_sensor.h"
#include "esphome/components/socket/tcp_client_link.h"
#ifdef USE_SOCKET_TCP_LISTENER
#include "esphome/components/socket/tcp_listener.h"
#endif
#include "esphome/components/uart/uart.h"
#include "esphome/core/component.h"
@@ -15,9 +18,14 @@ class UartTcp : public Component, public uart::UARTDevice {
public:
void set_host(const char *host) { this->link_.set_host(host); }
void set_port(uint16_t port) { this->link_.set_port(port); }
void set_server(bool server) { this->server_ = server; }
void set_reconnect_interval(uint32_t ms) { this->link_.set_reconnect_interval(ms); }
void set_connected_sensor(binary_sensor::BinarySensor *sensor) { this->connected_sensor_ = sensor; }
#ifdef USE_SOCKET_TCP_LISTENER
void set_server(bool server) { this->server_ = server; }
#ifdef USE_SOCKET_IPV4_ALLOW
void set_allow(const socket::Ipv4AllowEntry *entries, size_t count) { this->listener_.set_allow(entries, count); }
#endif
#endif
void setup() override;
void loop() override;
@@ -27,8 +35,6 @@ class UartTcp : public Component, public uart::UARTDevice {
protected:
void sync_link_();
void try_listen_();
void accept_client_();
void read_socket_();
void read_uart_();
void discard_uart_();
@@ -36,7 +42,9 @@ class UartTcp : public Component, public uart::UARTDevice {
static constexpr size_t READ_CHUNK = 128;
socket::TcpClientLink link_;
std::unique_ptr<socket::ListenSocket> listen_;
#ifdef USE_SOCKET_TCP_LISTENER
socket::TcpListener listener_;
#endif
binary_sensor::BinarySensor *connected_sensor_{nullptr};
bool server_{false};
// The link state loop() saw last; edges clear the buffer and publish the sensor.
+2
View File
@@ -419,8 +419,10 @@
#define USE_SENDSPIN_VISUALIZER
#define USE_SENDSPIN_PORT 8928 // NOLINT
#define USE_SOCKET_IMPL_BSD_SOCKETS
#define USE_SOCKET_IPV4_ALLOW
#define USE_SOCKET_IPV4_RESOLVE
#define USE_SOCKET_TCP_CLIENT_LINK
#define USE_SOCKET_TCP_LISTENER
#define USE_LWIP_FAST_SELECT
#define USE_SPEAKER
+3
View File
@@ -7,5 +7,8 @@ uart_tcp:
uart_id: uart_bus
role: server
port: 502
allowed_ips:
- 192.168.1.10
- 192.168.1.0/24
connected:
name: UART TCP Connected
@@ -18,5 +18,13 @@ uart_tcp:
uart_id: uart_bus
role: server
port: 18126
allowed_ips:
- 127.0.0.1
connected:
name: Bridge Connected
- id: denied_bridge
uart_id: uart_bus
role: server
port: 18127
allowed_ips:
- 192.0.2.1
+10
View File
@@ -25,6 +25,7 @@ async def test_uart_tcp_bridge(
unused_tcp_port_factory,
) -> None:
server_port = unused_tcp_port_factory()
denied_port = unused_tcp_port_factory()
controller_fd, device_fd = os.openpty()
os.set_blocking(controller_fd, False)
# uart's validate_port wants a two segment device path; Linux ptys live at
@@ -32,6 +33,7 @@ async def test_uart_tcp_bridge(
pty_link = f"/tmp/uart-tcp-pty-{os.getpid()}"
pathlib.Path(pty_link).symlink_to(os.ttyname(device_fd))
yaml_config = yaml_config.replace("port: 18126", f"port: {server_port}")
yaml_config = yaml_config.replace("port: 18127", f"port: {denied_port}")
yaml_config = yaml_config.replace("PTY_PATH", pty_link)
lines = LineWaiter()
@@ -103,6 +105,14 @@ async def test_uart_tcp_bridge(
await writer.drain()
assert await read_uart(5) == b"down2"
writer.close()
# A peer outside the allow list is rejected and closed.
denied_reader, denied_writer = await asyncio.open_connection(
"127.0.0.1", denied_port
)
await lines.wait_for("Rejected 127.0.0.1")
assert await asyncio.wait_for(denied_reader.read(8), 10) == b""
denied_writer.close()
finally:
loop.remove_reader(controller_fd)
os.close(controller_fd)
@@ -18,9 +18,13 @@ def test_network_order_swaps_to_sockaddr_value() -> None:
def test_add_ipv4_allow_emits_nothing_for_an_empty_list() -> None:
setter = MagicMock()
with patch.object(socket.cg, "add") as add:
with (
patch.object(socket.cg, "add") as add,
patch.object(socket.cg, "add_define") as add_define,
):
socket.add_ipv4_allow(setter, [], "bridge")
add.assert_not_called()
add_define.assert_not_called()
setter.assert_not_called()
@@ -29,6 +33,7 @@ def test_add_ipv4_allow_wires_the_setter_with_cleared_host_bits() -> None:
networks = [IPv4Network("192.168.175.33/24", strict=False)]
with (
patch.object(socket.cg, "add") as add,
patch.object(socket.cg, "add_define") as add_define,
patch.object(socket.cg, "progmem_array") as array,
):
socket.add_ipv4_allow(setter, networks, "bridge")
@@ -37,6 +42,7 @@ def test_add_ipv4_allow_wires_the_setter_with_cleared_host_bits() -> None:
assert str(socket._network_order(IPv4Address("255.255.255.0"))) in rendered
setter.assert_called_once_with(array.return_value, 1)
add.assert_called_once()
add_define.assert_called_once_with("USE_SOCKET_IPV4_ALLOW")
def test_schema_caps_the_list_length() -> None:
@@ -13,6 +13,7 @@ def test_helper_files_filtered_until_required() -> None:
filtered = socket.FILTER_SOURCE_FILES()
assert "ipv4_resolve.cpp" in filtered
assert "tcp_client_link.cpp" in filtered
assert "tcp_listener.cpp" in filtered
mock_core.defines = {Define("USE_SOCKET_IPV4_RESOLVE")}
filtered = socket.FILTER_SOURCE_FILES()
@@ -22,10 +23,12 @@ def test_helper_files_filtered_until_required() -> None:
mock_core.defines = {
Define("USE_SOCKET_IPV4_RESOLVE"),
Define("USE_SOCKET_TCP_CLIENT_LINK"),
Define("USE_SOCKET_TCP_LISTENER"),
}
filtered = socket.FILTER_SOURCE_FILES()
assert "ipv4_resolve.cpp" not in filtered
assert "tcp_client_link.cpp" not in filtered
assert "tcp_listener.cpp" not in filtered
def test_require_tcp_client_link_pulls_in_the_resolver() -> None:
@@ -36,3 +39,13 @@ def test_require_tcp_client_link_pulls_in_the_resolver() -> None:
"USE_SOCKET_IPV4_RESOLVE",
"USE_SOCKET_TCP_CLIENT_LINK",
}
def test_require_tcp_listener_pulls_in_the_link() -> None:
with patch.object(socket.cg, "add_define") as add_define:
socket.require_tcp_listener()
assert {call.args[0] for call in add_define.call_args_list} == {
"USE_SOCKET_IPV4_RESOLVE",
"USE_SOCKET_TCP_CLIENT_LINK",
"USE_SOCKET_TCP_LISTENER",
}