1
0
mirror of https://github.com/elisspace/core.git synced 2026-08-30 16:12:01 +00:00

Support multiple sites on the webserver

aiohttp already supports multiple TCPSites for a single server

This allows Home Assistant to listen on multiple ports with
different ssl configuration per port so that devices that
need a less secure configuration can be isolated to a single
port.

The motivation for this is that I did not want to downgrade
my entire ssl configuration for a single device that needs
a less secure configuration, or to support an ONVIF camera
that does not do SSL
This commit is contained in:
J. Nick Koston
2023-05-19 13:01:34 -05:00
parent 3cdcfb1654
commit d8fa06de67

View File

@@ -2,6 +2,7 @@
from __future__ import annotations
import asyncio
from dataclasses import dataclass
import datetime
from ipaddress import IPv4Network, IPv6Network, ip_network
import logging
@@ -87,18 +88,38 @@ STORAGE_KEY: Final = DOMAIN
STORAGE_VERSION: Final = 1
SAVE_DELAY: Final = 180
CONF_SERVERS = "servers"
SERVER_SCHEMA = {
vol.Optional(CONF_SERVER_HOST): vol.All(
cv.ensure_list, vol.Length(min=1), [cv.string]
),
vol.Optional(CONF_SERVER_PORT, default=SERVER_PORT): cv.port,
vol.Optional(CONF_SSL_CERTIFICATE): cv.isfile,
vol.Optional(CONF_SSL_PEER_CERTIFICATE): cv.isfile,
vol.Optional(CONF_SSL_KEY): cv.isfile,
vol.Optional(CONF_SSL_PROFILE, default=SSL_MODERN): vol.In(
[SSL_INTERMEDIATE, SSL_MODERN]
),
}
def _has_all_unique_ports(servers: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Validate that each http service has a unique port."""
ports = [list[CONF_SERVER_PORT] for list in servers]
vol.Schema(vol.Unique())(ports)
return servers
HTTP_SCHEMA: Final = vol.All(
cv.deprecated(CONF_BASE_URL),
vol.Schema(
{
vol.Optional(CONF_SERVER_HOST): vol.All(
cv.ensure_list, vol.Length(min=1), [cv.string]
),
vol.Optional(CONF_SERVER_PORT, default=SERVER_PORT): cv.port,
**SERVER_SCHEMA,
vol.Optional(CONF_BASE_URL): cv.string,
vol.Optional(CONF_SSL_CERTIFICATE): cv.isfile,
vol.Optional(CONF_SSL_PEER_CERTIFICATE): cv.isfile,
vol.Optional(CONF_SSL_KEY): cv.isfile,
vol.Optional(CONF_SERVERS): vol.All(
cv.ensure_list, [vol.Schema(SERVER_SCHEMA)], _has_all_unique_ports
),
vol.Optional(CONF_CORS_ORIGINS, default=DEFAULT_CORS): vol.All(
cv.ensure_list, [cv.string]
),
@@ -110,13 +131,11 @@ HTTP_SCHEMA: Final = vol.All(
CONF_LOGIN_ATTEMPTS_THRESHOLD, default=NO_LOGIN_ATTEMPT_THRESHOLD
): vol.Any(cv.positive_int, NO_LOGIN_ATTEMPT_THRESHOLD),
vol.Optional(CONF_IP_BAN_ENABLED, default=True): cv.boolean,
vol.Optional(CONF_SSL_PROFILE, default=SSL_MODERN): vol.In(
[SSL_INTERMEDIATE, SSL_MODERN]
),
}
),
)
CONFIG_SCHEMA: Final = vol.Schema({DOMAIN: HTTP_SCHEMA}, extra=vol.ALLOW_EXTRA)
@@ -161,6 +180,33 @@ class ApiConfig:
self.use_ssl = use_ssl
@dataclass
class SiteServerConfig:
"""Configuration for a single TCPSite."""
server_host: list[str]
server_port: int
ssl_certificate: str | None
ssl_peer_certificate: str | None
ssl_key: str | None
ssl_profile: str
ssl_context: ssl.SSLContext | None = None
def _create_site_server_config_from_dict(
conf: ConfData | dict[str, Any]
) -> SiteServerConfig:
"""Create a SiteServerConfig from a dict."""
return SiteServerConfig(
server_host=conf[CONF_SERVER_HOST],
server_port=conf[CONF_SERVER_PORT],
ssl_certificate=conf.get(CONF_SSL_CERTIFICATE),
ssl_peer_certificate=conf.get(CONF_SSL_PEER_CERTIFICATE),
ssl_key=conf.get(CONF_SSL_KEY),
ssl_profile=conf[CONF_SSL_PROFILE],
)
async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool:
"""Set up the HTTP API and debug interface."""
conf: ConfData | None = config.get(DOMAIN)
@@ -168,27 +214,27 @@ async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool:
if conf is None:
conf = cast(ConfData, HTTP_SCHEMA({}))
server_host = conf.get(CONF_SERVER_HOST)
server_port = conf[CONF_SERVER_PORT]
ssl_certificate = conf.get(CONF_SSL_CERTIFICATE)
ssl_peer_certificate = conf.get(CONF_SSL_PEER_CERTIFICATE)
ssl_key = conf.get(CONF_SSL_KEY)
# configuration options that affect all TCPSites
cors_origins = conf[CONF_CORS_ORIGINS]
use_x_forwarded_for = conf.get(CONF_USE_X_FORWARDED_FOR, False)
trusted_proxies = conf.get(CONF_TRUSTED_PROXIES) or []
is_ban_enabled = conf[CONF_IP_BAN_ENABLED]
login_threshold = conf[CONF_LOGIN_ATTEMPTS_THRESHOLD]
ssl_profile = conf[CONF_SSL_PROFILE]
site_configs: list[SiteServerConfig] = []
if CONF_SERVERS in config[DOMAIN]:
site_configs = [
_create_site_server_config_from_dict(conf)
for conf in config[DOMAIN][CONF_SERVERS]
]
else:
site_configs = [_create_site_server_config_from_dict(conf)]
server = HomeAssistantHTTP(
hass,
server_host=server_host,
server_port=server_port,
ssl_certificate=ssl_certificate,
ssl_peer_certificate=ssl_peer_certificate,
ssl_key=ssl_key,
site_configs=site_configs,
trusted_proxies=trusted_proxies,
ssl_profile=ssl_profile,
)
await server.async_initialize(
cors_origins=cors_origins,
@@ -212,16 +258,19 @@ async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool:
async_when_setup_or_start(hass, "frontend", start_server)
hass.http = server
primary_server_conf = site_configs[0]
local_ip = await async_get_source_ip(hass)
host = local_ip
if server_host is not None:
# Assume the first server host name provided as API host
host = server_host[0]
if primary_server_host := primary_server_conf.server_host:
primary_host = primary_server_host[0]
else:
primary_host = local_ip
hass.config.api = ApiConfig(
local_ip, host, server_port, ssl_certificate is not None
local_ip,
primary_host,
primary_server_conf.server_port,
primary_server_conf.ssl_certificate is not None,
)
return True
@@ -286,13 +335,8 @@ class HomeAssistantHTTP:
def __init__(
self,
hass: HomeAssistant,
ssl_certificate: str | None,
ssl_peer_certificate: str | None,
ssl_key: str | None,
server_host: list[str] | None,
server_port: int,
site_configs: list[SiteServerConfig],
trusted_proxies: list[IPv4Network | IPv6Network],
ssl_profile: str,
) -> None:
"""Initialize the HTTP Home Assistant server."""
self.app = HomeAssistantApplication(
@@ -304,16 +348,10 @@ class HomeAssistantHTTP:
},
)
self.hass = hass
self.ssl_certificate = ssl_certificate
self.ssl_peer_certificate = ssl_peer_certificate
self.ssl_key = ssl_key
self.server_host = server_host
self.server_port = server_port
self.site_configs = site_configs
self.trusted_proxies = trusted_proxies
self.ssl_profile = ssl_profile
self.runner: web.AppRunner | None = None
self.site: HomeAssistantTCPSite | None = None
self.context: ssl.SSLContext | None = None
self.sites: list[HomeAssistantTCPSite] = []
async def async_initialize(
self,
@@ -341,10 +379,8 @@ class HomeAssistantHTTP:
setup_cors(self.app, cors_origins)
if self.ssl_certificate:
self.context = await self.hass.async_add_executor_job(
self._create_ssl_context
)
if any(site.ssl_certificate for site in self.site_configs):
await self.hass.async_add_executor_job(self._create_ssl_contexts)
def register_view(self, view: HomeAssistantView | type[HomeAssistantView]) -> None:
"""Register a view with the WSGI server.
@@ -417,53 +453,55 @@ class HomeAssistantHTTP:
self.app.router.add_route("GET", url_path, serve_file)
)
def _create_ssl_context(self) -> ssl.SSLContext | None:
context: ssl.SSLContext | None = None
assert self.ssl_certificate is not None
try:
if self.ssl_profile == SSL_INTERMEDIATE:
context = ssl_util.server_context_intermediate()
else:
context = ssl_util.server_context_modern()
context.load_cert_chain(self.ssl_certificate, self.ssl_key)
except OSError as error:
if not self.hass.config.safe_mode:
raise HomeAssistantError(
f"Could not use SSL certificate from {self.ssl_certificate}:"
f" {error}"
) from error
_LOGGER.error(
"Could not read SSL certificate from %s: %s",
self.ssl_certificate,
error,
)
def _create_ssl_contexts(self) -> None:
for site in self.site_configs:
context: ssl.SSLContext | None = None
assert site.ssl_certificate is not None
try:
context = self._create_emergency_ssl_context()
except OSError as error2:
if site.ssl_profile == SSL_INTERMEDIATE:
context = ssl_util.server_context_intermediate()
else:
context = ssl_util.server_context_modern()
context.load_cert_chain(site.ssl_certificate, site.ssl_key)
except OSError as error:
if not self.hass.config.safe_mode:
raise HomeAssistantError(
f"Could not use SSL certificate from {site.ssl_certificate}:"
f" {error}"
) from error
_LOGGER.error(
"Could not create an emergency self signed ssl certificate: %s",
error2,
"Could not read SSL certificate from %s: %s",
site.ssl_certificate,
error,
)
context = None
else:
_LOGGER.critical(
"Home Assistant is running in safe mode with an emergency self"
" signed ssl certificate because the configured SSL certificate was"
" not usable"
)
return context
try:
context = self._create_emergency_ssl_context()
except OSError as os_error:
_LOGGER.error(
"Could not create an emergency self signed ssl certificate: %s",
os_error,
)
continue
else:
_LOGGER.critical(
"Home Assistant is running in safe mode with an emergency self"
" signed ssl certificate because the configured SSL certificate was"
" not usable"
)
site.ssl_context = context
continue
if self.ssl_peer_certificate:
if context is None:
raise HomeAssistantError(
"Failed to create ssl context, no fallback available because a peer"
" certificate is required."
)
if site.ssl_peer_certificate:
if context is None:
raise HomeAssistantError(
"Failed to create ssl context, no fallback available because a peer"
" certificate is required."
)
context.verify_mode = ssl.CERT_REQUIRED
context.load_verify_locations(self.ssl_peer_certificate)
context.verify_mode = ssl.CERT_REQUIRED
context.load_verify_locations(site.ssl_peer_certificate)
return context
site.ssl_context = context
def _create_emergency_ssl_context(self) -> ssl.SSLContext:
"""Create an emergency ssl certificate so we can still startup."""
@@ -527,22 +565,38 @@ class HomeAssistantHTTP:
)
await self.runner.setup()
self.site = HomeAssistantTCPSite(
self.runner, self.server_host, self.server_port, ssl_context=self.context
)
try:
await self.site.start()
except OSError as error:
_LOGGER.error(
"Failed to create HTTP server at port %d: %s", self.server_port, error
sites = [
HomeAssistantTCPSite(
self.runner,
site_config.server_host,
site_config.server_port,
ssl_context=site_config.ssl_context,
)
for site_config in self.site_configs
]
_LOGGER.info("Now listening on port %d", self.server_port)
results = await asyncio.gather(
*(site.start() for site in sites), return_exceptions=True
)
for idx, result in enumerate(results):
site_config = self.site_configs[idx]
if isinstance(result, Exception):
_LOGGER.error(
"Failed to create HTTP server at port %d: %s",
site_config.server_port,
result,
)
continue
self.sites.append(sites[idx])
_LOGGER.info("Now listening on port %d", site_config.server_port)
async def stop(self) -> None:
"""Stop the aiohttp server."""
if self.site is not None:
await self.site.stop()
if self.sites:
await asyncio.gather(*[site.stop() for site in self.sites])
if self.runner is not None:
await self.runner.cleanup()