fix: Round 2 review findings — async SSRF, error sanitization, constants

Security:
- Make validate_endpoint async with hass.async_add_executor_job for DNS resolution
- Use _RestrictedIPError instead of fragile string matching for IP check flow
- Add is_multicast/is_unspecified to SSRF IP restriction checks
- Remove resolved private IP from error messages (generic message)
- Remove raw endpoint from __init__.py error log
- Sanitize error messages in metrics: strip URLs, API keys, Gemini key patterns
- Truncate API error response bodies before logging (512 chars)
- Use generic error messages in Gemini exception handlers (no str(e) interpolation)
- Pass api_key as explicit APIClient constructor parameter (not from header)
- Add defensive validation: Gemini provider requires api_key at construction

Code quality:
- Add constants: DEFAULT_INSTANCE_NAME, MIN/MAX_CONTEXT_MESSAGES, MIN/MAX_HISTORY_SIZE
- Replace all hardcoded schema ranges with named constants
- Move datetime import to module level in history.py
- Remove unused full_history/full_history_available keys
- Chain socket.gaierror properly with raise...from
This commit is contained in:
SMKRV
2026-03-12 01:57:49 +03:00
parent c54bfcff3b
commit 31560a8835
8 changed files with 72 additions and 38 deletions
+3 -2
View File
@@ -263,9 +263,9 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
model = config.get(CONF_MODEL, get_default_model(api_provider)) model = config.get(CONF_MODEL, get_default_model(api_provider))
raw_endpoint = config.get(CONF_API_ENDPOINT, get_default_endpoint(api_provider)) raw_endpoint = config.get(CONF_API_ENDPOINT, get_default_endpoint(api_provider))
try: try:
endpoint = validate_endpoint(raw_endpoint) endpoint = await validate_endpoint(hass, raw_endpoint)
except ValueError as err: except ValueError as err:
_LOGGER.error("Invalid API endpoint %s: %s", raw_endpoint, err) _LOGGER.error("Invalid API endpoint: %s", err)
raise ConfigEntryNotReady(f"Invalid API endpoint: {err}") from err raise ConfigEntryNotReady(f"Invalid API endpoint: {err}") from err
# API key can now be updated via options # API key can now be updated via options
api_key = config.get(CONF_API_KEY, entry.data.get(CONF_API_KEY)) api_key = config.get(CONF_API_KEY, entry.data.get(CONF_API_KEY))
@@ -292,6 +292,7 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
api_provider=api_provider, api_provider=api_provider,
model=model, model=model,
api_timeout=api_timeout, api_timeout=api_timeout,
api_key=api_key,
) )
coordinator = HATextAICoordinator( coordinator = HATextAICoordinator(
+11 -7
View File
@@ -39,6 +39,7 @@ class APIClient:
api_provider: str, api_provider: str,
model: str, model: str,
api_timeout: int = DEFAULT_API_TIMEOUT, api_timeout: int = DEFAULT_API_TIMEOUT,
api_key: Optional[str] = None,
) -> None: ) -> None:
"""Initialize API client.""" """Initialize API client."""
self.session = session self.session = session
@@ -48,6 +49,9 @@ class APIClient:
self.model = model self.model = model
self.api_timeout = api_timeout self.api_timeout = api_timeout
self.timeout = ClientTimeout(total=api_timeout) self.timeout = ClientTimeout(total=api_timeout)
self._api_key = api_key
if self.api_provider == API_PROVIDER_GEMINI and not api_key:
raise ValueError("Gemini provider requires api_key parameter")
self._closed = False self._closed = False
async def __aenter__(self): async def __aenter__(self):
@@ -119,7 +123,8 @@ class APIClient:
raise HomeAssistantError("API rate limit exceeded") raise HomeAssistantError("API rate limit exceeded")
# Client/server errors — don't retry # Client/server errors — don't retry
_LOGGER.error("API error (status %d): %s", response.status, error_data) truncated_error = str(error_data)[:512]
_LOGGER.error("API error (status %d): %s", response.status, truncated_error)
raise HomeAssistantError(f"API error: status {response.status}") raise HomeAssistantError(f"API error: status {response.status}")
except asyncio.TimeoutError: except asyncio.TimeoutError:
@@ -368,8 +373,7 @@ class APIClient:
genai = await asyncio.to_thread(import_genai) genai = await asyncio.to_thread(import_genai)
# Extract API key from headers (Bearer token) api_key = self._api_key
api_key = self.headers.get("Authorization", "").replace("Bearer ", "")
def create_client(): def create_client():
if self.endpoint and self.endpoint != "https://generativelanguage.googleapis.com/v1beta": if self.endpoint and self.endpoint != "https://generativelanguage.googleapis.com/v1beta":
@@ -502,11 +506,11 @@ class APIClient:
} }
except ImportError as e: except ImportError as e:
_LOGGER.error(f"Google Gemini library not installed: {str(e)}") _LOGGER.error("Google Gemini library not installed: %s", e)
raise HomeAssistantError(f"Missing dependency: {str(e)}. Please install google-genai.") raise HomeAssistantError("Missing dependency: google-genai. Please install it.")
except Exception as e: except Exception as e:
_LOGGER.error(f"Gemini API error: {str(e)}") _LOGGER.error("Gemini API error: %s", e)
raise HomeAssistantError(f"Gemini API error: {str(e)}") raise HomeAssistantError("Gemini API request failed")
async def shutdown(self) -> None: async def shutdown(self) -> None:
"""Shutdown API client.""" """Shutdown API client."""
+12 -7
View File
@@ -45,8 +45,13 @@ from .const import (
MIN_API_TIMEOUT, MIN_API_TIMEOUT,
MAX_API_TIMEOUT, MAX_API_TIMEOUT,
DEFAULT_NAME_PREFIX, DEFAULT_NAME_PREFIX,
DEFAULT_INSTANCE_NAME,
DEFAULT_MAX_HISTORY, DEFAULT_MAX_HISTORY,
CONF_MAX_HISTORY_SIZE, CONF_MAX_HISTORY_SIZE,
MIN_CONTEXT_MESSAGES,
MAX_CONTEXT_MESSAGES,
MIN_HISTORY_SIZE,
MAX_HISTORY_SIZE,
) )
from homeassistant.util import dt as dt_util from homeassistant.util import dt as dt_util
@@ -91,7 +96,7 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
"""Build provider configuration schema with optional defaults from data.""" """Build provider configuration schema with optional defaults from data."""
defaults = data or {} defaults = data or {}
return vol.Schema({ return vol.Schema({
vol.Required(CONF_NAME, default=defaults.get(CONF_NAME, "my_assistant")): str, vol.Required(CONF_NAME, default=defaults.get(CONF_NAME, DEFAULT_INSTANCE_NAME)): str,
vol.Required(CONF_API_KEY): str, vol.Required(CONF_API_KEY): str,
vol.Required(CONF_MODEL, default=defaults.get(CONF_MODEL, get_default_model(self._provider))): str, vol.Required(CONF_MODEL, default=defaults.get(CONF_MODEL, get_default_model(self._provider))): str,
vol.Required(CONF_API_ENDPOINT, default=defaults.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str, vol.Required(CONF_API_ENDPOINT, default=defaults.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str,
@@ -116,14 +121,14 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
default=defaults.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES) default=defaults.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
): vol.All( ): vol.All(
vol.Coerce(int), vol.Coerce(int),
vol.Range(min=1, max=20) vol.Range(min=MIN_CONTEXT_MESSAGES, max=MAX_CONTEXT_MESSAGES)
), ),
vol.Optional( vol.Optional(
CONF_MAX_HISTORY_SIZE, CONF_MAX_HISTORY_SIZE,
default=defaults.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY) default=defaults.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY)
): vol.All( ): vol.All(
vol.Coerce(int), vol.Coerce(int),
vol.Range(min=1, max=100) vol.Range(min=MIN_HISTORY_SIZE, max=MAX_HISTORY_SIZE)
), ),
}) })
@@ -225,7 +230,7 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
return False return False
try: try:
endpoint = validate_endpoint(user_input[CONF_API_ENDPOINT]) endpoint = await validate_endpoint(self.hass, user_input[CONF_API_ENDPOINT])
except ValueError as err: except ValueError as err:
_LOGGER.error("Endpoint validation failed: %s", err) _LOGGER.error("Endpoint validation failed: %s", err)
self._errors["base"] = "cannot_connect" self._errors["base"] = "cannot_connect"
@@ -313,7 +318,7 @@ class OptionsFlowHandler(config_entries.OptionsFlow):
return False return False
try: try:
endpoint = validate_endpoint(endpoint) endpoint = await validate_endpoint(self.hass, endpoint)
except ValueError as err: except ValueError as err:
_LOGGER.error("Endpoint validation failed: %s", err) _LOGGER.error("Endpoint validation failed: %s", err)
self._errors["base"] = "cannot_connect" self._errors["base"] = "cannot_connect"
@@ -507,13 +512,13 @@ class OptionsFlowHandler(config_entries.OptionsFlow):
default=data.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES) default=data.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
): vol.All( ): vol.All(
vol.Coerce(int), vol.Coerce(int),
vol.Range(min=1, max=20) vol.Range(min=MIN_CONTEXT_MESSAGES, max=MAX_CONTEXT_MESSAGES)
), ),
vol.Optional( vol.Optional(
CONF_MAX_HISTORY_SIZE, CONF_MAX_HISTORY_SIZE,
default=data.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY) default=data.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY)
): vol.All( ): vol.All(
vol.Coerce(int), vol.Coerce(int),
vol.Range(min=1, max=100) vol.Range(min=MIN_HISTORY_SIZE, max=MAX_HISTORY_SIZE)
), ),
}) })
+5
View File
@@ -83,7 +83,12 @@ DEFAULT_API_TIMEOUT: Final = 30
DEFAULT_MAX_HISTORY: Final = 50 DEFAULT_MAX_HISTORY: Final = 50
DEFAULT_NAME: Final = "HA Text AI" DEFAULT_NAME: Final = "HA Text AI"
DEFAULT_NAME_PREFIX = "ha_text_ai" DEFAULT_NAME_PREFIX = "ha_text_ai"
DEFAULT_INSTANCE_NAME: Final = "my_assistant"
DEFAULT_CONTEXT_MESSAGES: Final = 5 DEFAULT_CONTEXT_MESSAGES: Final = 5
MIN_CONTEXT_MESSAGES: Final = 1
MAX_CONTEXT_MESSAGES: Final = 20
MIN_HISTORY_SIZE: Final = 1
MAX_HISTORY_SIZE: Final = 100
TRUNCATION_INDICATOR = " ... " TRUNCATION_INDICATOR = " ... "
@@ -406,7 +406,6 @@ class HATextAICoordinator(DataUpdateCoordinator):
"history_info": { "history_info": {
"total_entries": 0, "total_entries": 0,
"displayed_entries": 0, "displayed_entries": 0,
"full_history_available": True,
}, },
"normalized_name": self.normalized_name, "normalized_name": self.normalized_name,
} }
+1 -4
View File
@@ -13,6 +13,7 @@ import logging
import os import os
import shutil import shutil
import traceback import traceback
from datetime import datetime
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
import aiofiles import aiofiles
@@ -386,8 +387,6 @@ class HistoryManager:
if start_date: if start_date:
try: try:
from datetime import datetime
start_dt = datetime.fromisoformat( start_dt = datetime.fromisoformat(
start_date.replace("Z", "+00:00") start_date.replace("Z", "+00:00")
) )
@@ -441,12 +440,10 @@ class HistoryManager:
history_info = { history_info = {
"total_entries": len(self._conversation_history), "total_entries": len(self._conversation_history),
"displayed_entries": len(limited_history), "displayed_entries": len(limited_history),
"full_history_available": True,
} }
return { return {
"entries": limited_history, "entries": limited_history,
"full_history": list(self._conversation_history),
"info": history_info, "info": history_info,
} }
+10 -1
View File
@@ -11,6 +11,7 @@ from __future__ import annotations
import json import json
import logging import logging
import os import os
import re
import traceback import traceback
from typing import Any, Dict from typing import Any, Dict
@@ -121,11 +122,19 @@ class MetricsManager:
self._performance_metrics["failed_requests"] += 1 self._performance_metrics["failed_requests"] += 1
await self._save_metrics() await self._save_metrics()
error_msg = str(error)
# Strip URLs, API keys, and query parameters from error messages
error_msg = re.sub(r'https?://\S+', '[URL]', error_msg)
error_msg = re.sub(r'[?&]key=[^\s&]+', '?key=***', error_msg)
error_msg = re.sub(r'AIza[A-Za-z0-9_-]+', '***', error_msg)
if len(error_msg) > 256:
error_msg = error_msg[:256] + "..."
error_details: Dict[str, Any] = { error_details: Dict[str, Any] = {
"timestamp": dt_util.utcnow().isoformat(), "timestamp": dt_util.utcnow().isoformat(),
"model": model, "model": model,
"instance": self.instance_name, "instance": self.instance_name,
"error_message": str(error), "error_message": error_msg,
"error_type": type(error).__name__, "error_type": type(error).__name__,
"traceback": traceback.format_exc() "traceback": traceback.format_exc()
if _LOGGER.isEnabledFor(logging.DEBUG) if _LOGGER.isEnabledFor(logging.DEBUG)
+30 -16
View File
@@ -12,6 +12,7 @@ from typing import Any
from urllib.parse import urlparse from urllib.parse import urlparse
from homeassistant.const import CONF_API_KEY from homeassistant.const import CONF_API_KEY
from homeassistant.core import HomeAssistant
def normalize_name(name: str) -> str: def normalize_name(name: str) -> str:
@@ -29,11 +30,27 @@ def safe_log_data(
return {k: "***" if k in sensitive_keys else v for k, v in data.items()} return {k: "***" if k in sensitive_keys else v for k, v in data.items()}
class _RestrictedIPError(ValueError):
"""Raised when an IP address is in a restricted range."""
def validate_endpoint(endpoint: str) -> str:
def _check_ip_restricted(addr: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool:
"""Check if an IP address is in a restricted range."""
return (
addr.is_private
or addr.is_reserved
or addr.is_loopback
or addr.is_link_local
or addr.is_multicast
or addr.is_unspecified
)
async def validate_endpoint(hass: HomeAssistant, endpoint: str) -> str:
"""Validate API endpoint URL for security. """Validate API endpoint URL for security.
Ensures HTTPS-only and blocks private/reserved IP ranges (SSRF protection). Ensures HTTPS-only and blocks private/reserved IP ranges (SSRF protection).
Uses async DNS resolution to avoid blocking the event loop.
Returns the validated endpoint stripped of trailing slashes. Returns the validated endpoint stripped of trailing slashes.
Raises: Raises:
@@ -51,28 +68,25 @@ def validate_endpoint(endpoint: str) -> str:
# Block private/reserved IPs (direct IP or resolved hostname) # Block private/reserved IPs (direct IP or resolved hostname)
try: try:
addr = ipaddress.ip_address(hostname) addr = ipaddress.ip_address(hostname)
if addr.is_private or addr.is_reserved or addr.is_loopback or addr.is_link_local: if _check_ip_restricted(addr):
raise ValueError("Private/reserved IP addresses are not allowed") raise _RestrictedIPError("Private/reserved IP addresses are not allowed")
except ValueError as e: except _RestrictedIPError:
if "not allowed" in str(e): raise
raise except ValueError:
# Not an IP literal — resolve hostname and check all resolved IPs # Not an IP literal — resolve hostname and check all resolved IPs
# to prevent DNS rebinding attacks # to prevent DNS rebinding attacks
try: try:
addrinfos = socket.getaddrinfo(hostname, None) addrinfos = await hass.async_add_executor_job(
socket.getaddrinfo, hostname, None
)
for family, _type, _proto, _canonname, sockaddr in addrinfos: for family, _type, _proto, _canonname, sockaddr in addrinfos:
ip_str = sockaddr[0] ip_str = sockaddr[0]
resolved_addr = ipaddress.ip_address(ip_str) resolved_addr = ipaddress.ip_address(ip_str)
if ( if _check_ip_restricted(resolved_addr):
resolved_addr.is_private
or resolved_addr.is_reserved
or resolved_addr.is_loopback
or resolved_addr.is_link_local
):
raise ValueError( raise ValueError(
f"Hostname {hostname} resolves to private/reserved IP {ip_str}" "Hostname resolves to a restricted IP range"
) )
except socket.gaierror: except socket.gaierror as err:
raise ValueError(f"Cannot resolve hostname: {hostname}") raise ValueError(f"Cannot resolve hostname: {hostname}") from err
return endpoint.rstrip("/") return endpoint.rstrip("/")