[socket] Add an IPv4 lookup next to set_sockaddr (#19909)

Co-authored-by: pre-commit-ci-lite[bot] <117423508+pre-commit-ci-lite[bot]@users.noreply.github.com>
Co-authored-by: J. Nick Koston <nick@koston.org>
Co-authored-by: J. Nick Koston <nick@home-assistant.io>
This commit is contained in:
Bascht74
2026-10-01 11:36:09 -05:00
committed by GitHub
co-authored by pre-commit-ci-lite[bot] J. Nick Koston J. Nick Koston
parent 54f6fb0eb2
commit eb423be9d2
8 changed files with 527 additions and 0 deletions
+205
View File
@@ -0,0 +1,205 @@
#include "ipv4_resolve.h"
#if defined(USE_SOCKET_IMPL_LWIP_TCP) || defined(USE_SOCKET_IMPL_LWIP_SOCKETS) || defined(USE_SOCKET_IMPL_BSD_SOCKETS)
#include "socket.h"
#include "esphome/core/helpers.h"
#include "esphome/core/log.h"
#include <cstring>
#if !defined(USE_HOST) && !defined(USE_ZEPHYR)
#include "lwip/dns.h"
#else
#include <netdb.h>
#endif
namespace esphome::socket {
static const char *const TAG = "socket";
bool Ipv4Resolve::consume_failure() {
#if defined(IPV4_RESOLVE_ATOMIC_STATE)
uint8_t expected = STATE_FAILED;
return this->state_word_.compare_exchange_strong(expected, STATE_IDLE);
#else
if (this->state_word_ != STATE_FAILED) {
return false;
}
this->state_word_ = STATE_IDLE;
return true;
#endif
}
void Ipv4Resolve::forget() {
#if defined(IPV4_RESOLVE_VOLATILE_EPOCH)
// The generation moves first; RESOLVING stays set until the callback
// drops the stale result, so start() cannot replace that lookup.
this->epoch_ = this->epoch_ + 1;
this->addr_word_ = 0;
if (this->state_word_ != STATE_RESOLVING) {
this->state_word_ = STATE_IDLE;
}
#elif defined(IPV4_RESOLVE_ATOMIC_STATE)
// PUBLISHING stays, so start() cannot queue a second callback
// while this one is storing the address.
uint8_t expected = STATE_RESOLVING;
if (this->state_word_.compare_exchange_strong(expected, STATE_IDLE)) {
this->set_addr_(0);
return;
}
if (this->state_() == STATE_PUBLISHING) {
return;
}
this->set_state_(STATE_IDLE);
this->set_addr_(0);
#else
this->set_state_(STATE_IDLE);
this->set_addr_(0);
#endif
}
socklen_t Ipv4Resolve::to_sockaddr(struct sockaddr *dest, socklen_t destlen, uint16_t port) const {
if (this->state_() != STATE_RESOLVED || destlen < sizeof(sockaddr_in)) {
return 0;
}
auto *in = reinterpret_cast<sockaddr_in *>(dest);
memset(in, 0, sizeof(sockaddr_in));
in->sin_family = AF_INET;
in->sin_port = htons(port);
in->sin_addr.s_addr = this->addr_();
return sizeof(sockaddr_in);
}
#if !defined(USE_HOST) && !defined(USE_ZEPHYR)
#if defined(IPV4_RESOLVE_VOLATILE_EPOCH)
bool Ipv4Resolve::drop_stale_(uint32_t expected) {
if (this->epoch_ == expected) {
return false;
}
this->state_word_ = STATE_IDLE;
this->addr_word_ = 0;
return true;
}
#endif
void Ipv4Resolve::dns_found(const char *name, const ip_addr_t *addr, void *arg) {
auto *self = static_cast<Ipv4Resolve *>(arg);
#if defined(IPV4_RESOLVE_VOLATILE_EPOCH)
const uint32_t expected = self->pending_epoch_;
if (self->drop_stale_(expected)) {
return;
}
#endif
if (addr != nullptr && IP_IS_V4(addr)) {
#if defined(IPV4_RESOLVE_ATOMIC_STATE)
// Only the callback that wins the exchange owns the lookup.
uint8_t expected = STATE_RESOLVING;
if (!self->state_word_.compare_exchange_strong(expected, STATE_PUBLISHING)) {
return;
}
self->set_addr_(ip4_addr_get_u32(ip_2_ip4(addr)));
expected = STATE_PUBLISHING;
// Lost ownership; the leftover address is gated by to_sockaddr()'s state check.
if (!self->state_word_.compare_exchange_strong(expected, STATE_RESOLVED)) {
return;
}
#elif defined(IPV4_RESOLVE_VOLATILE_EPOCH)
self->addr_word_ = ip4_addr_get_u32(ip_2_ip4(addr));
self->state_word_ = STATE_RESOLVED;
if (self->drop_stale_(expected)) {
return;
}
#else
// Do not publish over a newer start().
if (self->state_() != STATE_RESOLVING) {
return;
}
self->set_addr_(ip4_addr_get_u32(ip_2_ip4(addr)));
self->set_state_(STATE_RESOLVED);
#endif
} else {
ESP_LOGW(self->tag_ != nullptr ? self->tag_ : TAG, "DNS failed for %s", name);
#if defined(IPV4_RESOLVE_ATOMIC_STATE)
uint8_t expected = STATE_RESOLVING;
self->state_word_.compare_exchange_strong(expected, STATE_FAILED);
#elif defined(IPV4_RESOLVE_VOLATILE_EPOCH)
self->state_word_ = STATE_FAILED;
if (self->drop_stale_(expected)) {
return;
}
#else
if (self->state_() != STATE_RESOLVING) {
return;
}
self->set_state_(STATE_FAILED);
#endif
}
}
#endif
void Ipv4Resolve::start(const char *host, uint16_t port, const char *tag) {
const uint8_t state = this->state_();
if (state == STATE_RESOLVED || state == STATE_RESOLVING || state == STATE_PUBLISHING) {
return;
}
this->set_state_(STATE_IDLE);
this->tag_ = tag;
struct sockaddr_storage literal;
if (set_sockaddr(reinterpret_cast<struct sockaddr *>(&literal), sizeof(literal), host, port) != 0) {
if (literal.ss_family == AF_INET) {
auto *in = reinterpret_cast<sockaddr_in *>(&literal);
this->set_addr_(in->sin_addr.s_addr);
this->set_state_(STATE_RESOLVED);
return;
}
this->set_state_(STATE_FAILED);
ESP_LOGW(tag, "Not an IPv4 address: %s", host);
return;
}
#if !defined(USE_HOST) && !defined(USE_ZEPHYR)
ip_addr_t cached;
err_t err;
{
LwIPLock lock;
#if defined(IPV4_RESOLVE_VOLATILE_EPOCH)
this->pending_epoch_ = this->epoch_;
#endif
this->set_state_(STATE_RESOLVING);
err = dns_gethostbyname_addrtype(host, &cached, &Ipv4Resolve::dns_found, this, LWIP_DNS_ADDRTYPE_IPV4);
if (err != ERR_INPROGRESS && this->state_() == STATE_RESOLVING) {
this->set_state_(STATE_IDLE);
}
}
if (err == ERR_OK && IP_IS_V4(&cached)) {
this->set_addr_(ip4_addr_get_u32(ip_2_ip4(&cached)));
this->set_state_(STATE_RESOLVED);
return;
}
if (err == ERR_INPROGRESS || this->state_() == STATE_RESOLVED || this->state_() == STATE_PUBLISHING) {
return;
}
#else
struct addrinfo hints {};
hints.ai_family = AF_INET;
hints.ai_socktype = SOCK_STREAM;
struct addrinfo *res = nullptr;
if (getaddrinfo(host, nullptr, &hints, &res) == 0 && res != nullptr) {
auto *in = reinterpret_cast<struct sockaddr_in *>(res->ai_addr);
if (res->ai_family == AF_INET) {
this->set_addr_(in->sin_addr.s_addr);
this->set_state_(STATE_RESOLVED);
}
freeaddrinfo(res);
if (this->ready()) {
return;
}
}
#endif
this->set_state_(STATE_FAILED);
ESP_LOGW(tag, "Could not resolve %s", host);
}
} // namespace esphome::socket
#endif
+103
View File
@@ -0,0 +1,103 @@
#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 <cstdint>
// MULTI_ATOMICS: one atomic state word, races closed by compare_exchange.
// SINGLE: volatile state, the callback never runs beside loop().
// MULTI_NO_ATOMICS: volatile state plus a generation; BK72xx has no
// compare_exchange and the DNS callback runs on the tcpip thread.
#if defined(ESPHOME_THREAD_MULTI_NO_ATOMICS)
#define IPV4_RESOLVE_VOLATILE_EPOCH
#elif defined(ESPHOME_THREAD_SINGLE)
#define IPV4_RESOLVE_VOLATILE
#else
#define IPV4_RESOLVE_ATOMIC_STATE
#include <atomic>
#endif
#if !defined(USE_HOST) && !defined(USE_ZEPHYR)
#include "lwip/ip_addr.h"
#endif
namespace esphome::socket {
/// One IPv4 literal or hostname. Must outlive a pending lookup.
class Ipv4Resolve {
public:
static constexpr uint8_t STATE_IDLE = 0;
static constexpr uint8_t STATE_RESOLVING = 1;
static constexpr uint8_t STATE_RESOLVED = 2;
static constexpr uint8_t STATE_FAILED = 3;
// The callback holds this between winning the lookup and storing the address.
static constexpr uint8_t STATE_PUBLISHING = 4;
/// Drop the stored address so the next start() resolves again.
/// A result already publishing may still land, so ready() can be true
/// right after this; after changing hosts, forget() until ready() is false.
void forget();
/// Drop a failed lookup so the next start() tries again.
bool consume_failure();
bool ready() const { return this->state_() == STATE_RESOLVED; }
/// Write the stored address into dest. Returns 0 until ready() is true.
socklen_t to_sockaddr(struct sockaddr *dest, socklen_t destlen, uint16_t port) const;
/// Resolve host; tag names the failure log. On host and Zephyr this
/// blocks in getaddrinfo().
void start(const char *host, uint16_t port, const char *tag);
private:
#if !defined(USE_HOST) && !defined(USE_ZEPHYR)
static void dns_found(const char *name, const ip_addr_t *addr, void *arg);
#if defined(IPV4_RESOLVE_VOLATILE_EPOCH)
bool drop_stale_(uint32_t expected);
#endif
#endif
uint8_t state_() const {
#if defined(IPV4_RESOLVE_ATOMIC_STATE)
return this->state_word_.load();
#else
return this->state_word_;
#endif
}
uint32_t addr_() const {
#if defined(IPV4_RESOLVE_ATOMIC_STATE)
return this->addr_word_.load();
#else
return this->addr_word_;
#endif
}
void set_state_(uint8_t state) {
#if defined(IPV4_RESOLVE_ATOMIC_STATE)
this->state_word_.store(state);
#else
this->state_word_ = state;
#endif
}
void set_addr_(uint32_t addr) {
#if defined(IPV4_RESOLVE_ATOMIC_STATE)
this->addr_word_.store(addr);
#else
this->addr_word_ = addr;
#endif
}
const char *tag_{nullptr};
#if defined(IPV4_RESOLVE_ATOMIC_STATE)
std::atomic<uint32_t> addr_word_{0};
std::atomic<uint8_t> state_word_{STATE_IDLE};
#elif defined(IPV4_RESOLVE_VOLATILE_EPOCH)
volatile uint32_t addr_word_{0};
volatile uint32_t epoch_{0};
volatile uint32_t pending_epoch_{0};
volatile uint8_t state_word_{STATE_IDLE};
#else
volatile uint32_t addr_word_{0};
volatile uint8_t state_word_{STATE_IDLE};
#endif
};
} // namespace esphome::socket
#endif
@@ -0,0 +1,71 @@
#include <gtest/gtest.h>
#include <cstring>
#include "esphome/components/socket/ipv4_resolve.h"
#ifdef USE_HOST
namespace esphome::socket::testing {
TEST(Ipv4Resolve, LiteralIsReadyWithoutDns) {
Ipv4Resolve lookup;
lookup.start("192.168.1.1", 1, "test");
EXPECT_TRUE(lookup.ready());
struct sockaddr_storage addr {};
socklen_t len = lookup.to_sockaddr(reinterpret_cast<struct sockaddr *>(&addr), sizeof(addr), 6053);
ASSERT_EQ(len, sizeof(sockaddr_in));
auto *in = reinterpret_cast<sockaddr_in *>(&addr);
EXPECT_EQ(in->sin_family, AF_INET);
EXPECT_EQ(ntohs(in->sin_port), 6053);
EXPECT_EQ(in->sin_addr.s_addr, htonl(0xC0A80101));
}
TEST(Ipv4Resolve, ForgetDropsTheLiteral) {
Ipv4Resolve lookup;
lookup.start("10.0.0.5", 80, "test");
ASSERT_TRUE(lookup.ready());
lookup.forget();
EXPECT_FALSE(lookup.ready());
struct sockaddr_storage addr {};
EXPECT_EQ(lookup.to_sockaddr(reinterpret_cast<struct sockaddr *>(&addr), sizeof(addr), 80), 0u);
}
TEST(Ipv4Resolve, Ipv6LiteralIsRejected) {
Ipv4Resolve lookup;
lookup.start("::1", 443, "test");
EXPECT_FALSE(lookup.ready());
EXPECT_TRUE(lookup.consume_failure());
}
TEST(Ipv4Resolve, ShortBufferWritesNothing) {
Ipv4Resolve lookup;
lookup.start("192.0.2.10", 502, "test");
ASSERT_TRUE(lookup.ready());
struct sockaddr_in addr {};
EXPECT_EQ(lookup.to_sockaddr(reinterpret_cast<struct sockaddr *>(&addr), sizeof(addr) - 1, 502), 0u);
}
TEST(Ipv4Resolve, RetryAfterFailureKeepsTheAddress) {
Ipv4Resolve lookup;
lookup.start("::1", 443, "test");
EXPECT_FALSE(lookup.ready());
lookup.start("192.168.1.1", 1, "test");
EXPECT_TRUE(lookup.ready());
EXPECT_FALSE(lookup.consume_failure());
EXPECT_TRUE(lookup.ready());
}
TEST(Ipv4Resolve, ForgetDropsAStaleFailure) {
Ipv4Resolve lookup;
lookup.start("::1", 443, "test");
lookup.forget();
EXPECT_FALSE(lookup.consume_failure());
EXPECT_FALSE(lookup.ready());
}
} // namespace esphome::socket::testing
#endif
@@ -0,0 +1,22 @@
import esphome.codegen as cg
import esphome.config_validation as cv
from esphome.const import CONF_ID
from esphome.types import ConfigType
AUTO_LOAD = ["socket"]
ipv4_resolve_test_component_ns = cg.esphome_ns.namespace("ipv4_resolve_test_component")
Ipv4ResolveTestComponent = ipv4_resolve_test_component_ns.class_(
"Ipv4ResolveTestComponent", cg.Component
)
CONFIG_SCHEMA = cv.Schema(
{
cv.GenerateID(): cv.declare_id(Ipv4ResolveTestComponent),
}
).extend(cv.COMPONENT_SCHEMA)
async def to_code(config: ConfigType) -> None:
var = cg.new_Pvariable(config[CONF_ID])
await cg.register_component(var, config)
@@ -0,0 +1,46 @@
#include "ipv4_resolve_test_component.h"
#include "esphome/components/socket/ipv4_resolve.h"
#include "esphome/core/log.h"
#include <cstring>
namespace esphome::ipv4_resolve_test_component {
static const char *const TAG = "ipv4_resolve_test";
static bool check_sockaddr(socket::Ipv4Resolve &lookup, uint16_t port, uint32_t expected) {
struct sockaddr_storage addr {};
socklen_t len = lookup.to_sockaddr(reinterpret_cast<struct sockaddr *>(&addr), sizeof(addr), port);
if (len != sizeof(sockaddr_in)) {
return false;
}
auto *in = reinterpret_cast<sockaddr_in *>(&addr);
return in->sin_family == AF_INET && ntohs(in->sin_port) == port && in->sin_addr.s_addr == htonl(expected);
}
void Ipv4ResolveTestComponent::setup() {
ESP_LOGI(TAG, "IPv4 resolve test starting");
socket::Ipv4Resolve lookup;
struct sockaddr_storage addr {};
lookup.start("192.168.1.1", 1, TAG);
bool ok = lookup.ready() && check_sockaddr(lookup, 6053, 0xC0A80101);
ESP_LOGI(TAG, "Literal resolve: %s", ok ? LOG_STR_LITERAL("PASSED") : LOG_STR_LITERAL("FAILED"));
lookup.forget();
ok = !lookup.ready() && lookup.to_sockaddr(reinterpret_cast<struct sockaddr *>(&addr), sizeof(addr), 80) == 0;
ESP_LOGI(TAG, "Forget drops address: %s", ok ? LOG_STR_LITERAL("PASSED") : LOG_STR_LITERAL("FAILED"));
lookup.start("::1", 443, TAG);
ok = !lookup.ready() && lookup.consume_failure();
ESP_LOGI(TAG, "IPv6 literal rejected: %s", ok ? LOG_STR_LITERAL("PASSED") : LOG_STR_LITERAL("FAILED"));
lookup.start("localhost", 6053, TAG);
ok = lookup.ready() && check_sockaddr(lookup, 6053, 0x7F000001);
ESP_LOGI(TAG, "Hostname resolve: %s", ok ? LOG_STR_LITERAL("PASSED") : LOG_STR_LITERAL("FAILED"));
ESP_LOGI(TAG, "IPv4 resolve test complete");
}
} // namespace esphome::ipv4_resolve_test_component
@@ -0,0 +1,12 @@
#pragma once
#include "esphome/core/component.h"
namespace esphome::ipv4_resolve_test_component {
class Ipv4ResolveTestComponent : public Component {
public:
void setup() override;
};
} // namespace esphome::ipv4_resolve_test_component
@@ -0,0 +1,17 @@
esphome:
name: socket-ipv4-resolve-test
host:
api:
logger:
level: INFO
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
components: [ipv4_resolve_test_component]
ipv4_resolve_test_component:
@@ -0,0 +1,51 @@
"""Integration test for the socket Ipv4Resolve helper on host."""
from __future__ import annotations
import asyncio
import pytest
from .types import APIClientConnectedFactory, RunCompiledFunction
CHECKS = (
"Literal resolve",
"Forget drops address",
"IPv6 literal rejected",
"Hostname resolve",
)
@pytest.mark.asyncio
async def test_socket_ipv4_resolve(
yaml_config: str,
run_compiled: RunCompiledFunction,
api_client_connected: APIClientConnectedFactory,
) -> None:
"""Exercise Ipv4Resolve literals, forget, failure, and getaddrinfo on host."""
test_complete = asyncio.Event()
results: dict[str, bool] = {}
def on_log_line(line: str) -> None:
if "IPv4 resolve test complete" in line:
test_complete.set()
return
for check in CHECKS:
if f"{check}:" in line:
results[check] = "PASSED" in line
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-ipv4-resolve-test"
try:
await asyncio.wait_for(test_complete.wait(), timeout=10.0)
except TimeoutError:
pytest.fail("IPv4 resolve test timed out")
for check in CHECKS:
assert results.get(check), f"{check} check failed or never ran"