From b6c79fdc9ae41ec6d6920827190ea23d38e8835e Mon Sep 17 00:00:00 2001 From: Bascht74 <66269310+Bascht74@users.noreply.github.com> Date: Fri, 2 Oct 2026 20:07:00 +0200 Subject: [PATCH] [socket] Share the server side as TcpListener (#20037) Co-authored-by: J. Nick Koston --- esphome/components/socket/__init__.py | 20 +++++ esphome/components/socket/ipv4_allow.h | 13 ++- esphome/components/socket/tcp_listener.cpp | 90 +++++++++++++++++++ esphome/components/socket/tcp_listener.h | 60 +++++++++++++ esphome/components/uart_tcp/__init__.py | 23 ++--- esphome/components/uart_tcp/uart_tcp.cpp | 66 +++----------- esphome/components/uart_tcp/uart_tcp.h | 16 +++- esphome/core/defines.h | 2 + tests/components/uart_tcp/common.yaml | 3 + .../integration/fixtures/uart_tcp_bridge.yaml | 8 ++ tests/integration/test_uart_tcp_bridge.py | 10 +++ .../socket/test_socket_ipv4_allow.py | 8 +- .../socket/test_socket_source_filter.py | 13 +++ 13 files changed, 262 insertions(+), 70 deletions(-) create mode 100644 esphome/components/socket/tcp_listener.cpp create mode 100644 esphome/components/socket/tcp_listener.h diff --git a/esphome/components/socket/__init__.py b/esphome/components/socket/__init__.py index 850f7ebcbc..269791f5e7 100644 --- a/esphome/components/socket/__init__.py +++ b/esphome/components/socket/__init__.py @@ -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", } ) diff --git a/esphome/components/socket/ipv4_allow.h b/esphome/components/socket/ipv4_allow.h index adb00b5449..f2f66edc05 100644 --- a/esphome/components/socket/ipv4_allow.h +++ b/esphome/components/socket/ipv4_allow.h @@ -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}; diff --git a/esphome/components/socket/tcp_listener.cpp b/esphome/components/socket/tcp_listener.cpp new file mode 100644 index 0000000000..e1f00e2927 --- /dev/null +++ b/esphome/components/socket/tcp_listener.cpp @@ -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 +#include + +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(&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(&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(&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(&peer); + char text[SOCKADDR_STR_LEN]; + format_sockaddr_to(sa, peer_len, std::span(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(&e.addr); + ESP_LOGCONFIG(this->tag_, " Allowed IP: %u.%u.%u.%u/%u", b[0], b[1], b[2], b[3], + static_cast(__builtin_popcount(e.mask))); + } +#endif +} + +} // namespace esphome::socket + +#endif diff --git a/esphome/components/socket/tcp_listener.h b/esphome/components/socket/tcp_listener.h new file mode 100644 index 0000000000..a42080e641 --- /dev/null +++ b/esphome/components/socket/tcp_listener.h @@ -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 +#include + +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 listen_; + const char *tag_{nullptr}; +#ifdef USE_SOCKET_IPV4_ALLOW + uint32_t last_reject_log_ms_{0}; + Ipv4Allow allow_; +#endif +}; + +} // namespace esphome::socket + +#endif diff --git a/esphome/components/uart_tcp/__init__.py b/esphome/components/uart_tcp/__init__.py index e7a85a0f2d..cd56e2c8bc 100644 --- a/esphome/components/uart_tcp/__init__.py +++ b/esphome/components/uart_tcp/__init__.py @@ -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: diff --git a/esphome/components/uart_tcp/uart_tcp.cpp b/esphome/components/uart_tcp/uart_tcp.cpp index 2531b73b20..cf15802b55 100644 --- a/esphome/components/uart_tcp/uart_tcp.cpp +++ b/esphome/components/uart_tcp/uart_tcp.cpp @@ -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(&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(&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_(); } diff --git a/esphome/components/uart_tcp/uart_tcp.h b/esphome/components/uart_tcp/uart_tcp.h index 2a59bf968f..1e19166f87 100644 --- a/esphome/components/uart_tcp/uart_tcp.h +++ b/esphome/components/uart_tcp/uart_tcp.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.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 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. diff --git a/esphome/core/defines.h b/esphome/core/defines.h index 21c4e3d996..97b1395553 100644 --- a/esphome/core/defines.h +++ b/esphome/core/defines.h @@ -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 diff --git a/tests/components/uart_tcp/common.yaml b/tests/components/uart_tcp/common.yaml index 8010878e89..eacad3dd6d 100644 --- a/tests/components/uart_tcp/common.yaml +++ b/tests/components/uart_tcp/common.yaml @@ -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 diff --git a/tests/integration/fixtures/uart_tcp_bridge.yaml b/tests/integration/fixtures/uart_tcp_bridge.yaml index c61add801f..2ab70a03b6 100644 --- a/tests/integration/fixtures/uart_tcp_bridge.yaml +++ b/tests/integration/fixtures/uart_tcp_bridge.yaml @@ -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 diff --git a/tests/integration/test_uart_tcp_bridge.py b/tests/integration/test_uart_tcp_bridge.py index b98d3ffac9..909df32410 100644 --- a/tests/integration/test_uart_tcp_bridge.py +++ b/tests/integration/test_uart_tcp_bridge.py @@ -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) diff --git a/tests/unit_tests/components/socket/test_socket_ipv4_allow.py b/tests/unit_tests/components/socket/test_socket_ipv4_allow.py index bc1185c26f..8a2f60a4b9 100644 --- a/tests/unit_tests/components/socket/test_socket_ipv4_allow.py +++ b/tests/unit_tests/components/socket/test_socket_ipv4_allow.py @@ -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: diff --git a/tests/unit_tests/components/socket/test_socket_source_filter.py b/tests/unit_tests/components/socket/test_socket_source_filter.py index 3967c9568d..c46a3f1e8f 100644 --- a/tests/unit_tests/components/socket/test_socket_source_filter.py +++ b/tests/unit_tests/components/socket/test_socket_source_filter.py @@ -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", + }