diff --git a/esphome/components/socket/tcp_client_link.cpp b/esphome/components/socket/tcp_client_link.cpp new file mode 100644 index 0000000000..f2c1b8e4c1 --- /dev/null +++ b/esphome/components/socket/tcp_client_link.cpp @@ -0,0 +1,150 @@ +#include "tcp_client_link.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 { + +// After this long in SYN, the stack's own retries are cut short. +static constexpr uint32_t CONNECT_TIMEOUT_MS = 10000; + +// Non-blocking options and TCP keepalive for a bridged stream socket. +// Keepalive is best-effort: the raw lwIP implementation (ESP8266, RP2040) +// rejects it, so a half-open link there is only detected by a failed write. +static void set_stream_options(Socket *sock) { + int yes = 1; + sock->setblocking(false); + sock->setsockopt(IPPROTO_TCP, TCP_NODELAY, &yes, sizeof(yes)); + sock->setsockopt(SOL_SOCKET, SO_KEEPALIVE, &yes, sizeof(yes)); +#ifdef TCP_KEEPIDLE + int idle = 30; + int interval = 10; + int count = 3; + sock->setsockopt(IPPROTO_TCP, TCP_KEEPIDLE, &idle, sizeof(idle)); + sock->setsockopt(IPPROTO_TCP, TCP_KEEPINTVL, &interval, sizeof(interval)); + sock->setsockopt(IPPROTO_TCP, TCP_KEEPCNT, &count, sizeof(count)); +#endif +} + +void TcpClientLink::begin(const char *tag) { + this->tag_ = tag; + // The first attempt must not wait out a full interval. + this->last_attempt_ms_ = App.get_loop_component_start_time() - this->reconnect_interval_ms_; +} + +void TcpClientLink::poll_slow_() { + if (this->sock_ == nullptr) { + this->try_connect_(); + return; + } + int err = 0; + switch (poll_connect(*this->sock_, err)) { + case ConnectPollResult::CONNECT_POLL_RESULT_PENDING: + // Give up before the stack's SYN retries do, so the interval stays honest + // and the next attempt resolves the host again. + if (App.get_loop_component_start_time() - this->last_attempt_ms_ >= + std::max(this->reconnect_interval_ms_, CONNECT_TIMEOUT_MS)) { + this->drop_(LOG_STR("Connect failed"), ETIMEDOUT); + } + return; + case ConnectPollResult::CONNECT_POLL_RESULT_ERROR: + this->drop_(LOG_STR("Connect failed"), err); + return; + default: + break; + } + this->connected_ = true; + ESP_LOGI(this->tag_, "Connected to %s:%u", this->host_.c_str(), this->port_); +} + +void TcpClientLink::try_connect_() { + if (this->resolved_.consume_failure()) { + this->note_attempt(); + return; + } + this->resolved_.start(this->host_.c_str(), this->port_, this->tag_); + if (!this->resolved_.ready()) { + return; + } + struct sockaddr_storage dest; + socklen_t dest_len = + this->resolved_.to_sockaddr(reinterpret_cast(&dest), sizeof(dest), this->port_); + if (dest_len == 0) { + this->note_attempt(); + return; + } + this->sock_ = socket_loop_monitored(dest.ss_family, SOCK_STREAM, IPPROTO_TCP); + if (this->sock_ == nullptr) { + this->drop_(LOG_STR("Connect failed"), errno); + return; + } + set_stream_options(this->sock_.get()); + // Starts the pending-connect clock that poll() times out against. + this->note_attempt(); + // An immediate success is reported by the next poll(); poll_connect() sees it writable. + if (this->sock_->connect(reinterpret_cast(&dest), dest_len) != 0 && errno != EINPROGRESS) { + this->drop_(LOG_STR("Connect failed"), errno); + } +} + +void TcpClientLink::adopt(std::unique_ptr sock) { + this->close(); + set_stream_options(sock.get()); + this->sock_ = std::move(sock); + this->connected_ = true; +} + +ssize_t TcpClientLink::read(uint8_t *buf, size_t len) { + if (!this->connected_) { + return 0; + } + ssize_t count = this->sock_->read(buf, len); + if (count > 0) { + return count; + } + if (count == 0 || (errno != EAGAIN && errno != EWOULDBLOCK)) { + this->drop_(LOG_STR("Connection lost"), count == 0 ? 0 : errno); + return -1; + } + return 0; +} + +ssize_t TcpClientLink::write(const uint8_t *buf, size_t len) { + if (!this->connected_ || len == 0) { + return 0; + } + ssize_t sent = this->sock_->write(buf, len); + if (sent >= 0) { + return sent; + } + if (errno == EAGAIN || errno == EWOULDBLOCK) { + return 0; + } + this->drop_(LOG_STR("Connection lost"), errno); + return -1; +} + +void TcpClientLink::close() { + if (this->sock_ != nullptr) { + this->sock_->shutdown(SHUT_RDWR); + this->sock_->close(); + this->sock_.reset(); + } + this->connected_ = false; + this->resolved_.forget(); +} + +void TcpClientLink::drop_(const LogString *what, int err) { + ESP_LOGW(this->tag_, "%s: %d", LOG_STR_ARG(what), err); + this->close(); + this->note_attempt(); +} + +} // namespace esphome::socket + +#endif diff --git a/esphome/components/socket/tcp_client_link.h b/esphome/components/socket/tcp_client_link.h new file mode 100644 index 0000000000..065c4df562 --- /dev/null +++ b/esphome/components/socket/tcp_client_link.h @@ -0,0 +1,74 @@ +#pragma once + +#include "headers.h" + +#if defined(USE_SOCKET_IMPL_LWIP_TCP) || defined(USE_SOCKET_IMPL_LWIP_SOCKETS) || defined(USE_SOCKET_IMPL_BSD_SOCKETS) + +#include "ipv4_resolve.h" +#include "socket.h" +#include "esphome/core/application.h" +#include "esphome/core/log.h" +#include "esphome/core/string_ref.h" + +#include +#include + +namespace esphome::socket { + +/// A reconnecting TCP stream driven from loop(). Owns the socket, the DNS +/// lookup and the retry backoff. A fatal read/write error closes the link +/// and schedules the next attempt; the caller sees the edge via connected(). +class TcpClientLink { + public: + void set_host(const char *host) { this->host_ = StringRef(host); } + void set_port(uint16_t port) { this->port_ = port; } + void set_reconnect_interval(uint32_t ms) { this->reconnect_interval_ms_ = ms; } + const char *host() const { return this->host_.c_str(); } + uint16_t port() const { return this->port_; } + uint32_t reconnect_interval() const { return this->reconnect_interval_ms_; } + + /// Call from setup(). tag names this link's log lines. + void begin(const char *tag); + /// Connect state machine; call every loop while acting as a client. + /// Inline no-op while connected or waiting out the backoff. + void poll() { + if (this->connected_ || (this->sock_ == nullptr && this->in_backoff())) { + return; + } + this->poll_slow_(); + } + /// Take over an accepted socket (the server side of a bridge). + void adopt(std::unique_ptr sock); + /// Returns bytes moved, 0 when nothing can move now, -1 when the link dropped. + ssize_t read(uint8_t *buf, size_t len); + ssize_t write(const uint8_t *buf, size_t len); + /// Close without scheduling a reconnect (shutdown). + void close(); + + bool connected() const { return this->connected_; } + bool ready() const { return this->sock_ != nullptr && this->sock_->ready(); } + /// Shared retry clock, also usable for a listen socket. + void note_attempt() { this->last_attempt_ms_ = App.get_loop_component_start_time(); } + bool in_backoff() const { + return App.get_loop_component_start_time() - this->last_attempt_ms_ < this->reconnect_interval_ms_; + } + + protected: + void poll_slow_(); + void try_connect_(); + /// Close after a failure, log what and errno, schedule the next attempt. + void drop_(const LogString *what, int err); + + StringRef host_; + std::unique_ptr sock_; + const char *tag_{nullptr}; + uint32_t last_attempt_ms_{0}; + uint32_t reconnect_interval_ms_{5000}; + Ipv4Resolve resolved_; + uint16_t port_{0}; + bool connected_{false}; +}; + +} // namespace esphome::socket + +#endif diff --git a/tests/integration/fixtures/external_components/tcp_client_link_test_component/__init__.py b/tests/integration/fixtures/external_components/tcp_client_link_test_component/__init__.py new file mode 100644 index 0000000000..8a7703d025 --- /dev/null +++ b/tests/integration/fixtures/external_components/tcp_client_link_test_component/__init__.py @@ -0,0 +1,35 @@ +import esphome.codegen as cg +from esphome.components.const import CONF_HOST +import esphome.config_validation as cv +from esphome.const import CONF_ID, CONF_PORT +from esphome.types import ConfigType + +AUTO_LOAD = ["socket"] + +CONF_RECONNECT_INTERVAL = "reconnect_interval" + +tcp_client_link_test_component_ns = cg.esphome_ns.namespace( + "tcp_client_link_test_component" +) +TcpClientLinkTestComponent = tcp_client_link_test_component_ns.class_( + "TcpClientLinkTestComponent", cg.Component +) + +CONFIG_SCHEMA = cv.Schema( + { + cv.GenerateID(): cv.declare_id(TcpClientLinkTestComponent), + cv.Required(CONF_HOST): cv.string, + cv.Required(CONF_PORT): cv.port, + cv.Optional( + CONF_RECONNECT_INTERVAL, default="1s" + ): cv.positive_time_period_milliseconds, + } +).extend(cv.COMPONENT_SCHEMA) + + +async def to_code(config: ConfigType) -> None: + var = cg.new_Pvariable(config[CONF_ID]) + await cg.register_component(var, config) + cg.add(var.set_host(config[CONF_HOST])) + cg.add(var.set_port(config[CONF_PORT])) + cg.add(var.set_reconnect_interval(config[CONF_RECONNECT_INTERVAL])) diff --git a/tests/integration/fixtures/external_components/tcp_client_link_test_component/tcp_client_link_test_component.cpp b/tests/integration/fixtures/external_components/tcp_client_link_test_component/tcp_client_link_test_component.cpp new file mode 100644 index 0000000000..7f2af8add8 --- /dev/null +++ b/tests/integration/fixtures/external_components/tcp_client_link_test_component/tcp_client_link_test_component.cpp @@ -0,0 +1,28 @@ +#include "tcp_client_link_test_component.h" +#include "esphome/core/log.h" + +namespace esphome::tcp_client_link_test_component { + +static const char *const TAG = "tcp_link_test"; + +void TcpClientLinkTestComponent::setup() { this->link_.begin(TAG); } + +void TcpClientLinkTestComponent::loop() { + this->link_.poll(); + bool up = this->link_.connected(); + if (up != this->was_up_) { + this->was_up_ = up; + ESP_LOGI(TAG, "Link %s", up ? LOG_STR_LITERAL("up") : LOG_STR_LITERAL("down")); + } + if (!up || !this->link_.ready()) { + return; + } + uint8_t buf[64]; + ssize_t count = this->link_.read(buf, sizeof(buf)); + if (count > 0) { + ESP_LOGI(TAG, "Echoing %d bytes", static_cast(count)); + this->link_.write(buf, static_cast(count)); + } +} + +} // namespace esphome::tcp_client_link_test_component diff --git a/tests/integration/fixtures/external_components/tcp_client_link_test_component/tcp_client_link_test_component.h b/tests/integration/fixtures/external_components/tcp_client_link_test_component/tcp_client_link_test_component.h new file mode 100644 index 0000000000..829f15b8a3 --- /dev/null +++ b/tests/integration/fixtures/external_components/tcp_client_link_test_component/tcp_client_link_test_component.h @@ -0,0 +1,24 @@ +#pragma once + +#include "esphome/components/socket/tcp_client_link.h" +#include "esphome/core/component.h" + +namespace esphome::tcp_client_link_test_component { + +/// Echoes every byte the link receives back to the peer and logs link edges. +class TcpClientLinkTestComponent : public Component { + 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_reconnect_interval(uint32_t ms) { this->link_.set_reconnect_interval(ms); } + + void setup() override; + void loop() override; + void on_shutdown() override { this->link_.close(); } + + protected: + socket::TcpClientLink link_; + bool was_up_{false}; +}; + +} // namespace esphome::tcp_client_link_test_component diff --git a/tests/integration/fixtures/socket_tcp_client_link.yaml b/tests/integration/fixtures/socket_tcp_client_link.yaml new file mode 100644 index 0000000000..2ed7f775ff --- /dev/null +++ b/tests/integration/fixtures/socket_tcp_client_link.yaml @@ -0,0 +1,20 @@ +esphome: + name: socket-tcp-client-link-test + +host: + +api: + +logger: + level: INFO + +external_components: + - source: + type: local + path: EXTERNAL_COMPONENT_PATH + components: [tcp_client_link_test_component] + +tcp_client_link_test_component: + host: 127.0.0.1 + port: 18123 + reconnect_interval: 1s diff --git a/tests/integration/test_socket_tcp_client_link.py b/tests/integration/test_socket_tcp_client_link.py new file mode 100644 index 0000000000..31cc344f67 --- /dev/null +++ b/tests/integration/test_socket_tcp_client_link.py @@ -0,0 +1,88 @@ +"""Integration test for socket::TcpClientLink on host. + +Pytest runs a real TCP server; the device echoes through the link. +Covers connect, read, write, a server-initiated drop and the reconnect. +""" + +from __future__ import annotations + +import asyncio +import contextlib + +import pytest + +from .types import APIClientConnectedFactory, RunCompiledFunction + +PAYLOAD = b"hello link" + + +@pytest.mark.asyncio +async def test_socket_tcp_client_link( + yaml_config: str, + run_compiled: RunCompiledFunction, + api_client_connected: APIClientConnectedFactory, + unused_tcp_port_factory, +) -> None: + server_port = unused_tcp_port_factory() + yaml_config = yaml_config.replace("port: 18123", f"port: {server_port}") + + echoed: list[bytes] = [] + echo_done = asyncio.Event() + reconnected = asyncio.Event() + link_down = asyncio.Event() + second_link_up = asyncio.Event() + link_up_count = 0 + + def on_log_line(line: str) -> None: + nonlocal link_up_count + if "Link up" in line: + link_up_count += 1 + if link_up_count >= 2: + second_link_up.set() + elif "Link down" in line: + link_down.set() + + async def handle( + reader: asyncio.StreamReader, writer: asyncio.StreamWriter + ) -> None: + if not echo_done.is_set(): + writer.write(PAYLOAD) + await writer.drain() + with contextlib.suppress(TimeoutError, asyncio.IncompleteReadError): + echoed.append( + await asyncio.wait_for(reader.readexactly(len(PAYLOAD)), 10) + ) + echo_done.set() + # Drop the connection so the link has to reconnect. + writer.close() + return + reconnected.set() + + server = await asyncio.start_server(handle, "127.0.0.1", server_port) + try: + async with ( + run_compiled(yaml_config, line_callback=on_log_line), + api_client_connected() as client, + ): + device_info = await client.device_info() + assert device_info is not None + assert device_info.name == "socket-tcp-client-link-test" + + try: + await asyncio.wait_for(echo_done.wait(), timeout=15.0) + except TimeoutError: + pytest.fail("Link never connected or echoed") + assert echoed and echoed[0] == PAYLOAD, "Echo payload mismatch" + + try: + await asyncio.wait_for(link_down.wait(), timeout=15.0) + except TimeoutError: + pytest.fail("Link never reported the drop") + try: + await asyncio.wait_for(reconnected.wait(), timeout=15.0) + await asyncio.wait_for(second_link_up.wait(), timeout=15.0) + except TimeoutError: + pytest.fail("Link did not reconnect after the server dropped it") + finally: + server.close() + await server.wait_closed()