diff --git a/custom_components/ha_text_ai/__init__.py b/custom_components/ha_text_ai/__init__.py index 179b843..cbb47fd 100644 --- a/custom_components/ha_text_ai/__init__.py +++ b/custom_components/ha_text_ai/__init__.py @@ -11,7 +11,6 @@ from __future__ import annotations import logging import os import shutil -import hashlib from datetime import datetime, timedelta from typing import Any, Dict, TypeVar @@ -27,6 +26,8 @@ from homeassistant.helpers import aiohttp_client from .coordinator import HATextAICoordinator from .api_client import APIClient +from .utils import get_file_hash, safe_log_data +from .providers import get_default_endpoint, get_default_model, build_auth_headers from .const import ( DOMAIN, PLATFORMS, @@ -38,19 +39,11 @@ from .const import ( CONF_API_TIMEOUT, CONF_API_PROVIDER, CONF_CONTEXT_MESSAGES, - API_PROVIDER_OPENAI, API_PROVIDER_ANTHROPIC, API_PROVIDER_DEEPSEEK, API_PROVIDER_GEMINI, - DEFAULT_MODEL, - DEFAULT_DEEPSEEK_MODEL, - DEFAULT_GEMINI_MODEL, DEFAULT_TEMPERATURE, DEFAULT_MAX_TOKENS, - DEFAULT_OPENAI_ENDPOINT, - DEFAULT_ANTHROPIC_ENDPOINT, - DEFAULT_DEEPSEEK_ENDPOINT, - DEFAULT_GEMINI_ENDPOINT, DEFAULT_REQUEST_INTERVAL, DEFAULT_API_TIMEOUT, DEFAULT_CONTEXT_MESSAGES, @@ -105,14 +98,6 @@ def get_coordinator_by_instance(hass: HomeAssistant, instance: str) -> HATextAIC raise HomeAssistantError(f"Instance {instance} not found") -def get_file_hash(file_path: str) -> str: - """Calculate SHA256 hash of file.""" - sha256_hash = hashlib.sha256() - with open(file_path, "rb") as f: - for byte_block in iter(lambda: f.read(4096), b""): - sha256_hash.update(byte_block) - return sha256_hash.hexdigest() - async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool: """Set up the Home Assistant Text AI component.""" # Initialize domain data storage @@ -322,23 +307,8 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool: session = aiohttp_client.async_get_clientsession(hass) - # Get default endpoint based on provider - default_endpoint = { - API_PROVIDER_OPENAI: DEFAULT_OPENAI_ENDPOINT, - API_PROVIDER_ANTHROPIC: DEFAULT_ANTHROPIC_ENDPOINT, - API_PROVIDER_DEEPSEEK: DEFAULT_DEEPSEEK_ENDPOINT, - API_PROVIDER_GEMINI: DEFAULT_GEMINI_ENDPOINT, - }.get(api_provider, DEFAULT_OPENAI_ENDPOINT) - - # Get default model based on provider - default_model = ( - DEFAULT_DEEPSEEK_MODEL if api_provider == API_PROVIDER_DEEPSEEK else - DEFAULT_GEMINI_MODEL if api_provider == API_PROVIDER_GEMINI else - DEFAULT_MODEL - ) - - model = config.get(CONF_MODEL, default_model) - endpoint = config.get(CONF_API_ENDPOINT, default_endpoint).rstrip('/') + model = config.get(CONF_MODEL, get_default_model(api_provider)) + endpoint = config.get(CONF_API_ENDPOINT, get_default_endpoint(api_provider)).rstrip('/') # API key can now be updated via options api_key = config.get(CONF_API_KEY, entry.data.get(CONF_API_KEY)) instance_name = entry.data.get(CONF_NAME, entry.entry_id) @@ -350,16 +320,7 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool: context_messages = config.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES) is_anthropic = api_provider == API_PROVIDER_ANTHROPIC - headers = { - "Content-Type": "application/json", - "Accept": "application/json" - } - - if is_anthropic: - headers["x-api-key"] = api_key - headers["anthropic-version"] = "2023-06-01" - else: - headers["Authorization"] = f"Bearer {api_key}" + headers = build_auth_headers(api_provider, api_key) if not await async_check_api(session, endpoint, headers, api_provider, api_timeout): raise ConfigEntryNotReady("API connection failed") diff --git a/custom_components/ha_text_ai/config_flow.py b/custom_components/ha_text_ai/config_flow.py index fe6b00a..73fa2e4 100644 --- a/custom_components/ha_text_ai/config_flow.py +++ b/custom_components/ha_text_ai/config_flow.py @@ -33,17 +33,10 @@ from .const import ( API_PROVIDER_DEEPSEEK, API_PROVIDER_GEMINI, API_PROVIDERS, - DEFAULT_MODEL, - DEFAULT_DEEPSEEK_MODEL, - DEFAULT_GEMINI_MODEL, DEFAULT_TEMPERATURE, DEFAULT_MAX_TOKENS, DEFAULT_REQUEST_INTERVAL, DEFAULT_API_TIMEOUT, - DEFAULT_OPENAI_ENDPOINT, - DEFAULT_ANTHROPIC_ENDPOINT, - DEFAULT_DEEPSEEK_ENDPOINT, - DEFAULT_GEMINI_ENDPOINT, DEFAULT_CONTEXT_MESSAGES, MIN_TEMPERATURE, MAX_TEMPERATURE, @@ -56,17 +49,12 @@ from .const import ( DEFAULT_MAX_HISTORY, CONF_MAX_HISTORY_SIZE, ) +from .utils import normalize_name # noqa: F401 — re-exported for backward compat +from .providers import get_default_endpoint, get_default_model _LOGGER = logging.getLogger(__name__) -def normalize_name(name: str) -> str: - """Normalize name to conform to HA naming convention using underscores.""" - normalized = ''.join(c if c.isalnum() or c == '_' else '_' for c in name) - normalized = '_'.join(filter(None, normalized.split('_'))) - return normalized.lower() - - class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): """Handle a config flow for HA text AI.""" @@ -101,20 +89,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): self._errors = {} if user_input is None: - # Selecting an endpoint by provider - default_endpoint = { - API_PROVIDER_OPENAI: DEFAULT_OPENAI_ENDPOINT, - API_PROVIDER_ANTHROPIC: DEFAULT_ANTHROPIC_ENDPOINT, - API_PROVIDER_DEEPSEEK: DEFAULT_DEEPSEEK_ENDPOINT, - API_PROVIDER_GEMINI: DEFAULT_GEMINI_ENDPOINT, - }.get(self._provider, DEFAULT_OPENAI_ENDPOINT) - - # Selecting the default model by provider - default_model = ( - DEFAULT_DEEPSEEK_MODEL if self._provider == API_PROVIDER_DEEPSEEK else - DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else - DEFAULT_MODEL - ) + default_endpoint = get_default_endpoint(self._provider) + default_model = get_default_model(self._provider) return self.async_show_form( step_id="provider", @@ -176,8 +152,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): data_schema=vol.Schema({ vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str, vol.Required(CONF_API_KEY): str, - vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else DEFAULT_MODEL)): str, - vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): str, + vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, get_default_model(self._provider))): str, + vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str, vol.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All( vol.Coerce(float), vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE) @@ -222,8 +198,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): data_schema=vol.Schema({ vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str, vol.Required(CONF_API_KEY, default=input_copy.get(CONF_API_KEY, "")): str, - vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else DEFAULT_MODEL)): str, - vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): str, + vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, get_default_model(self._provider))): str, + vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str, vol.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All( vol.Coerce(float), vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE) @@ -270,8 +246,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): data_schema=vol.Schema({ vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str, vol.Required(CONF_API_KEY): str, - vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL)): str, - vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_GEMINI_ENDPOINT)): str, + vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, get_default_model(self._provider))): str, + vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str, # Other fields remain the same }), errors=self._errors @@ -284,8 +260,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): data_schema=vol.Schema({ vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str, vol.Required(CONF_API_KEY, default=input_copy.get(CONF_API_KEY, "")): str, - vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_MODEL)): str, - vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_OPENAI_ENDPOINT)): str, + vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, get_default_model(self._provider))): str, + vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str, vol.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All( vol.Coerce(float), vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE) @@ -327,8 +303,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): data_schema=vol.Schema({ vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str, vol.Required(CONF_API_KEY, default=input_copy.get(CONF_API_KEY, "")): str, - vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else DEFAULT_MODEL)): str, - vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): str, + vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, get_default_model(self._provider))): str, + vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str, # Other fields remain the same }), errors={"base": str(e)} @@ -436,11 +412,7 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): unique_id = f"{DOMAIN}_{normalized_name}_{self._provider}".lower() - default_model = ( - DEFAULT_DEEPSEEK_MODEL if self._provider == API_PROVIDER_DEEPSEEK else - DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else - DEFAULT_MODEL - ) + default_model = get_default_model(self._provider) entry_data = { CONF_API_PROVIDER: self._provider, @@ -486,20 +458,11 @@ class OptionsFlowHandler(config_entries.OptionsFlow): def _get_default_endpoint(self, provider: str) -> str: """Get default endpoint for provider.""" - return { - API_PROVIDER_OPENAI: DEFAULT_OPENAI_ENDPOINT, - API_PROVIDER_ANTHROPIC: DEFAULT_ANTHROPIC_ENDPOINT, - API_PROVIDER_DEEPSEEK: DEFAULT_DEEPSEEK_ENDPOINT, - API_PROVIDER_GEMINI: DEFAULT_GEMINI_ENDPOINT, - }.get(provider, DEFAULT_OPENAI_ENDPOINT) + return get_default_endpoint(provider) def _get_default_model(self, provider: str) -> str: """Get default model for provider.""" - return ( - DEFAULT_DEEPSEEK_MODEL if provider == API_PROVIDER_DEEPSEEK else - DEFAULT_GEMINI_MODEL if provider == API_PROVIDER_GEMINI else - DEFAULT_MODEL - ) + return get_default_model(provider) def _get_api_headers(self, api_key: str, provider: str) -> Dict[str, str]: """Get API headers based on provider.""" diff --git a/custom_components/ha_text_ai/const.py b/custom_components/ha_text_ai/const.py index 275e544..cb245a0 100644 --- a/custom_components/ha_text_ai/const.py +++ b/custom_components/ha_text_ai/const.py @@ -76,6 +76,7 @@ ICONS_SUBDOMAIN = "icons" # Default values DEFAULT_MODEL: Final = "gpt-4o-mini" +DEFAULT_ANTHROPIC_MODEL: Final = "claude-3-5-sonnet" DEFAULT_DEEPSEEK_MODEL: Final = "deepseek-chat" DEFAULT_GEMINI_MODEL: Final = "gemini-2.0-flash" DEFAULT_TEMPERATURE: Final = 0.1 diff --git a/custom_components/ha_text_ai/coordinator.py b/custom_components/ha_text_ai/coordinator.py index f0902e9..d835878 100644 --- a/custom_components/ha_text_ai/coordinator.py +++ b/custom_components/ha_text_ai/coordinator.py @@ -26,7 +26,7 @@ from homeassistant.util import dt as dt_util from homeassistant.exceptions import HomeAssistantError from homeassistant.const import CONF_NAME -from .config_flow import normalize_name +from .utils import normalize_name from .const import ( DOMAIN, STATE_READY, @@ -83,9 +83,6 @@ class HATextAICoordinator(DataUpdateCoordinator): self.instance_name = instance_name self.normalized_name = None - # Use the normalize_name function from config_flow to ensure consistency - from .config_flow import normalize_name - self.normalized_name = normalize_name(instance_name) self._metrics_file = os.path.join( hass.config.path(".storage"), diff --git a/custom_components/ha_text_ai/providers.py b/custom_components/ha_text_ai/providers.py new file mode 100644 index 0000000..e6e70aa --- /dev/null +++ b/custom_components/ha_text_ai/providers.py @@ -0,0 +1,95 @@ +""" +Provider registry for HA Text AI integration. + +Centralizes provider-specific configuration to avoid dispatch duplication +across __init__.py, config_flow.py, and api_client.py. + +@license: CC BY-NC-SA 4.0 International +@author: SMKRV +@github: https://github.com/smkrv/ha-text-ai +@source: https://github.com/smkrv/ha-text-ai +""" +from typing import Any + +from .const import ( + API_PROVIDER_OPENAI, + API_PROVIDER_ANTHROPIC, + API_PROVIDER_DEEPSEEK, + API_PROVIDER_GEMINI, + DEFAULT_MODEL, + DEFAULT_ANTHROPIC_MODEL, + DEFAULT_DEEPSEEK_MODEL, + DEFAULT_GEMINI_MODEL, + DEFAULT_OPENAI_ENDPOINT, + DEFAULT_ANTHROPIC_ENDPOINT, + DEFAULT_DEEPSEEK_ENDPOINT, + DEFAULT_GEMINI_ENDPOINT, +) + +PROVIDER_REGISTRY: dict[str, dict[str, Any]] = { + API_PROVIDER_OPENAI: { + "default_model": DEFAULT_MODEL, + "default_endpoint": DEFAULT_OPENAI_ENDPOINT, + "auth_header": "Authorization", + "auth_prefix": "Bearer ", + "check_path": "/models", + }, + API_PROVIDER_ANTHROPIC: { + "default_model": DEFAULT_ANTHROPIC_MODEL, + "default_endpoint": DEFAULT_ANTHROPIC_ENDPOINT, + "auth_header": "x-api-key", + "auth_prefix": "", + "check_path": "/v1/models", + "extra_headers": { + "anthropic-version": "2023-06-01", + }, + }, + API_PROVIDER_DEEPSEEK: { + "default_model": DEFAULT_DEEPSEEK_MODEL, + "default_endpoint": DEFAULT_DEEPSEEK_ENDPOINT, + "auth_header": "Authorization", + "auth_prefix": "Bearer ", + "check_path": "/models", + }, + API_PROVIDER_GEMINI: { + "default_model": DEFAULT_GEMINI_MODEL, + "default_endpoint": DEFAULT_GEMINI_ENDPOINT, + "auth_header": "Authorization", + "auth_prefix": "Bearer ", + "check_path": None, # Gemini does not support /models check + }, +} + + +def get_provider_config(provider: str) -> dict[str, Any]: + """Get full provider configuration. + + Raises ValueError for unknown providers to avoid sending + credentials to the wrong endpoint. + """ + if provider not in PROVIDER_REGISTRY: + raise ValueError(f"Unknown API provider: {provider}") + return PROVIDER_REGISTRY[provider] + + +def get_default_endpoint(provider: str) -> str: + """Get default API endpoint for a provider.""" + return get_provider_config(provider)["default_endpoint"] + + +def get_default_model(provider: str) -> str: + """Get default model for a provider.""" + return get_provider_config(provider)["default_model"] + + +def build_auth_headers(provider: str, api_key: str) -> dict[str, str]: + """Build authentication headers for a provider.""" + config = get_provider_config(provider) + headers = { + "Content-Type": "application/json", + "Accept": "application/json", + } + headers[config["auth_header"]] = f"{config['auth_prefix']}{api_key}" + if "extra_headers" in config: + headers.update(config["extra_headers"]) + return headers diff --git a/custom_components/ha_text_ai/utils.py b/custom_components/ha_text_ai/utils.py new file mode 100644 index 0000000..41adcfe --- /dev/null +++ b/custom_components/ha_text_ai/utils.py @@ -0,0 +1,36 @@ +""" +Utility functions for HA Text AI integration. + +@license: CC BY-NC-SA 4.0 International +@author: SMKRV +@github: https://github.com/smkrv/ha-text-ai +@source: https://github.com/smkrv/ha-text-ai +""" +import hashlib +from typing import Any + +from homeassistant.const import CONF_API_KEY + + +def normalize_name(name: str) -> str: + """Normalize name to conform to HA naming convention using underscores.""" + normalized = ''.join(c if c.isalnum() or c == '_' else '_' for c in name) + normalized = '_'.join(filter(None, normalized.split('_'))) + return normalized.lower() + + +def get_file_hash(file_path: str) -> str: + """Calculate SHA256 hash of file.""" + sha256_hash = hashlib.sha256() + with open(file_path, "rb") as f: + for byte_block in iter(lambda: f.read(4096), b""): + sha256_hash.update(byte_block) + return sha256_hash.hexdigest() + + +def safe_log_data( + data: dict[str, Any], + sensitive_keys: tuple[str, ...] = (CONF_API_KEY,), +) -> dict[str, Any]: + """Filter sensitive keys from data for safe logging.""" + return {k: "***" if k in sensitive_keys else v for k, v in data.items()} diff --git a/docs/specs/2026-03-12-v2.4.0-fix-plan.md b/docs/specs/2026-03-12-v2.4.0-fix-plan.md index fcbe2cd..b75b37e 100644 --- a/docs/specs/2026-03-12-v2.4.0-fix-plan.md +++ b/docs/specs/2026-03-12-v2.4.0-fix-plan.md @@ -805,7 +805,7 @@ chore: Bump version to 2.4.0 | Фаза | Статус | Дата начала | Дата завершения | Коммит | |------|--------|-------------|-----------------|--------| -| 0. Подготовка | [ ] | | | | +| 0. Подготовка | [x] | 2026-03-12 | 2026-03-12 | e7c8b22+ | | 1. Security | [ ] | | | | | 2. Critical Bugs | [ ] | | | | | 3. Concurrency | [ ] | | | |