Files
esphome/esphome/components/socket/__init__.py
T

279 lines
9.4 KiB
Python

from collections.abc import Callable, MutableMapping
from dataclasses import dataclass
from enum import StrEnum
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
from esphome.types import ConfigType
_LOGGER = logging.getLogger(__name__)
CODEOWNERS = ["@esphome/core"]
socket_ns = cg.esphome_ns.namespace("socket")
Ipv4AllowEntry = socket_ns.struct("Ipv4AllowEntry")
CONF_IMPLEMENTATION = "implementation"
IMPLEMENTATION_LWIP_TCP = "lwip_tcp"
IMPLEMENTATION_LWIP_SOCKETS = "lwip_sockets"
IMPLEMENTATION_BSD_SOCKETS = "bsd_sockets"
# Socket tracking infrastructure
# Components register their socket needs and platforms read this to configure appropriately
KEY_SOCKET_CONSUMERS_TCP = "socket_consumers_tcp"
KEY_SOCKET_CONSUMERS_UDP = "socket_consumers_udp"
KEY_SOCKET_CONSUMERS_TCP_LISTEN = "socket_consumers_tcp_listen"
# Recommended minimum socket counts.
# Platforms should apply these (or their own) on top of get_socket_counts().
# These cover minimal configs (e.g. api-only without web_server).
# When web_server is present, its 5 registered sockets push past the TCP minimum.
MIN_TCP_SOCKETS = 8
MIN_UDP_SOCKETS = 6
# Minimum listening sockets — at least api + ota baseline.
MIN_TCP_LISTEN_SOCKETS = 2
class SocketType(StrEnum):
TCP = "tcp"
UDP = "udp"
TCP_LISTEN = "tcp_listen"
_SOCKET_TYPE_KEYS = {
SocketType.TCP: KEY_SOCKET_CONSUMERS_TCP,
SocketType.UDP: KEY_SOCKET_CONSUMERS_UDP,
SocketType.TCP_LISTEN: KEY_SOCKET_CONSUMERS_TCP_LISTEN,
}
def consume_sockets(
value: int, consumer: str, socket_type: SocketType = SocketType.TCP
) -> Callable[[MutableMapping], MutableMapping]:
"""Register socket usage for a component.
Args:
value: Number of sockets needed by the component
consumer: Name of the component consuming the sockets
socket_type: Type of socket (SocketType.TCP, SocketType.UDP, or SocketType.TCP_LISTEN)
Returns:
A validator function that records the socket usage
"""
typed_key = _SOCKET_TYPE_KEYS[socket_type]
def _consume_sockets(config: MutableMapping) -> MutableMapping:
consumers: dict[str, int] = CORE.data.setdefault(typed_key, {})
consumers[consumer] = consumers.get(consumer, 0) + value
return config
return _consume_sockets
def _format_consumers(consumers: dict[str, int]) -> str:
"""Format consumer dict as 'name=count, ...' or 'none'."""
if not consumers:
return "none"
return ", ".join(f"{name}={count}" for name, count in sorted(consumers.items()))
@dataclass(frozen=True)
class SocketCounts:
"""Socket counts and component details for platform configuration."""
tcp: int
udp: int
tcp_listen: int
tcp_details: str
udp_details: str
tcp_listen_details: str
def get_socket_counts() -> SocketCounts:
"""Return socket counts and component details for platform configuration.
Platforms call this during code generation to configure lwIP socket limits.
All components will have registered their needs by then.
Platforms should apply their own minimums on top of these values.
"""
tcp_consumers = CORE.data.get(KEY_SOCKET_CONSUMERS_TCP, {})
udp_consumers = CORE.data.get(KEY_SOCKET_CONSUMERS_UDP, {})
tcp_listen_consumers = CORE.data.get(KEY_SOCKET_CONSUMERS_TCP_LISTEN, {})
tcp = sum(tcp_consumers.values())
udp = sum(udp_consumers.values())
tcp_listen = sum(tcp_listen_consumers.values())
tcp_details = _format_consumers(tcp_consumers)
udp_details = _format_consumers(udp_consumers)
tcp_listen_details = _format_consumers(tcp_listen_consumers)
_LOGGER.debug(
"Socket counts: TCP=%d (%s), UDP=%d (%s), TCP_LISTEN=%d (%s)",
tcp,
tcp_details,
udp,
udp_details,
tcp_listen,
tcp_listen_details,
)
return SocketCounts(
tcp, udp, tcp_listen, tcp_details, udp_details, tcp_listen_details
)
def require_wake_loop_threadsafe() -> None:
"""Deprecated: wake loop support is now always available on all platforms.
This function adds backward-compatible defines so external components
that check #ifdef USE_WAKE_LOOP_THREADSAFE / USE_SOCKET_SELECT_SUPPORT
continue to compile. Remove before 2026.12.0.
"""
# Remove before 2026.12.0
_LOGGER.warning(
"require_wake_loop_threadsafe() is deprecated and no longer needed. "
"Wake loop support is now always available. Remove this call and any "
"#ifdef USE_SOCKET_SELECT_SUPPORT / USE_WAKE_LOOP_THREADSAFE guards. "
"This will be removed in 2026.12.0."
)
# Add deprecated defines for backward compat with external component C++ code
cg.add_define("USE_WAKE_LOOP_THREADSAFE")
cg.add_define("USE_SOCKET_SELECT_SUPPORT")
# For an Ipv4Allow config option; a sanity cap on the list length.
IPV4_ALLOW_SCHEMA = cv.All(cv.ensure_list(cv.ipv4network), cv.Length(max=255))
_HOST = cv.Any(cv.domain, cv.hostname)
def ipv4_host(value: object) -> str:
"""Validate an IPv4 address or a hostname; the resolver behind it is IPv4 only."""
value = cv.string(value)
try:
cv.ipv6address(value)
except cv.Invalid:
return _HOST(value)
raise cv.Invalid(
"IPv6 addresses are not supported, use an IPv4 address or a hostname"
)
def _network_order(addr: IPv4Address) -> int:
"""The s_addr value for addr on the little endian targets."""
return int.from_bytes(addr.packed, "little")
def add_ipv4_allow(
setter: cg.MockObj, networks: list[IPv4Network], owner_id: ID | str
) -> None:
"""Emit a flash array for validated IPV4_ALLOW_SCHEMA entries and wire it to setter.
PROGMEM on esp8266. Emits nothing for an empty list.
"""
if not networks:
return
cg.add_define("USE_SOCKET_IPV4_ALLOW")
entries = [
cg.StructInitializer(
Ipv4AllowEntry,
("addr", _network_order(net.network_address)),
("mask", _network_order(net.netmask)),
)
for net in networks
]
arr_id = ID(f"{owner_id}_ipv4_allow", is_declaration=True, type=Ipv4AllowEntry)
arr = cg.progmem_array(arr_id, cg.ArrayInitializer(*entries))
cg.add(setter(arr, len(entries)))
def require_ipv4_resolve() -> None:
"""Compile the shared IPv4 lookup; call from a consumer's to_code."""
cg.add_define("USE_SOCKET_IPV4_RESOLVE")
def require_tcp_client_link() -> None:
"""Compile the reconnecting TCP client link; call from a consumer's to_code."""
require_ipv4_resolve()
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(
CONF_IMPLEMENTATION,
esp8266=IMPLEMENTATION_LWIP_TCP,
esp32=IMPLEMENTATION_BSD_SOCKETS,
rp2=IMPLEMENTATION_LWIP_TCP,
bk72xx=IMPLEMENTATION_LWIP_SOCKETS,
ln882x=IMPLEMENTATION_LWIP_SOCKETS,
rtl87xx=IMPLEMENTATION_LWIP_SOCKETS,
host=IMPLEMENTATION_BSD_SOCKETS,
nrf52=IMPLEMENTATION_BSD_SOCKETS,
): cv.one_of(
IMPLEMENTATION_LWIP_TCP,
IMPLEMENTATION_LWIP_SOCKETS,
IMPLEMENTATION_BSD_SOCKETS,
lower=True,
space="_",
),
}
)
async def to_code(config: ConfigType) -> None:
impl = config[CONF_IMPLEMENTATION]
if impl == IMPLEMENTATION_LWIP_TCP:
cg.add_define("USE_SOCKET_IMPL_LWIP_TCP")
elif impl == IMPLEMENTATION_LWIP_SOCKETS:
cg.add_define("USE_SOCKET_IMPL_LWIP_SOCKETS")
elif impl == IMPLEMENTATION_BSD_SOCKETS:
cg.add_define("USE_SOCKET_IMPL_BSD_SOCKETS")
if CORE.using_zephyr:
from esphome.components.zephyr import zephyr_add_prj_conf
zephyr_add_prj_conf("NET_SOCKETS", True)
zephyr_add_prj_conf("POSIX_API", True)
# ESP32 and LibreTiny both have LwIP >= 2.1.3 with lwip_socket_dbg_get_socket()
# and FreeRTOS task notifications — enable fast select to bypass lwip_select().
# Only when not using lwip_tcp, which does not provide select() support.
if (CORE.is_esp32 or CORE.is_libretiny) and impl != IMPLEMENTATION_LWIP_TCP:
cg.add_build_flag("-DUSE_LWIP_FAST_SELECT")
# Each implementation file is fully #ifdef'd on the define set in to_code
# for the selected implementation. The helper files compile only for
# consumers that called the matching require_ function.
FILTER_SOURCE_FILES = filter_source_files_from_defines(
{
"lwip_raw_tcp_impl.cpp": "USE_SOCKET_IMPL_LWIP_TCP",
"bsd_sockets_impl.cpp": "USE_SOCKET_IMPL_BSD_SOCKETS",
"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",
}
)