[socket] Add a reconnecting TCP client link (#19999)

Co-authored-by: J. Nick Koston <nick@home-assistant.io>
This commit is contained in:
Bascht74
2026-10-01 23:40:09 +00:00
committed by GitHub
co-authored by J. Nick Koston
parent 99070187b8
commit b324a7c6e0
7 changed files with 419 additions and 0 deletions
@@ -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 <algorithm>
#include <cerrno>
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<struct sockaddr *>(&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<struct sockaddr *>(&dest), dest_len) != 0 && errno != EINPROGRESS) {
this->drop_(LOG_STR("Connect failed"), errno);
}
}
void TcpClientLink::adopt(std::unique_ptr<Socket> 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
@@ -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 <cstdint>
#include <memory>
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<Socket> 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<Socket> 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
@@ -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]))
@@ -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<int>(count));
this->link_.write(buf, static_cast<size_t>(count));
}
}
} // namespace esphome::tcp_client_link_test_component
@@ -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
@@ -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
@@ -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()