mirror of
https://github.com/smkrv/ha-text-ai.git
synced 2026-07-21 22:54:00 +08:00
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:
@@ -263,9 +263,9 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
|
||||
model = config.get(CONF_MODEL, get_default_model(api_provider))
|
||||
raw_endpoint = config.get(CONF_API_ENDPOINT, get_default_endpoint(api_provider))
|
||||
try:
|
||||
endpoint = validate_endpoint(raw_endpoint)
|
||||
endpoint = await validate_endpoint(hass, raw_endpoint)
|
||||
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
|
||||
# API key can now be updated via options
|
||||
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,
|
||||
model=model,
|
||||
api_timeout=api_timeout,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
coordinator = HATextAICoordinator(
|
||||
|
||||
@@ -39,6 +39,7 @@ class APIClient:
|
||||
api_provider: str,
|
||||
model: str,
|
||||
api_timeout: int = DEFAULT_API_TIMEOUT,
|
||||
api_key: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Initialize API client."""
|
||||
self.session = session
|
||||
@@ -48,6 +49,9 @@ class APIClient:
|
||||
self.model = model
|
||||
self.api_timeout = 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
|
||||
|
||||
async def __aenter__(self):
|
||||
@@ -119,7 +123,8 @@ class APIClient:
|
||||
raise HomeAssistantError("API rate limit exceeded")
|
||||
|
||||
# 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}")
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
@@ -368,8 +373,7 @@ class APIClient:
|
||||
|
||||
genai = await asyncio.to_thread(import_genai)
|
||||
|
||||
# Extract API key from headers (Bearer token)
|
||||
api_key = self.headers.get("Authorization", "").replace("Bearer ", "")
|
||||
api_key = self._api_key
|
||||
|
||||
def create_client():
|
||||
if self.endpoint and self.endpoint != "https://generativelanguage.googleapis.com/v1beta":
|
||||
@@ -502,11 +506,11 @@ class APIClient:
|
||||
}
|
||||
|
||||
except ImportError as e:
|
||||
_LOGGER.error(f"Google Gemini library not installed: {str(e)}")
|
||||
raise HomeAssistantError(f"Missing dependency: {str(e)}. Please install google-genai.")
|
||||
_LOGGER.error("Google Gemini library not installed: %s", e)
|
||||
raise HomeAssistantError("Missing dependency: google-genai. Please install it.")
|
||||
except Exception as e:
|
||||
_LOGGER.error(f"Gemini API error: {str(e)}")
|
||||
raise HomeAssistantError(f"Gemini API error: {str(e)}")
|
||||
_LOGGER.error("Gemini API error: %s", e)
|
||||
raise HomeAssistantError("Gemini API request failed")
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
"""Shutdown API client."""
|
||||
|
||||
@@ -45,8 +45,13 @@ from .const import (
|
||||
MIN_API_TIMEOUT,
|
||||
MAX_API_TIMEOUT,
|
||||
DEFAULT_NAME_PREFIX,
|
||||
DEFAULT_INSTANCE_NAME,
|
||||
DEFAULT_MAX_HISTORY,
|
||||
CONF_MAX_HISTORY_SIZE,
|
||||
MIN_CONTEXT_MESSAGES,
|
||||
MAX_CONTEXT_MESSAGES,
|
||||
MIN_HISTORY_SIZE,
|
||||
MAX_HISTORY_SIZE,
|
||||
)
|
||||
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."""
|
||||
defaults = data or {}
|
||||
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_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,
|
||||
@@ -116,14 +121,14 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
||||
default=defaults.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=20)
|
||||
vol.Range(min=MIN_CONTEXT_MESSAGES, max=MAX_CONTEXT_MESSAGES)
|
||||
),
|
||||
vol.Optional(
|
||||
CONF_MAX_HISTORY_SIZE,
|
||||
default=defaults.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY)
|
||||
): vol.All(
|
||||
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
|
||||
|
||||
try:
|
||||
endpoint = validate_endpoint(user_input[CONF_API_ENDPOINT])
|
||||
endpoint = await validate_endpoint(self.hass, user_input[CONF_API_ENDPOINT])
|
||||
except ValueError as err:
|
||||
_LOGGER.error("Endpoint validation failed: %s", err)
|
||||
self._errors["base"] = "cannot_connect"
|
||||
@@ -313,7 +318,7 @@ class OptionsFlowHandler(config_entries.OptionsFlow):
|
||||
return False
|
||||
|
||||
try:
|
||||
endpoint = validate_endpoint(endpoint)
|
||||
endpoint = await validate_endpoint(self.hass, endpoint)
|
||||
except ValueError as err:
|
||||
_LOGGER.error("Endpoint validation failed: %s", err)
|
||||
self._errors["base"] = "cannot_connect"
|
||||
@@ -507,13 +512,13 @@ class OptionsFlowHandler(config_entries.OptionsFlow):
|
||||
default=data.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=20)
|
||||
vol.Range(min=MIN_CONTEXT_MESSAGES, max=MAX_CONTEXT_MESSAGES)
|
||||
),
|
||||
vol.Optional(
|
||||
CONF_MAX_HISTORY_SIZE,
|
||||
default=data.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=100)
|
||||
vol.Range(min=MIN_HISTORY_SIZE, max=MAX_HISTORY_SIZE)
|
||||
),
|
||||
})
|
||||
|
||||
@@ -83,7 +83,12 @@ DEFAULT_API_TIMEOUT: Final = 30
|
||||
DEFAULT_MAX_HISTORY: Final = 50
|
||||
DEFAULT_NAME: Final = "HA Text AI"
|
||||
DEFAULT_NAME_PREFIX = "ha_text_ai"
|
||||
DEFAULT_INSTANCE_NAME: Final = "my_assistant"
|
||||
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 = " ... "
|
||||
|
||||
|
||||
@@ -406,7 +406,6 @@ class HATextAICoordinator(DataUpdateCoordinator):
|
||||
"history_info": {
|
||||
"total_entries": 0,
|
||||
"displayed_entries": 0,
|
||||
"full_history_available": True,
|
||||
},
|
||||
"normalized_name": self.normalized_name,
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ import logging
|
||||
import os
|
||||
import shutil
|
||||
import traceback
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import aiofiles
|
||||
@@ -386,8 +387,6 @@ class HistoryManager:
|
||||
|
||||
if start_date:
|
||||
try:
|
||||
from datetime import datetime
|
||||
|
||||
start_dt = datetime.fromisoformat(
|
||||
start_date.replace("Z", "+00:00")
|
||||
)
|
||||
@@ -441,12 +440,10 @@ class HistoryManager:
|
||||
history_info = {
|
||||
"total_entries": len(self._conversation_history),
|
||||
"displayed_entries": len(limited_history),
|
||||
"full_history_available": True,
|
||||
}
|
||||
|
||||
return {
|
||||
"entries": limited_history,
|
||||
"full_history": list(self._conversation_history),
|
||||
"info": history_info,
|
||||
}
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from __future__ import annotations
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import traceback
|
||||
from typing import Any, Dict
|
||||
|
||||
@@ -121,11 +122,19 @@ class MetricsManager:
|
||||
self._performance_metrics["failed_requests"] += 1
|
||||
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] = {
|
||||
"timestamp": dt_util.utcnow().isoformat(),
|
||||
"model": model,
|
||||
"instance": self.instance_name,
|
||||
"error_message": str(error),
|
||||
"error_message": error_msg,
|
||||
"error_type": type(error).__name__,
|
||||
"traceback": traceback.format_exc()
|
||||
if _LOGGER.isEnabledFor(logging.DEBUG)
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from homeassistant.const import CONF_API_KEY
|
||||
from homeassistant.core import HomeAssistant
|
||||
|
||||
|
||||
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()}
|
||||
|
||||
|
||||
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.
|
||||
|
||||
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.
|
||||
|
||||
Raises:
|
||||
@@ -51,28 +68,25 @@ def validate_endpoint(endpoint: str) -> str:
|
||||
# Block private/reserved IPs (direct IP or resolved hostname)
|
||||
try:
|
||||
addr = ipaddress.ip_address(hostname)
|
||||
if addr.is_private or addr.is_reserved or addr.is_loopback or addr.is_link_local:
|
||||
raise ValueError("Private/reserved IP addresses are not allowed")
|
||||
except ValueError as e:
|
||||
if "not allowed" in str(e):
|
||||
if _check_ip_restricted(addr):
|
||||
raise _RestrictedIPError("Private/reserved IP addresses are not allowed")
|
||||
except _RestrictedIPError:
|
||||
raise
|
||||
except ValueError:
|
||||
# Not an IP literal — resolve hostname and check all resolved IPs
|
||||
# to prevent DNS rebinding attacks
|
||||
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:
|
||||
ip_str = sockaddr[0]
|
||||
resolved_addr = ipaddress.ip_address(ip_str)
|
||||
if (
|
||||
resolved_addr.is_private
|
||||
or resolved_addr.is_reserved
|
||||
or resolved_addr.is_loopback
|
||||
or resolved_addr.is_link_local
|
||||
):
|
||||
if _check_ip_restricted(resolved_addr):
|
||||
raise ValueError(
|
||||
f"Hostname {hostname} resolves to private/reserved IP {ip_str}"
|
||||
"Hostname resolves to a restricted IP range"
|
||||
)
|
||||
except socket.gaierror:
|
||||
raise ValueError(f"Cannot resolve hostname: {hostname}")
|
||||
except socket.gaierror as err:
|
||||
raise ValueError(f"Cannot resolve hostname: {hostname}") from err
|
||||
|
||||
return endpoint.rstrip("/")
|
||||
|
||||
Reference in New Issue
Block a user