Prefix the poll result enumerators, report a reset after connect as ECONNRESET, and test the set_sockaddr contract on host

This commit is contained in:
J. Nick Koston
2026-09-05 15:39:31 +02:00
parent 457e224e19
commit 329dec0e9a
7 changed files with 84 additions and 19 deletions
@@ -107,15 +107,15 @@ void AsyncClient::loop() {
if (connecting_) {
int err = 0;
switch (socket::poll_connect(*socket_, err)) {
case socket::ConnectPollResult::CONNECT_POLL_PENDING:
case socket::ConnectPollResult::CONNECT_POLL_RESULT_PENDING:
break;
case socket::ConnectPollResult::CONNECT_POLL_CONNECTED:
case socket::ConnectPollResult::CONNECT_POLL_RESULT_CONNECTED:
connecting_ = false;
connected_ = true;
if (connect_cb_)
connect_cb_(connect_arg_, this);
break;
case socket::ConnectPollResult::CONNECT_POLL_ERROR:
case socket::ConnectPollResult::CONNECT_POLL_RESULT_ERROR:
ESP_LOGW(TAG, "Connection failed: %d", err);
close();
if (error_cb_)
+3 -3
View File
@@ -207,9 +207,9 @@ static constexpr size_t SOCKADDR_STR_LEN = 16; // INET_ADDRSTRLEN
/// Outcome of polling a non-blocking connect(); see socket::poll_connect().
enum class ConnectPollResult : uint8_t {
CONNECT_POLL_PENDING,
CONNECT_POLL_CONNECTED,
CONNECT_POLL_ERROR,
CONNECT_POLL_RESULT_PENDING,
CONNECT_POLL_RESULT_CONNECTED,
CONNECT_POLL_RESULT_ERROR,
};
} // namespace esphome::socket
@@ -496,21 +496,23 @@ int LWIPRawImpl::connect(const struct sockaddr *addr, socklen_t addrlen) {
ConnectPollResult LWIPRawImpl::poll_connect(int &err_out) const {
// pcb_ first; see the ordering note on the declaration
if (this->pcb_ == nullptr) {
err_out = this->connect_err_ == 0 || this->connect_err_ == EINPROGRESS ? ECONNRESET : this->connect_err_;
return ConnectPollResult::CONNECT_POLL_ERROR;
// Only a recorded connect failure carries its own reason
const bool failed = this->connect_err_ == ECONNREFUSED || this->connect_err_ == ETIMEDOUT;
err_out = failed ? this->connect_err_ : ECONNRESET;
return ConnectPollResult::CONNECT_POLL_RESULT_ERROR;
}
switch (this->connect_err_) {
case EINPROGRESS:
yield_to_sys(); // so the SYN-ACK is processed between polls
return ConnectPollResult::CONNECT_POLL_PENDING;
return ConnectPollResult::CONNECT_POLL_RESULT_PENDING;
case EISCONN:
return ConnectPollResult::CONNECT_POLL_CONNECTED;
return ConnectPollResult::CONNECT_POLL_RESULT_CONNECTED;
case 0:
err_out = EINVAL; // no connect was started
return ConnectPollResult::CONNECT_POLL_ERROR;
return ConnectPollResult::CONNECT_POLL_RESULT_ERROR;
default:
err_out = this->connect_err_;
return ConnectPollResult::CONNECT_POLL_ERROR;
return ConnectPollResult::CONNECT_POLL_RESULT_ERROR;
}
}
+6 -6
View File
@@ -208,7 +208,7 @@ ConnectPollResult poll_connect(Socket &sock, int &err_out) {
if (fd < 0 || fd >= FD_SETSIZE) {
// FD_SET on either is undefined behavior
err_out = EBADF;
return ConnectPollResult::CONNECT_POLL_ERROR;
return ConnectPollResult::CONNECT_POLL_RESULT_ERROR;
}
// Connect completion is a write event; the main loop only selects on reads
fd_set writefds;
@@ -224,22 +224,22 @@ ConnectPollResult poll_connect(Socket &sock, int &err_out) {
#endif
if (ret < 0) {
err_out = errno;
return ConnectPollResult::CONNECT_POLL_ERROR;
return ConnectPollResult::CONNECT_POLL_RESULT_ERROR;
}
if (ret == 0) {
return ConnectPollResult::CONNECT_POLL_PENDING;
return ConnectPollResult::CONNECT_POLL_RESULT_PENDING;
}
int error = 0;
socklen_t len = sizeof(error);
if (sock.getsockopt(SOL_SOCKET, SO_ERROR, &error, &len) != 0) {
err_out = errno;
return ConnectPollResult::CONNECT_POLL_ERROR;
return ConnectPollResult::CONNECT_POLL_RESULT_ERROR;
}
if (error != 0) {
err_out = error;
return ConnectPollResult::CONNECT_POLL_ERROR;
return ConnectPollResult::CONNECT_POLL_RESULT_ERROR;
}
return ConnectPollResult::CONNECT_POLL_CONNECTED;
return ConnectPollResult::CONNECT_POLL_RESULT_CONNECTED;
}
#endif
+1 -1
View File
@@ -147,7 +147,7 @@ socklen_t set_sockaddr_any(struct sockaddr *addr, socklen_t addrlen, uint16_t po
/// Check a non-blocking connect() for completion without blocking. Only
/// meaningful after connect() returned -1 with errno EINPROGRESS. On
/// CONNECT_POLL_ERROR, err_out holds the socket's SO_ERROR (or errno when the
/// CONNECT_POLL_RESULT_ERROR, err_out holds the socket's SO_ERROR (or errno when the
/// poll itself failed) on fd based implementations, and the failure recorded
/// by the lwip callbacks on the raw lwip implementation.
#ifdef USE_SOCKET_IMPL_LWIP_TCP
@@ -0,0 +1,18 @@
esphome:
name: socket-set-sockaddr
on_boot:
then:
- lambda: |-
// Exercise the contract callers rely on: 0 for text that is not an
// address, the sockaddr length otherwise, broadcast included
struct sockaddr_storage addr;
auto *sa = reinterpret_cast<struct sockaddr *>(&addr);
ESP_LOGI("test", "SET_SOCKADDR invalid=%u valid=%u broadcast=%u",
(unsigned) socket::set_sockaddr(sa, sizeof(addr), "not an address", 1234),
(unsigned) socket::set_sockaddr(sa, sizeof(addr), "192.0.2.1", 1234),
(unsigned) socket::set_sockaddr(sa, sizeof(addr), "255.255.255.255", 1234));
host:
api:
logger:
level: INFO
@@ -0,0 +1,45 @@
"""Integration test for the socket::set_sockaddr failure contract.
Callers skip an address when set_sockaddr returns 0, so text that is not an
address must return 0 while a real address, including the broadcast address
that inet_addr() would have reported as a failure, returns its length.
"""
import asyncio
import re
import pytest
from .types import APIClientConnectedFactory, RunCompiledFunction
@pytest.mark.asyncio
async def test_socket_set_sockaddr(
yaml_config: str,
run_compiled: RunCompiledFunction,
api_client_connected: APIClientConnectedFactory,
) -> None:
"""set_sockaddr reports an invalid address with 0 and accepts broadcast."""
loop = asyncio.get_running_loop()
result: asyncio.Future[tuple[int, int, int]] = loop.create_future()
def on_log_line(line: str) -> None:
match = re.search(
r"SET_SOCKADDR invalid=(\d+) valid=(\d+) broadcast=(\d+)", line
)
if match and not result.done():
result.set_result(tuple(int(g) for g in match.groups()))
async with (
run_compiled(yaml_config, line_callback=on_log_line),
api_client_connected() as client,
):
assert (await client.device_info()).name == "socket-set-sockaddr"
try:
invalid, valid, broadcast = await asyncio.wait_for(result, timeout=10.0)
except TimeoutError:
pytest.fail("SET_SOCKADDR marker never appeared")
assert invalid == 0
assert valid > 0
assert broadcast == valid