[tcp_uart] Add a server role with an IPv4 allow list (#20026)

This commit is contained in:
Bascht74
2026-10-02 13:47:10 -05:00
committed by GitHub
parent de0a0b2af4
commit 5328e5f5e0
8 changed files with 205 additions and 36 deletions
+49 -20
View File
@@ -5,6 +5,7 @@ from esphome.components.const import (
CONF_HOST,
CONF_PARITY,
CONF_RECONNECT_INTERVAL,
CONF_ROLE,
CONF_STOP_BITS,
)
import esphome.config_validation as cv
@@ -25,43 +26,71 @@ MULTI_CONF = True
tcp_uart_ns = cg.esphome_ns.namespace("tcp_uart")
TcpUart = tcp_uart_ns.class_("TcpUart", uart.UARTComponent, cg.Component)
CONF_ALLOWED_IPS = "allowed_ips"
CONF_CONNECTED = "connected"
BASE_SCHEMA = cv.Schema(
{
cv.GenerateID(): cv.declare_id(TcpUart),
cv.Required(CONF_PORT): cv.port,
cv.Optional(CONF_BAUD_RATE, default=9600): cv.int_range(min=1),
cv.Optional(CONF_DATA_BITS, default=8): cv.int_range(min=5, max=8),
cv.Optional(CONF_PARITY, default="NONE"): cv.enum(
uart.UART_PARITY_OPTIONS, upper=True
),
cv.Optional(CONF_STOP_BITS, default=1): cv.one_of(1, 2, int=True),
cv.Optional(
CONF_RECONNECT_INTERVAL, default="5s"
): cv.positive_time_period_milliseconds,
cv.Optional(CONF_CONNECTED): binary_sensor.binary_sensor_schema(
device_class=DEVICE_CLASS_CONNECTIVITY,
entity_category=ENTITY_CATEGORY_DIAGNOSTIC,
),
}
).extend(cv.COMPONENT_SCHEMA)
CONFIG_SCHEMA = cv.All(
cv.Schema(
cv.typed_schema(
{
cv.GenerateID(): cv.declare_id(TcpUart),
cv.Required(CONF_HOST): cv.string,
cv.Required(CONF_PORT): cv.port,
cv.Optional(CONF_BAUD_RATE, default=9600): cv.int_range(min=1),
cv.Optional(CONF_DATA_BITS, default=8): cv.int_range(min=5, max=8),
cv.Optional(CONF_PARITY, default="NONE"): cv.enum(
uart.UART_PARITY_OPTIONS, upper=True
"client": BASE_SCHEMA.extend(
{
cv.Required(CONF_HOST): cv.string,
}
),
cv.Optional(CONF_STOP_BITS, default=1): cv.one_of(1, 2, int=True),
cv.Optional(
CONF_RECONNECT_INTERVAL, default="5s"
): cv.positive_time_period_milliseconds,
cv.Optional(CONF_CONNECTED): binary_sensor.binary_sensor_schema(
device_class=DEVICE_CLASS_CONNECTIVITY,
entity_category=ENTITY_CATEGORY_DIAGNOSTIC,
"server": BASE_SCHEMA.extend(
{
cv.Optional(CONF_ALLOWED_IPS): socket.IPV4_ALLOW_SCHEMA,
}
),
}
).extend(cv.COMPONENT_SCHEMA),
socket.consume_sockets(1, "tcp_uart"),
},
key=CONF_ROLE,
default_type="client",
lower=True,
),
socket.consume_role_sockets("tcp_uart"),
)
async def to_code(config: ConfigType) -> None:
socket.require_tcp_client_link()
var = cg.new_Pvariable(config[CONF_ID], config[CONF_HOST], config[CONF_PORT])
var = cg.new_Pvariable(config[CONF_ID])
await cg.register_component(var, config)
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]))
# The socket is not clocked. These only satisfy UARTComponent and a consumer check.
cg.add(var.set_baud_rate(config[CONF_BAUD_RATE]))
cg.add(var.set_data_bits(config[CONF_DATA_BITS]))
cg.add(var.set_stop_bits(config[CONF_STOP_BITS]))
cg.add(var.set_parity(config[CONF_PARITY]))
if (host := config.get(CONF_HOST)) is not None:
cg.add(var.set_host(host))
binary_sensors = binary_sensor.sub_binary_sensors(config)
await binary_sensors(CONF_CONNECTED, var.set_connected_sensor)
+27 -2
View File
@@ -14,6 +14,9 @@ static constexpr uint32_t DROP_LOG_INTERVAL_MS = 5000;
void TcpUart::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);
}
@@ -22,12 +25,24 @@ void TcpUart::setup() {
void TcpUart::dump_config() {
ESP_LOGCONFIG(TAG,
"TCP UART:\n"
" Host: %s:%u\n"
" %s: %s:%u\n"
" Reconnect Interval: %" PRIu32 "ms",
this->link_.host(), this->link_.port(), this->link_.reconnect_interval());
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 TcpUart::on_shutdown() {
this->link_.close();
#ifdef USE_SOCKET_TCP_LISTENER
this->listener_.close();
#endif
}
void TcpUart::sync_link_() {
bool up = this->link_.connected();
this->link_was_up_ = up;
@@ -63,7 +78,17 @@ void TcpUart::read_socket_() {
}
void TcpUart::loop() {
#ifdef USE_SOCKET_TCP_LISTENER
if (this->server_) {
// link_was_up_ holds the accept until the previous drop's edge has run,
// so the sensor and the cleared RX buffer always see the disconnect.
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_();
}
+18 -7
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_component.h"
#include "esphome/core/component.h"
@@ -9,22 +12,26 @@
namespace esphome::tcp_uart {
/// TCP client presented as a UART. Bytes are copied unchanged.
/// TCP client or server presented as a UART. Bytes are copied unchanged.
class TcpUart : public uart::UARTComponent, public Component {
public:
TcpUart(const char *host, uint16_t port) {
this->link_.set_host(host);
this->link_.set_port(port);
this->rx_buffer_size_ = RX_BUFFER_SIZE;
}
TcpUart() { this->rx_buffer_size_ = RX_BUFFER_SIZE; }
void set_host(const char *host) { this->link_.set_host(host); }
void set_port(uint16_t port) { this->link_.set_port(port); }
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;
void dump_config() override;
void on_shutdown() override { this->link_.close(); }
void on_shutdown() override;
float get_setup_priority() const override { return setup_priority::AFTER_WIFI; }
void write_array(const uint8_t *data, size_t len) override;
@@ -47,11 +54,15 @@ class TcpUart : public uart::UARTComponent, public Component {
static constexpr size_t RX_BUFFER_SIZE = 1024;
socket::TcpClientLink link_;
#ifdef USE_SOCKET_TCP_LISTENER
socket::TcpListener listener_;
#endif
binary_sensor::BinarySensor *connected_sensor_{nullptr};
uint32_t last_drop_log_ms_{0};
// rx_[rx_start_, rx_end_) holds unread bytes; read_socket_() compacts to the front.
uint16_t rx_start_{0};
uint16_t rx_end_{0};
bool server_{false};
// The link state loop() saw last; edges clear rx_ and publish the sensor.
bool link_was_up_{false};
// A read stopped before EAGAIN. ready() stays false until new data arrives.
+9 -6
View File
@@ -3,18 +3,21 @@ wifi:
password: password1
tcp_uart:
- id: tcp_uart_1
host: 192.0.2.10
port: 502
- id: tcp_uart_server
role: server
port: 5020
reconnect_interval: 10s
allowed_ips:
- 192.0.2.20
- 192.0.2.0/24
connected:
name: TCP UART Connected
name: TCP UART Server Connected
interval:
- interval: 60s
then:
- lambda: |-
uint8_t byte;
if (id(tcp_uart_1).available() && id(tcp_uart_1).read_byte(&byte)) {
id(tcp_uart_1).write_byte(byte);
if (id(tcp_uart_server).available() && id(tcp_uart_server).read_byte(&byte)) {
id(tcp_uart_server).write_byte(byte);
}
@@ -0,0 +1,11 @@
wifi:
ssid: MySSID
password: password1
tcp_uart:
- id: tcp_uart_1
host: 192.0.2.10
port: 502
reconnect_interval: 10s
connected:
name: TCP UART Connected
@@ -13,7 +13,11 @@ namespace esphome::tcp_uart::testing {
class TcpUartUnderTest : public TcpUart {
public:
TcpUartUnderTest() : TcpUart("peer", 1) { this->link_.begin("flush_test"); }
TcpUartUnderTest() {
this->set_host("peer");
this->set_port(1);
this->link_.begin("flush_test");
}
socket::TcpClientLink &link() { return this->link_; }
};
@@ -0,0 +1,30 @@
esphome:
name: tcp-uart-server-test
host:
api:
logger:
level: INFO
tcp_uart:
- id: allowed_bus
role: server
port: 18126
allowed_ips:
- 127.0.0.1
- id: denied_bus
role: server
port: 18127
allowed_ips:
- 192.0.2.1
interval:
- interval: 50ms
then:
- lambda: |-
uint8_t b;
while (id(allowed_bus).read_byte(&b)) {
id(allowed_bus).write_byte(b);
}
+56
View File
@@ -0,0 +1,56 @@
"""Integration test for a tcp_uart server on host.
Pytest connects as the TCP client. One server allows 127.0.0.1 and echoes.
The other allows only 192.0.2.1, so the same client is closed.
"""
from __future__ import annotations
import asyncio
import contextlib
import pytest
from .log_utils import LineWaiter
from .types import APIClientConnectedFactory, RunCompiledFunction
PAYLOAD = b"ping!"
@pytest.mark.asyncio
async def test_tcp_uart_server(
yaml_config: str,
run_compiled: RunCompiledFunction,
api_client_connected: APIClientConnectedFactory,
unused_tcp_port_factory,
) -> None:
allowed_port = unused_tcp_port_factory()
denied_port = unused_tcp_port_factory()
yaml_config = yaml_config.replace("port: 18126", f"port: {allowed_port}")
yaml_config = yaml_config.replace("port: 18127", f"port: {denied_port}")
lines = LineWaiter()
async with (
run_compiled(yaml_config, line_callback=lines.callback),
api_client_connected() as client,
):
device_info = await client.device_info()
assert device_info is not None
assert device_info.name == "tcp-uart-server-test"
await lines.wait_for(f"Listening on {allowed_port}")
await lines.wait_for(f"Listening on {denied_port}")
reader, writer = await asyncio.open_connection("127.0.0.1", allowed_port)
await lines.wait_for("Client connected from 127.0.0.1")
writer.write(PAYLOAD)
await writer.drain()
assert await asyncio.wait_for(reader.readexactly(len(PAYLOAD)), 10) == PAYLOAD
writer.close()
denied_reader, denied_writer = await asyncio.open_connection(
"127.0.0.1", denied_port
)
await lines.wait_for("Rejected 127.0.0.1")
with contextlib.suppress(ConnectionResetError):
assert await asyncio.wait_for(denied_reader.read(8), 10) == b""
denied_writer.close()