mirror of
https://github.com/smkrv/ha-text-ai.git
synced 2026-07-21 22:54:00 +08:00
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:
@@ -11,7 +11,6 @@ from __future__ import annotations
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import shutil
|
import shutil
|
||||||
import hashlib
|
|
||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
from typing import Any, Dict, TypeVar
|
from typing import Any, Dict, TypeVar
|
||||||
|
|
||||||
@@ -27,6 +26,8 @@ from homeassistant.helpers import aiohttp_client
|
|||||||
|
|
||||||
from .coordinator import HATextAICoordinator
|
from .coordinator import HATextAICoordinator
|
||||||
from .api_client import APIClient
|
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 (
|
from .const import (
|
||||||
DOMAIN,
|
DOMAIN,
|
||||||
PLATFORMS,
|
PLATFORMS,
|
||||||
@@ -38,19 +39,11 @@ from .const import (
|
|||||||
CONF_API_TIMEOUT,
|
CONF_API_TIMEOUT,
|
||||||
CONF_API_PROVIDER,
|
CONF_API_PROVIDER,
|
||||||
CONF_CONTEXT_MESSAGES,
|
CONF_CONTEXT_MESSAGES,
|
||||||
API_PROVIDER_OPENAI,
|
|
||||||
API_PROVIDER_ANTHROPIC,
|
API_PROVIDER_ANTHROPIC,
|
||||||
API_PROVIDER_DEEPSEEK,
|
API_PROVIDER_DEEPSEEK,
|
||||||
API_PROVIDER_GEMINI,
|
API_PROVIDER_GEMINI,
|
||||||
DEFAULT_MODEL,
|
|
||||||
DEFAULT_DEEPSEEK_MODEL,
|
|
||||||
DEFAULT_GEMINI_MODEL,
|
|
||||||
DEFAULT_TEMPERATURE,
|
DEFAULT_TEMPERATURE,
|
||||||
DEFAULT_MAX_TOKENS,
|
DEFAULT_MAX_TOKENS,
|
||||||
DEFAULT_OPENAI_ENDPOINT,
|
|
||||||
DEFAULT_ANTHROPIC_ENDPOINT,
|
|
||||||
DEFAULT_DEEPSEEK_ENDPOINT,
|
|
||||||
DEFAULT_GEMINI_ENDPOINT,
|
|
||||||
DEFAULT_REQUEST_INTERVAL,
|
DEFAULT_REQUEST_INTERVAL,
|
||||||
DEFAULT_API_TIMEOUT,
|
DEFAULT_API_TIMEOUT,
|
||||||
DEFAULT_CONTEXT_MESSAGES,
|
DEFAULT_CONTEXT_MESSAGES,
|
||||||
@@ -105,14 +98,6 @@ def get_coordinator_by_instance(hass: HomeAssistant, instance: str) -> HATextAIC
|
|||||||
|
|
||||||
raise HomeAssistantError(f"Instance {instance} not found")
|
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:
|
async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool:
|
||||||
"""Set up the Home Assistant Text AI component."""
|
"""Set up the Home Assistant Text AI component."""
|
||||||
# Initialize domain data storage
|
# 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)
|
session = aiohttp_client.async_get_clientsession(hass)
|
||||||
|
|
||||||
# Get default endpoint based on provider
|
model = config.get(CONF_MODEL, get_default_model(api_provider))
|
||||||
default_endpoint = {
|
endpoint = config.get(CONF_API_ENDPOINT, get_default_endpoint(api_provider)).rstrip('/')
|
||||||
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('/')
|
|
||||||
# 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))
|
||||||
instance_name = entry.data.get(CONF_NAME, entry.entry_id)
|
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)
|
context_messages = config.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
|
||||||
is_anthropic = api_provider == API_PROVIDER_ANTHROPIC
|
is_anthropic = api_provider == API_PROVIDER_ANTHROPIC
|
||||||
|
|
||||||
headers = {
|
headers = build_auth_headers(api_provider, api_key)
|
||||||
"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}"
|
|
||||||
|
|
||||||
if not await async_check_api(session, endpoint, headers, api_provider, api_timeout):
|
if not await async_check_api(session, endpoint, headers, api_provider, api_timeout):
|
||||||
raise ConfigEntryNotReady("API connection failed")
|
raise ConfigEntryNotReady("API connection failed")
|
||||||
|
|||||||
@@ -33,17 +33,10 @@ from .const import (
|
|||||||
API_PROVIDER_DEEPSEEK,
|
API_PROVIDER_DEEPSEEK,
|
||||||
API_PROVIDER_GEMINI,
|
API_PROVIDER_GEMINI,
|
||||||
API_PROVIDERS,
|
API_PROVIDERS,
|
||||||
DEFAULT_MODEL,
|
|
||||||
DEFAULT_DEEPSEEK_MODEL,
|
|
||||||
DEFAULT_GEMINI_MODEL,
|
|
||||||
DEFAULT_TEMPERATURE,
|
DEFAULT_TEMPERATURE,
|
||||||
DEFAULT_MAX_TOKENS,
|
DEFAULT_MAX_TOKENS,
|
||||||
DEFAULT_REQUEST_INTERVAL,
|
DEFAULT_REQUEST_INTERVAL,
|
||||||
DEFAULT_API_TIMEOUT,
|
DEFAULT_API_TIMEOUT,
|
||||||
DEFAULT_OPENAI_ENDPOINT,
|
|
||||||
DEFAULT_ANTHROPIC_ENDPOINT,
|
|
||||||
DEFAULT_DEEPSEEK_ENDPOINT,
|
|
||||||
DEFAULT_GEMINI_ENDPOINT,
|
|
||||||
DEFAULT_CONTEXT_MESSAGES,
|
DEFAULT_CONTEXT_MESSAGES,
|
||||||
MIN_TEMPERATURE,
|
MIN_TEMPERATURE,
|
||||||
MAX_TEMPERATURE,
|
MAX_TEMPERATURE,
|
||||||
@@ -56,17 +49,12 @@ from .const import (
|
|||||||
DEFAULT_MAX_HISTORY,
|
DEFAULT_MAX_HISTORY,
|
||||||
CONF_MAX_HISTORY_SIZE,
|
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__)
|
_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):
|
class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
||||||
"""Handle a config flow for HA text AI."""
|
"""Handle a config flow for HA text AI."""
|
||||||
|
|
||||||
@@ -101,20 +89,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
self._errors = {}
|
self._errors = {}
|
||||||
|
|
||||||
if user_input is None:
|
if user_input is None:
|
||||||
# Selecting an endpoint by provider
|
default_endpoint = get_default_endpoint(self._provider)
|
||||||
default_endpoint = {
|
default_model = get_default_model(self._provider)
|
||||||
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
|
|
||||||
)
|
|
||||||
|
|
||||||
return self.async_show_form(
|
return self.async_show_form(
|
||||||
step_id="provider",
|
step_id="provider",
|
||||||
@@ -176,8 +152,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
data_schema=vol.Schema({
|
data_schema=vol.Schema({
|
||||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
||||||
vol.Required(CONF_API_KEY): 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_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, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): 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.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All(
|
||||||
vol.Coerce(float),
|
vol.Coerce(float),
|
||||||
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
||||||
@@ -222,8 +198,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
data_schema=vol.Schema({
|
data_schema=vol.Schema({
|
||||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
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_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_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, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): 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.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All(
|
||||||
vol.Coerce(float),
|
vol.Coerce(float),
|
||||||
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
||||||
@@ -270,8 +246,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
data_schema=vol.Schema({
|
data_schema=vol.Schema({
|
||||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
||||||
vol.Required(CONF_API_KEY): str,
|
vol.Required(CONF_API_KEY): str,
|
||||||
vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL)): 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, DEFAULT_GEMINI_ENDPOINT)): str,
|
vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str,
|
||||||
# Other fields remain the same
|
# Other fields remain the same
|
||||||
}),
|
}),
|
||||||
errors=self._errors
|
errors=self._errors
|
||||||
@@ -284,8 +260,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
data_schema=vol.Schema({
|
data_schema=vol.Schema({
|
||||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
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_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_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, DEFAULT_OPENAI_ENDPOINT)): 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.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All(
|
||||||
vol.Coerce(float),
|
vol.Coerce(float),
|
||||||
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
||||||
@@ -327,8 +303,8 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
data_schema=vol.Schema({
|
data_schema=vol.Schema({
|
||||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
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_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_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, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): str,
|
vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, get_default_endpoint(self._provider))): str,
|
||||||
# Other fields remain the same
|
# Other fields remain the same
|
||||||
}),
|
}),
|
||||||
errors={"base": str(e)}
|
errors={"base": str(e)}
|
||||||
@@ -436,11 +412,7 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
|
|
||||||
unique_id = f"{DOMAIN}_{normalized_name}_{self._provider}".lower()
|
unique_id = f"{DOMAIN}_{normalized_name}_{self._provider}".lower()
|
||||||
|
|
||||||
default_model = (
|
default_model = get_default_model(self._provider)
|
||||||
DEFAULT_DEEPSEEK_MODEL if self._provider == API_PROVIDER_DEEPSEEK else
|
|
||||||
DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else
|
|
||||||
DEFAULT_MODEL
|
|
||||||
)
|
|
||||||
|
|
||||||
entry_data = {
|
entry_data = {
|
||||||
CONF_API_PROVIDER: self._provider,
|
CONF_API_PROVIDER: self._provider,
|
||||||
@@ -486,20 +458,11 @@ class OptionsFlowHandler(config_entries.OptionsFlow):
|
|||||||
|
|
||||||
def _get_default_endpoint(self, provider: str) -> str:
|
def _get_default_endpoint(self, provider: str) -> str:
|
||||||
"""Get default endpoint for provider."""
|
"""Get default endpoint for provider."""
|
||||||
return {
|
return get_default_endpoint(provider)
|
||||||
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)
|
|
||||||
|
|
||||||
def _get_default_model(self, provider: str) -> str:
|
def _get_default_model(self, provider: str) -> str:
|
||||||
"""Get default model for provider."""
|
"""Get default model for provider."""
|
||||||
return (
|
return get_default_model(provider)
|
||||||
DEFAULT_DEEPSEEK_MODEL if provider == API_PROVIDER_DEEPSEEK else
|
|
||||||
DEFAULT_GEMINI_MODEL if provider == API_PROVIDER_GEMINI else
|
|
||||||
DEFAULT_MODEL
|
|
||||||
)
|
|
||||||
|
|
||||||
def _get_api_headers(self, api_key: str, provider: str) -> Dict[str, str]:
|
def _get_api_headers(self, api_key: str, provider: str) -> Dict[str, str]:
|
||||||
"""Get API headers based on provider."""
|
"""Get API headers based on provider."""
|
||||||
|
|||||||
@@ -76,6 +76,7 @@ ICONS_SUBDOMAIN = "icons"
|
|||||||
|
|
||||||
# Default values
|
# Default values
|
||||||
DEFAULT_MODEL: Final = "gpt-4o-mini"
|
DEFAULT_MODEL: Final = "gpt-4o-mini"
|
||||||
|
DEFAULT_ANTHROPIC_MODEL: Final = "claude-3-5-sonnet"
|
||||||
DEFAULT_DEEPSEEK_MODEL: Final = "deepseek-chat"
|
DEFAULT_DEEPSEEK_MODEL: Final = "deepseek-chat"
|
||||||
DEFAULT_GEMINI_MODEL: Final = "gemini-2.0-flash"
|
DEFAULT_GEMINI_MODEL: Final = "gemini-2.0-flash"
|
||||||
DEFAULT_TEMPERATURE: Final = 0.1
|
DEFAULT_TEMPERATURE: Final = 0.1
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from homeassistant.util import dt as dt_util
|
|||||||
from homeassistant.exceptions import HomeAssistantError
|
from homeassistant.exceptions import HomeAssistantError
|
||||||
from homeassistant.const import CONF_NAME
|
from homeassistant.const import CONF_NAME
|
||||||
|
|
||||||
from .config_flow import normalize_name
|
from .utils import normalize_name
|
||||||
from .const import (
|
from .const import (
|
||||||
DOMAIN,
|
DOMAIN,
|
||||||
STATE_READY,
|
STATE_READY,
|
||||||
@@ -83,9 +83,6 @@ class HATextAICoordinator(DataUpdateCoordinator):
|
|||||||
self.instance_name = instance_name
|
self.instance_name = instance_name
|
||||||
self.normalized_name = None
|
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.normalized_name = normalize_name(instance_name)
|
||||||
self._metrics_file = os.path.join(
|
self._metrics_file = os.path.join(
|
||||||
hass.config.path(".storage"),
|
hass.config.path(".storage"),
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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()}
|
||||||
@@ -805,7 +805,7 @@ chore: Bump version to 2.4.0
|
|||||||
|
|
||||||
| Фаза | Статус | Дата начала | Дата завершения | Коммит |
|
| Фаза | Статус | Дата начала | Дата завершения | Коммит |
|
||||||
|------|--------|-------------|-----------------|--------|
|
|------|--------|-------------|-----------------|--------|
|
||||||
| 0. Подготовка | [ ] | | | |
|
| 0. Подготовка | [x] | 2026-03-12 | 2026-03-12 | e7c8b22+ |
|
||||||
| 1. Security | [ ] | | | |
|
| 1. Security | [ ] | | | |
|
||||||
| 2. Critical Bugs | [ ] | | | |
|
| 2. Critical Bugs | [ ] | | | |
|
||||||
| 3. Concurrency | [ ] | | | |
|
| 3. Concurrency | [ ] | | | |
|
||||||
|
|||||||
Reference in New Issue
Block a user