diff --git a/esphome/components/tcp_uart/__init__.py b/esphome/components/tcp_uart/__init__.py index e46e4dbfd0..dbe805c281 100644 --- a/esphome/components/tcp_uart/__init__.py +++ b/esphome/components/tcp_uart/__init__.py @@ -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) diff --git a/esphome/components/tcp_uart/tcp_uart.cpp b/esphome/components/tcp_uart/tcp_uart.cpp index b24280c8b1..5913610ee8 100644 --- a/esphome/components/tcp_uart/tcp_uart.cpp +++ b/esphome/components/tcp_uart/tcp_uart.cpp @@ -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_(); } diff --git a/esphome/components/tcp_uart/tcp_uart.h b/esphome/components/tcp_uart/tcp_uart.h index 7ce822c2fb..29f7af564a 100644 --- a/esphome/components/tcp_uart/tcp_uart.h +++ b/esphome/components/tcp_uart/tcp_uart.h @@ -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. diff --git a/tests/components/tcp_uart/common.yaml b/tests/components/tcp_uart/common.yaml index f3c1d061aa..2c2d264f7f 100644 --- a/tests/components/tcp_uart/common.yaml +++ b/tests/components/tcp_uart/common.yaml @@ -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); } diff --git a/tests/components/tcp_uart/test-client.esp32-idf.yaml b/tests/components/tcp_uart/test-client.esp32-idf.yaml new file mode 100644 index 0000000000..c63c6a7c7d --- /dev/null +++ b/tests/components/tcp_uart/test-client.esp32-idf.yaml @@ -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 diff --git a/tests/components/tcp_uart/test_flush_host.cpp b/tests/components/tcp_uart/test_flush_host.cpp index 9c6c63110c..5573d587ce 100644 --- a/tests/components/tcp_uart/test_flush_host.cpp +++ b/tests/components/tcp_uart/test_flush_host.cpp @@ -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_; } }; diff --git a/tests/integration/fixtures/tcp_uart_server.yaml b/tests/integration/fixtures/tcp_uart_server.yaml new file mode 100644 index 0000000000..20bc221d9d --- /dev/null +++ b/tests/integration/fixtures/tcp_uart_server.yaml @@ -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); + } diff --git a/tests/integration/test_tcp_uart_server.py b/tests/integration/test_tcp_uart_server.py new file mode 100644 index 0000000000..9a69ec42fd --- /dev/null +++ b/tests/integration/test_tcp_uart_server.py @@ -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()