refactor: Phase 0 — extract utils.py and providers.py, centralize provider dispatch

- Create utils.py: normalize_name, get_file_hash, safe_log_data
- Create providers.py: PROVIDER_REGISTRY with get_default_endpoint,
  get_default_model, build_auth_headers
- Add DEFAULT_ANTHROPIC_MODEL constant (was incorrectly using gpt-4o-mini)
- Replace all inline dispatch tables in config_flow.py and __init__.py
- Fix circular import: coordinator.py now imports from utils, not config_flow
- Fix NameError in error paths: replace bare constant refs with provider functions
- Raise ValueError for unknown providers instead of silent OpenAI fallback
This commit is contained in:
SMKRV
2026-03-12 00:55:59 +03:00
parent e7c8b22fde
commit ce0a75f219
7 changed files with 156 additions and 103 deletions
+5 -44
View File
@@ -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")
+17 -54
View File
@@ -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."""
+1
View File
@@ -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
+1 -4
View File
@@ -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"),
+95
View File
@@ -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
+36
View File
@@ -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()}
+1 -1
View File
@@ -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 | [ ] | | | |