mirror of
https://github.com/smkrv/ha-text-ai.git
synced 2026-07-22 07:03:58 +08:00
- service registration extracted to idempotent _async_register_services,
called from both async_setup and async_setup_entry: unloading the last
entry unregisters services, and a reload (every options change) never
re-ran async_setup, leaving the integration without services
- get_history handler wraps the list in {"history": [...]}: HA rejects
non-dict action responses, so every return_response call failed with a
server error since the service gained SupportsResponse (v2.4.x) - the
path never worked, no consumer could depend on the old shape
- README: get_history example shows response_variable usage and the
actual limit clamp semantics
Both found by live smoke test in HA 2026.7 (Docker), not by static review
427 lines
17 KiB
Python
427 lines
17 KiB
Python
"""
|
|
The HA Text AI integration.
|
|
|
|
@license: MIT (https://opensource.org/licenses/MIT)
|
|
@author: SMKRV
|
|
@github: https://github.com/smkrv/ha-text-ai
|
|
@source: https://github.com/smkrv/ha-text-ai
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
from typing import Any
|
|
|
|
import asyncio
|
|
|
|
import voluptuous as vol
|
|
|
|
from homeassistant.config_entries import ConfigEntry
|
|
from homeassistant.const import CONF_API_KEY, CONF_NAME
|
|
from homeassistant.core import HomeAssistant, ServiceCall, SupportsResponse
|
|
from homeassistant.exceptions import ConfigEntryNotReady, HomeAssistantError
|
|
from homeassistant.helpers import config_validation as cv
|
|
from homeassistant.util import dt as dt_util
|
|
|
|
from .coordinator import HATextAICoordinator
|
|
from .api_client import APIClient
|
|
from .utils import create_pinned_session, normalize_name, safe_log_data, validate_endpoint
|
|
from .providers import get_default_endpoint, get_default_model, build_auth_headers
|
|
from .const import (
|
|
DOMAIN,
|
|
PLATFORMS,
|
|
CONF_MODEL,
|
|
CONF_TEMPERATURE,
|
|
CONF_MAX_TOKENS,
|
|
CONF_API_ENDPOINT,
|
|
CONF_REQUEST_INTERVAL,
|
|
CONF_API_TIMEOUT,
|
|
CONF_API_PROVIDER,
|
|
CONF_CONTEXT_MESSAGES,
|
|
DEFAULT_TEMPERATURE,
|
|
DEFAULT_MAX_TOKENS,
|
|
DEFAULT_REQUEST_INTERVAL,
|
|
DEFAULT_API_TIMEOUT,
|
|
DEFAULT_CONTEXT_MESSAGES,
|
|
SERVICE_ASK_QUESTION,
|
|
SERVICE_CLEAR_HISTORY,
|
|
SERVICE_GET_HISTORY,
|
|
SERVICE_SET_SYSTEM_PROMPT,
|
|
DEFAULT_MAX_HISTORY,
|
|
CONF_MAX_HISTORY_SIZE,
|
|
CONF_ALLOW_LOCAL_NETWORK,
|
|
DEFAULT_ALLOW_LOCAL_NETWORK,
|
|
CONF_DISABLE_THINKING,
|
|
DEFAULT_DISABLE_THINKING,
|
|
)
|
|
|
|
_LOGGER = logging.getLogger(__name__)
|
|
|
|
CONFIG_SCHEMA = cv.config_entry_only_config_schema(DOMAIN)
|
|
|
|
SERVICE_SCHEMA_ASK_QUESTION = vol.Schema({
|
|
vol.Required("instance"): cv.string,
|
|
vol.Required("question"): vol.All(cv.string, vol.Length(min=1, max=100000)),
|
|
vol.Optional("system_prompt"): vol.All(cv.string, vol.Length(max=50000)),
|
|
vol.Optional("model"): cv.string,
|
|
vol.Optional("temperature"): vol.All(
|
|
vol.Coerce(float), vol.Range(min=0.0, max=2.0)
|
|
),
|
|
vol.Optional("max_tokens"): cv.positive_int,
|
|
vol.Optional("context_messages"): cv.positive_int,
|
|
vol.Optional("structured_output", default=False): cv.boolean,
|
|
vol.Optional("json_schema"): vol.All(cv.string, vol.Length(max=50000)),
|
|
vol.Optional("disable_thinking"): cv.boolean,
|
|
})
|
|
|
|
SERVICE_SCHEMA_SET_SYSTEM_PROMPT = vol.Schema({
|
|
vol.Required("instance"): cv.string,
|
|
vol.Required("prompt"): cv.string,
|
|
})
|
|
|
|
SERVICE_SCHEMA_GET_HISTORY = vol.Schema({
|
|
vol.Required("instance"): cv.string,
|
|
# No default and no schema max: omitting limit returns the full history
|
|
# (pre-2.5.0 behavior) and oversized values are clamped to
|
|
# ABSOLUTE_MAX_HISTORY_SIZE in history.async_get_history instead of
|
|
# failing the whole service call.
|
|
vol.Optional("limit"): vol.All(cv.positive_int, vol.Range(min=1)),
|
|
vol.Optional("filter_model"): cv.string,
|
|
vol.Optional("start_date"): cv.string,
|
|
vol.Optional("include_metadata"): cv.boolean,
|
|
vol.Optional("sort_order"): vol.In(["newest", "oldest"]),
|
|
})
|
|
|
|
def get_coordinator_by_instance(hass: HomeAssistant, instance: str) -> HATextAICoordinator:
|
|
"""Get coordinator by instance name or normalized name.
|
|
|
|
Accepts instance_name, normalized_name, or sensor entity_id.
|
|
"""
|
|
if instance.startswith("sensor."):
|
|
instance = instance.replace("sensor.ha_text_ai_", "", 1)
|
|
|
|
normalized_input = normalize_name(instance)
|
|
|
|
for entry_id, coord in hass.data[DOMAIN].items():
|
|
if not isinstance(coord, HATextAICoordinator):
|
|
continue
|
|
if (
|
|
coord.instance_name.lower() == instance.lower()
|
|
or coord.normalized_name == normalized_input
|
|
):
|
|
return coord
|
|
|
|
raise HomeAssistantError(f"Instance {instance} not found")
|
|
|
|
async def async_setup(hass: HomeAssistant, config: dict[str, Any]) -> bool:
|
|
"""Set up the Home Assistant Text AI component."""
|
|
# Initialize domain data storage
|
|
hass.data.setdefault(DOMAIN, {})
|
|
_async_register_services(hass)
|
|
return True
|
|
|
|
def _async_register_services(hass: HomeAssistant) -> None:
|
|
"""Register domain services; safe to call again after unload.
|
|
|
|
Unloading the last config entry unregisters the services, and a config
|
|
entry reload (every options change does one) runs unload + setup_entry
|
|
without re-running async_setup — so setup_entry must be able to bring
|
|
the services back.
|
|
"""
|
|
if hass.services.has_service(DOMAIN, SERVICE_ASK_QUESTION):
|
|
return
|
|
|
|
async def async_ask_question(call: ServiceCall) -> dict:
|
|
"""Handle ask_question service with response data."""
|
|
try:
|
|
coordinator = get_coordinator_by_instance(hass, call.data["instance"])
|
|
response = await coordinator.async_ask_question(
|
|
question=call.data["question"],
|
|
model=call.data.get("model"),
|
|
temperature=call.data.get("temperature"),
|
|
max_tokens=call.data.get("max_tokens"),
|
|
system_prompt=call.data.get("system_prompt"),
|
|
context_messages=call.data.get("context_messages"),
|
|
structured_output=call.data.get("structured_output", False),
|
|
json_schema=call.data.get("json_schema"),
|
|
disable_thinking=call.data.get("disable_thinking"),
|
|
)
|
|
|
|
# Return structured response data
|
|
return {
|
|
"response_text": response.get("content", ""),
|
|
"tokens_used": response.get("tokens", {}).get("total", 0),
|
|
"prompt_tokens": response.get("tokens", {}).get("prompt", 0),
|
|
"completion_tokens": response.get("tokens", {}).get("completion", 0),
|
|
"model_used": response.get("model", call.data.get("model", coordinator.model)),
|
|
"instance": call.data["instance"],
|
|
"question": call.data["question"],
|
|
"timestamp": response.get("timestamp"),
|
|
"success": True
|
|
}
|
|
except Exception as err:
|
|
_LOGGER.error("Error asking question: %s", str(err))
|
|
# Return error response
|
|
return {
|
|
"response_text": "",
|
|
"tokens_used": 0,
|
|
"prompt_tokens": 0,
|
|
"completion_tokens": 0,
|
|
"model_used": call.data.get("model", ""),
|
|
"instance": call.data["instance"],
|
|
"question": call.data["question"],
|
|
"timestamp": dt_util.utcnow().isoformat(),
|
|
"success": False,
|
|
"error": str(err),
|
|
"error_type": type(err).__name__
|
|
}
|
|
|
|
async def async_clear_history(call: ServiceCall) -> None:
|
|
"""Handle clear_history service."""
|
|
try:
|
|
coordinator = get_coordinator_by_instance(hass, call.data["instance"])
|
|
await coordinator.async_clear_history()
|
|
except Exception as err:
|
|
_LOGGER.error("Error clearing history: %s", str(err))
|
|
raise HomeAssistantError(f"Failed to clear history: {str(err)}") from err
|
|
|
|
async def async_get_history(call: ServiceCall) -> dict:
|
|
"""Handle get_history service."""
|
|
try:
|
|
coordinator = get_coordinator_by_instance(hass, call.data["instance"])
|
|
history = await coordinator.async_get_history(
|
|
limit=call.data.get("limit"),
|
|
filter_model=call.data.get("filter_model"),
|
|
start_date=call.data.get("start_date"),
|
|
include_metadata=call.data.get("include_metadata", False),
|
|
sort_order=call.data.get("sort_order", "newest")
|
|
)
|
|
# HA requires action responses to be dicts. The bare list made
|
|
# every return_response call fail with a server error, so this
|
|
# path never worked before and the wrapper breaks no consumer.
|
|
return {"history": history}
|
|
except Exception as err:
|
|
_LOGGER.error("Error getting history: %s", str(err))
|
|
raise HomeAssistantError(f"Failed to get history: {str(err)}") from err
|
|
|
|
async def async_set_system_prompt(call: ServiceCall) -> None:
|
|
"""Handle set_system_prompt service."""
|
|
try:
|
|
coordinator = get_coordinator_by_instance(hass, call.data["instance"])
|
|
await coordinator.async_set_system_prompt(call.data["prompt"])
|
|
except Exception as err:
|
|
_LOGGER.error("Error setting system prompt: %s", str(err))
|
|
raise HomeAssistantError(f"Failed to set system prompt: {str(err)}") from err
|
|
|
|
# Register services
|
|
hass.services.async_register(
|
|
DOMAIN,
|
|
SERVICE_ASK_QUESTION,
|
|
async_ask_question,
|
|
schema=SERVICE_SCHEMA_ASK_QUESTION,
|
|
supports_response=SupportsResponse.OPTIONAL
|
|
)
|
|
|
|
hass.services.async_register(
|
|
DOMAIN,
|
|
SERVICE_CLEAR_HISTORY,
|
|
async_clear_history,
|
|
schema=vol.Schema({vol.Required("instance"): cv.string})
|
|
)
|
|
|
|
hass.services.async_register(
|
|
DOMAIN,
|
|
SERVICE_GET_HISTORY,
|
|
async_get_history,
|
|
schema=SERVICE_SCHEMA_GET_HISTORY,
|
|
supports_response=SupportsResponse.OPTIONAL
|
|
)
|
|
|
|
hass.services.async_register(
|
|
DOMAIN,
|
|
SERVICE_SET_SYSTEM_PROMPT,
|
|
async_set_system_prompt,
|
|
schema=SERVICE_SCHEMA_SET_SYSTEM_PROMPT
|
|
)
|
|
|
|
async def async_check_api(session, endpoint: str, headers: dict, provider: str, api_timeout: int = DEFAULT_API_TIMEOUT) -> bool:
|
|
"""Check API availability using provider registry configuration."""
|
|
try:
|
|
from .providers import get_provider_config
|
|
provider_config = get_provider_config(provider)
|
|
check_path = provider_config.get("check_path")
|
|
|
|
if check_path is None:
|
|
# Provider does not support /models check (e.g. Gemini)
|
|
auth_header = provider_config["auth_header"]
|
|
auth_value = headers.get(auth_header, "").replace(provider_config.get("auth_prefix", ""), "")
|
|
if auth_value:
|
|
return True
|
|
_LOGGER.error("API key is missing or empty for %s", provider)
|
|
return False
|
|
|
|
check_url = f"{endpoint}{check_path}"
|
|
|
|
async with asyncio.timeout(api_timeout):
|
|
async with session.get(
|
|
check_url, headers=headers, allow_redirects=False
|
|
) as response:
|
|
if response.status == 200:
|
|
return True
|
|
elif response.status == 401:
|
|
_LOGGER.error("Invalid API key")
|
|
return False
|
|
elif response.status == 429:
|
|
_LOGGER.warning("Rate limit exceeded during API check")
|
|
return False
|
|
else:
|
|
_LOGGER.error("API check failed with status: %d", response.status)
|
|
return False
|
|
except Exception as ex:
|
|
_LOGGER.error("API check error: %s", str(ex))
|
|
return False
|
|
|
|
async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
|
|
"""Set up HA Text AI from a config entry."""
|
|
_LOGGER.debug("Setting up HA Text AI entry: %s", safe_log_data(dict(entry.data)))
|
|
|
|
session = None
|
|
try:
|
|
# Get provider from data or options (options takes precedence)
|
|
config = {**entry.data, **entry.options}
|
|
api_provider = config.get(CONF_API_PROVIDER)
|
|
|
|
if not api_provider:
|
|
_LOGGER.error("API provider not specified")
|
|
raise ConfigEntryNotReady("API provider is required")
|
|
|
|
model = config.get(CONF_MODEL, get_default_model(api_provider))
|
|
raw_endpoint = config.get(CONF_API_ENDPOINT, get_default_endpoint(api_provider))
|
|
allow_local = config.get(CONF_ALLOW_LOCAL_NETWORK, DEFAULT_ALLOW_LOCAL_NETWORK)
|
|
if allow_local:
|
|
_LOGGER.info(
|
|
"Local network mode enabled for endpoint %s — "
|
|
"SSRF protection relaxed for self-hosted proxies",
|
|
raw_endpoint,
|
|
)
|
|
try:
|
|
endpoint, resolved_ips = await validate_endpoint(
|
|
hass, raw_endpoint, allow_local=allow_local
|
|
)
|
|
except ValueError as err:
|
|
_LOGGER.error("Invalid API endpoint: %s", err)
|
|
raise ConfigEntryNotReady(f"Invalid API endpoint: {err}") from err
|
|
|
|
# Pinned session closes DNS-rebinding TOCTOU and isolates cookies
|
|
# from other integrations sharing the same endpoint hostname.
|
|
# The integration owns this session: APIClient.shutdown() closes it
|
|
# on unload, the except handler below closes it on failed setup.
|
|
session = create_pinned_session(endpoint, resolved_ips)
|
|
# 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)
|
|
request_interval = config.get(CONF_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL)
|
|
api_timeout = config.get(CONF_API_TIMEOUT, DEFAULT_API_TIMEOUT)
|
|
max_tokens = config.get(CONF_MAX_TOKENS, DEFAULT_MAX_TOKENS)
|
|
temperature = config.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)
|
|
max_history_size = config.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY)
|
|
context_messages = config.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
|
|
disable_thinking = config.get(CONF_DISABLE_THINKING, DEFAULT_DISABLE_THINKING)
|
|
|
|
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")
|
|
|
|
_LOGGER.debug("Creating API client for %s with endpoint %s", api_provider, endpoint)
|
|
|
|
api_client = APIClient(
|
|
session=session,
|
|
endpoint=endpoint,
|
|
headers=headers,
|
|
api_provider=api_provider,
|
|
model=model,
|
|
api_timeout=api_timeout,
|
|
api_key=api_key,
|
|
)
|
|
|
|
coordinator = HATextAICoordinator(
|
|
hass=hass,
|
|
client=api_client,
|
|
model=model,
|
|
update_interval=request_interval,
|
|
instance_name=instance_name,
|
|
config_entry=entry,
|
|
max_tokens=max_tokens,
|
|
temperature=temperature,
|
|
max_history_size=max_history_size,
|
|
context_messages=context_messages,
|
|
api_timeout=api_timeout,
|
|
disable_thinking=disable_thinking,
|
|
)
|
|
|
|
# Initialize coordinator (directories, history, metrics)
|
|
await coordinator.async_initialize()
|
|
|
|
_LOGGER.debug("Created coordinator for %s", instance_name)
|
|
|
|
# Store coordinator
|
|
hass.data.setdefault(DOMAIN, {})
|
|
hass.data[DOMAIN][entry.entry_id] = coordinator
|
|
|
|
# A reload after the last entry was unloaded needs the services back.
|
|
_async_register_services(hass)
|
|
|
|
_LOGGER.debug("Stored coordinator in hass.data[%s][%s]", DOMAIN, entry.entry_id)
|
|
|
|
# Set up platforms
|
|
await hass.config_entries.async_forward_entry_setups(entry, PLATFORMS)
|
|
|
|
# Register update listener for options changes
|
|
entry.async_on_unload(entry.add_update_listener(async_update_options))
|
|
|
|
_LOGGER.debug("Setup completed for %s", instance_name)
|
|
|
|
return True
|
|
|
|
except Exception as err:
|
|
_LOGGER.exception("Error setting up HA Text AI: %s", err)
|
|
if session is not None and not session.closed:
|
|
await session.close()
|
|
raise
|
|
|
|
async def async_update_options(hass: HomeAssistant, entry: ConfigEntry) -> None:
|
|
"""Handle options update - reload the config entry."""
|
|
_LOGGER.info("Options updated for %s, reloading integration", entry.title)
|
|
await hass.config_entries.async_reload(entry.entry_id)
|
|
|
|
async def async_unload_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
|
|
"""Unload a config entry."""
|
|
try:
|
|
unload_ok = await hass.config_entries.async_unload_platforms(entry, PLATFORMS)
|
|
if unload_ok and entry.entry_id in hass.data[DOMAIN]:
|
|
coordinator = hass.data[DOMAIN].pop(entry.entry_id)
|
|
|
|
if hasattr(coordinator.client, 'shutdown'):
|
|
await coordinator.client.shutdown()
|
|
|
|
await coordinator.async_shutdown()
|
|
|
|
# When removing the last config entry, also unregister services and
|
|
# clear the domain bucket so HA doesn't show stale services in the UI.
|
|
if not hass.data.get(DOMAIN):
|
|
hass.data.pop(DOMAIN, None)
|
|
for service in (
|
|
SERVICE_ASK_QUESTION,
|
|
SERVICE_CLEAR_HISTORY,
|
|
SERVICE_GET_HISTORY,
|
|
SERVICE_SET_SYSTEM_PROMPT,
|
|
):
|
|
if hass.services.has_service(DOMAIN, service):
|
|
hass.services.async_remove(DOMAIN, service)
|
|
|
|
return unload_ok
|
|
|
|
except Exception as ex:
|
|
_LOGGER.exception("Error unloading entry: %s", str(ex))
|
|
return False
|