diff --git a/custom_components/ha_text_ai/__init__.py b/custom_components/ha_text_ai/__init__.py index a02649e..64ad642 100644 --- a/custom_components/ha_text_ai/__init__.py +++ b/custom_components/ha_text_ai/__init__.py @@ -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( diff --git a/custom_components/ha_text_ai/api_client.py b/custom_components/ha_text_ai/api_client.py index 733ff78..5851060 100644 --- a/custom_components/ha_text_ai/api_client.py +++ b/custom_components/ha_text_ai/api_client.py @@ -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.""" diff --git a/custom_components/ha_text_ai/config_flow.py b/custom_components/ha_text_ai/config_flow.py index 4cbad01..7afede8 100644 --- a/custom_components/ha_text_ai/config_flow.py +++ b/custom_components/ha_text_ai/config_flow.py @@ -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) ), }) diff --git a/custom_components/ha_text_ai/const.py b/custom_components/ha_text_ai/const.py index 73597d9..ea7b0f1 100644 --- a/custom_components/ha_text_ai/const.py +++ b/custom_components/ha_text_ai/const.py @@ -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 = " ... " diff --git a/custom_components/ha_text_ai/coordinator.py b/custom_components/ha_text_ai/coordinator.py index 6c8b977..b0670b0 100644 --- a/custom_components/ha_text_ai/coordinator.py +++ b/custom_components/ha_text_ai/coordinator.py @@ -406,7 +406,6 @@ class HATextAICoordinator(DataUpdateCoordinator): "history_info": { "total_entries": 0, "displayed_entries": 0, - "full_history_available": True, }, "normalized_name": self.normalized_name, } diff --git a/custom_components/ha_text_ai/history.py b/custom_components/ha_text_ai/history.py index cec95dc..f2d9ce8 100644 --- a/custom_components/ha_text_ai/history.py +++ b/custom_components/ha_text_ai/history.py @@ -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, } diff --git a/custom_components/ha_text_ai/metrics.py b/custom_components/ha_text_ai/metrics.py index 31f7503..2b88bce 100644 --- a/custom_components/ha_text_ai/metrics.py +++ b/custom_components/ha_text_ai/metrics.py @@ -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) diff --git a/custom_components/ha_text_ai/utils.py b/custom_components/ha_text_ai/utils.py index 1b126ea..93cca7d 100644 --- a/custom_components/ha_text_ai/utils.py +++ b/custom_components/ha_text_ai/utils.py @@ -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): - raise + 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("/")