Files
Home-Assistant/custom_components/ai_agent_ha/agent.py
T
2026-07-20 22:52:35 -04:00

4416 lines
202 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""The AI Agent implementation with multiple provider support.
Example config:
ai_agent_ha:
ai_provider: openai # or 'llama', 'gemini', 'openrouter', 'anthropic', 'alter', 'zai', 'local_ollama', 'openai_compatible'
llama_token: "..."
openai_token: "..."
gemini_token: "..."
openrouter_token: "..."
anthropic_token: "..."
alter_token: "..."
zai_token: "..."
zai_endpoint: "general" # or 'coding' for z.ai (3× usage, 1/7 cost)
local_ollama_url: "http://localhost:11434/api/generate" # Required for local_ollama provider
openai_compatible_url: "http://example.com/v1/" or "http://localhost/v1/" # (Url must end with /v1/)
# Model configuration (optional, defaults will be used if not specified)
models:
openai: "gpt-3.5-turbo" # or "gpt-4", "gpt-4-turbo", etc.
llama: "Llama-4-Maverick-17B-128E-Instruct-FP8"
gemini: "gemini-2.5-flash" # or "gemini-2.5-pro", "gemini-2.0-flash", etc.
openrouter: "openai/gpt-4o" # or any model available on OpenRouter
anthropic: "claude-sonnet-4-5-20250929" # or "claude-sonnet-4-20250514", "claude-3-5-sonnet-20241022", "claude-3-opus-20240229", etc.
alter: "your-model-name" # model name for Alter API
zai: "glm-4.7" # model name for z.ai API (glm-4.7, glm-4.6, glm-4.5, etc.)
local_ollama: "llama3.2" # model name for local_ollama provider (optional if your API doesn't require it)
openai_compatible: "model unique-id or your-model-name" # model name for your OpenAI-compatible endpoint
"""
import asyncio
import json
import logging
import os
import shutil
import tempfile
import time
from datetime import datetime, timedelta
from typing import Any, Dict, List, Optional, Union
from urllib.parse import quote
import aiohttp
import yaml # type: ignore[import-untyped]
from homeassistant.core import HomeAssistant
from homeassistant.helpers.storage import Store
from homeassistant.util import dt as dt_util
from .const import CONF_OPENAI_BASE_URL, CONF_WEATHER_ENTITY, DOMAIN
_LOGGER = logging.getLogger(__name__)
# === Security Utilities ===
def sanitize_for_logging(data: Any, mask: str = "***REDACTED***") -> Any:
"""Sanitize sensitive data for safe logging.
Recursively masks sensitive fields like API keys, tokens, passwords, etc.
This prevents accidental exposure of credentials in debug logs.
Args:
data: The data structure to sanitize (dict, list, str, etc.)
mask: The string to use for masking sensitive values
Returns:
A sanitized copy of the data with sensitive fields masked
Example:
>>> config = {"openai_token": "sk-abc123", "ai_provider": "openai"}
>>> sanitize_for_logging(config)
{"openai_token": "***REDACTED***", "ai_provider": "openai"}
"""
# Sensitive field patterns (case-insensitive)
sensitive_patterns = {
"token",
"key",
"password",
"secret",
"credential",
"auth",
"authorization",
"api_key",
"apikey",
"llama_token",
"openai_token",
"gemini_token",
"anthropic_token",
"openrouter_token",
"alter_token",
"zai_token",
}
if isinstance(data, dict):
sanitized = {}
for key, value in data.items():
# Check if key matches any sensitive pattern
key_lower = str(key).lower()
is_sensitive = any(pattern in key_lower for pattern in sensitive_patterns)
if is_sensitive:
sanitized[key] = mask
else:
# Recursively sanitize nested structures
sanitized[key] = sanitize_for_logging(value, mask)
return sanitized
elif isinstance(data, list):
return [sanitize_for_logging(item, mask) for item in data]
elif isinstance(data, tuple):
return tuple(sanitize_for_logging(item, mask) for item in data)
else:
# Primitive types (str, int, bool, etc.) - return as-is
return data
# === AI Client Abstractions ===
class NonRetryableAIError(Exception):
"""Provider error that cannot succeed on retry (e.g. deterministic HTTP 4xx).
Raised for client errors like "prompt is too long" where re-sending the
identical payload is guaranteed to fail again (see issue #80). Rate
limiting (429) and request timeouts (408) stay retryable.
"""
class RateLimitedAIError(Exception):
"""Provider rate limit (HTTP 429) with an optional server-suggested wait.
Carries the provider's retry-after hint so the retry loop can wait long
enough for a per-minute token window to reset instead of burning all
retries in a few seconds (see issue #80).
"""
def __init__(self, message: str, retry_after: Optional[float] = None):
super().__init__(message)
self.retry_after = retry_after
class BaseAIClient:
async def get_response(self, messages, **kwargs):
raise NotImplementedError
class LocalOllamaClient(BaseAIClient):
"""Client for Ollama-style local models using /api/generate style endpoints."""
def __init__(self, url, model=""):
self.url = url
self.model = model
async def get_response(self, messages, **kwargs):
_LOGGER.debug(
"Making request to local Ollama API with model: '%s' at URL: %s",
self.model or "[NO MODEL SPECIFIED]",
self.url,
)
if not self.model:
_LOGGER.warning(
"No model specified for local Ollama API request. Some APIs (like Ollama) require a model name."
)
headers = {"Content-Type": "application/json"}
# Format user prompt from messages
prompt = ""
for message in messages:
role = message.get("role", "")
content = message.get("content", "")
# Simple formatting: prefixing each message with its role
if role == "system":
prompt += f"System: {content}\n\n"
elif role == "user":
prompt += f"User: {content}\n\n"
elif role == "assistant":
prompt += f"Assistant: {content}\n\n"
# Add final prompt prefix for the assistant's response
prompt += "Assistant: "
# Build a generic payload that works with most local Ollama-style API servers
payload = {
"prompt": prompt,
"stream": False, # Disable streaming to get a single complete response
# max_tokens omitted - let local model use its default capacity
}
# Add model if specified
if self.model:
payload["model"] = self.model
# Note: Payloads don't contain auth tokens (those are in headers), but may contain user prompts
_LOGGER.debug(
"Local Ollama API request payload: %s", json.dumps(payload, indent=2)
)
# Ollama-specific validation
if "model" not in payload or not payload["model"]:
_LOGGER.warning(
"Missing 'model' field in request to local Ollama API. This may cause issues with Ollama."
)
elif self.url and "ollama" in self.url.lower():
_LOGGER.debug(
"Detected Ollama URL, ensuring model is specified: %s",
payload.get("model"),
)
async with aiohttp.ClientSession() as session:
async with session.post(
self.url,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
if resp.status != 200:
error_text = await resp.text()
_LOGGER.error(
"Local Ollama API error %d: %s", resp.status, error_text
)
# Provide more specific error messages for common Ollama issues
if resp.status == 404:
if "model" in payload and payload["model"]:
raise Exception(
f"Model '{payload['model']}' not found. Please ensure the model is installed in Ollama using: ollama pull {payload['model']}"
)
else:
raise Exception(
"Local Ollama API endpoint not found. Please check the URL and ensure Ollama is running."
)
elif resp.status == 400:
raise Exception(
f"Bad request to local Ollama API. Error: {error_text}"
)
else:
raise Exception(
f"Local Ollama API error {resp.status}: {error_text}"
)
try:
response_text = await resp.text()
_LOGGER.debug(
"Local Ollama API response (first 200 chars): %s",
response_text[:200],
)
_LOGGER.debug("Local Ollama API response status: %d", resp.status)
# Sanitize headers to avoid logging any auth tokens
_LOGGER.debug(
"Local Ollama API response headers: %s",
sanitize_for_logging(dict(resp.headers)),
)
# Try to parse as JSON
try:
data = json.loads(response_text)
# Try common response formats
# Ollama format - return only the response text
if "response" in data:
response_content = data["response"]
_LOGGER.debug(
"Extracted response content: %s",
(
response_content[:100]
if response_content
else "[EMPTY]"
),
)
# Check if response is empty or None
if not response_content or response_content.strip() == "":
_LOGGER.warning(
"Ollama returned empty response. Full data: %s",
data,
)
# Check if this is a loading response
if data.get("done_reason") == "load":
_LOGGER.warning(
"Ollama is still loading the model. Please wait and try again."
)
return json.dumps(
{
"request_type": "final_response",
"response": "The AI model is still loading. Please wait a moment and try again.",
}
)
elif data.get("done") is False:
_LOGGER.warning(
"Ollama response indicates it's not done yet."
)
return json.dumps(
{
"request_type": "final_response",
"response": "The AI is still processing your request. Please try again.",
}
)
else:
return json.dumps(
{
"request_type": "final_response",
"response": "The AI returned an empty response. Please try rephrasing your question.",
}
)
# Check if the response looks like JSON
response_content = response_content.strip()
if response_content.startswith(
"{"
) and response_content.endswith("}"):
try:
# Validate that it's actually JSON and contains valid request_type
parsed_json = json.loads(response_content)
if (
isinstance(parsed_json, dict)
and "request_type" in parsed_json
):
_LOGGER.debug(
"Local Ollama model provided valid JSON response"
)
return response_content
else:
_LOGGER.debug(
"JSON missing request_type, treating as plain text"
)
except json.JSONDecodeError:
_LOGGER.debug(
"Invalid JSON from local Ollama model, treating as plain text"
)
pass
# If it's plain text, wrap it in the expected JSON format
wrapped_response = {
"request_type": "final_response",
"response": response_content,
}
_LOGGER.debug("Wrapped plain text response in JSON format")
return json.dumps(wrapped_response)
# OpenAI-like format
elif "choices" in data and len(data["choices"]) > 0:
choice = data["choices"][0]
if "message" in choice and "content" in choice["message"]:
content = choice["message"]["content"]
elif "text" in choice:
content = choice["text"]
else:
content = str(data)
# Check if it's valid JSON with request_type
content = content.strip()
if content.startswith("{") and content.endswith("}"):
try:
parsed_json = json.loads(content)
if (
isinstance(parsed_json, dict)
and "request_type" in parsed_json
):
_LOGGER.debug(
"Local Ollama model provided valid JSON response (OpenAI format)"
)
return content
else:
_LOGGER.debug(
"JSON missing request_type, treating as plain text (OpenAI format)"
)
except json.JSONDecodeError:
_LOGGER.debug(
"Invalid JSON from local Ollama model, treating as plain text (OpenAI format)"
)
pass
# Wrap in expected format if plain text
wrapped_response = {
"request_type": "final_response",
"response": content,
}
return json.dumps(wrapped_response)
# Generic content field
elif "content" in data:
content = data["content"]
content = content.strip()
if content.startswith("{") and content.endswith("}"):
try:
parsed_json = json.loads(content)
if (
isinstance(parsed_json, dict)
and "request_type" in parsed_json
):
_LOGGER.debug(
"Local Ollama model provided valid JSON response (generic format)"
)
return content
else:
_LOGGER.debug(
"JSON missing request_type, treating as plain text (generic format)"
)
except json.JSONDecodeError:
_LOGGER.debug(
"Invalid JSON from local Ollama model, treating as plain text (generic format)"
)
pass
wrapped_response = {
"request_type": "final_response",
"response": content,
}
return json.dumps(wrapped_response)
# Handle case where no standard fields are found
_LOGGER.warning(
"No standard response fields found in local Ollama API response. Full response: %s",
data,
)
# Check for Ollama-specific edge cases
if data.get("done_reason") == "load":
return json.dumps(
{
"request_type": "final_response",
"response": "The AI model is still loading. Please wait a moment and try again.",
}
)
elif data.get("done") is False:
return json.dumps(
{
"request_type": "final_response",
"response": "The AI is still processing your request. Please try again.",
}
)
elif "message" in data:
# Some APIs use "message" field
message_content = data["message"]
if (
isinstance(message_content, dict)
and "content" in message_content
):
content = message_content["content"]
else:
content = str(message_content)
return json.dumps(
{"request_type": "final_response", "response": content}
)
# Return the whole data as string if we can't find a specific field
return json.dumps(
{
"request_type": "final_response",
"response": f"Received unexpected response format from local Ollama API: {str(data)}",
}
)
except json.JSONDecodeError:
# If not JSON, check if it's a JSON response that got corrupted by wrapping
response_text = response_text.strip()
if response_text.startswith("{") and response_text.endswith(
"}"
):
try:
parsed_json = json.loads(response_text)
if (
isinstance(parsed_json, dict)
and "request_type" in parsed_json
):
_LOGGER.debug(
"Local Ollama model provided valid JSON response (direct)"
)
return response_text
except json.JSONDecodeError:
pass
# If not valid JSON, wrap the raw text in expected format
_LOGGER.debug("Response is not JSON, wrapping plain text")
wrapped_response = {
"request_type": "final_response",
"response": response_text,
}
return json.dumps(wrapped_response)
except Exception as e:
_LOGGER.error(
"Failed to parse local Ollama API response: %s", str(e)
)
raise Exception(
f"Failed to parse local Ollama API response: {str(e)}"
)
class OpenaiCompatibleClient(BaseAIClient):
"""Client for OpenAI-compatible endpoints (e.g., LM Studio, vLLM, etc.).
Expected URL format: http://example.com/v1/
This client sends chat completions requests to: {url}/chat/completions
No API key is required by default, but can be provided if needed.
"""
def __init__(self, base_url, model="", api_key=None):
# Ensure base_url ends with /v1/ style segment
base_url = (base_url or "").strip().rstrip("/")
if not base_url:
raise Exception("openai_compatible_url is required and must not be empty")
self.base_url = base_url
# Derive chat completions endpoint
self.api_url = f"{self.base_url}/chat/completions"
self.model = model
self.api_key = api_key or "" # Optional; many local endpoints dont require it
async def get_response(self, messages, **kwargs):
_LOGGER.debug(
"Making request to OpenAI-compatible endpoint at %s with model: %s",
self.api_url,
self.model or "[NO MODEL SPECIFIED]",
)
if not self.model:
_LOGGER.warning(
"No model specified for OpenAI-compatible request. Some servers require a model name."
)
headers = {
"Content-Type": "application/json",
}
# Add Authorization header only if an API key is set
if self.api_key:
headers["Authorization"] = f"Bearer {self.api_key}"
payload = {
"model": self.model,
"messages": messages,
"temperature": 0.7,
"top_p": 0.9,
# max_tokens omitted - let server/model use its default capacity
}
_LOGGER.debug(
"OpenAI-compatible request payload: %s",
json.dumps(payload, indent=2),
)
async with aiohttp.ClientSession() as session:
async with session.post(
self.api_url,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
response_text = await resp.text()
_LOGGER.debug("OpenAI-compatible API response status: %d", resp.status)
_LOGGER.debug(
"OpenAI-compatible API response (first 500 chars): %s",
response_text[:500],
)
if resp.status != 200:
_LOGGER.error(
"OpenAI-compatible API error %d: %s",
resp.status,
response_text,
)
raise Exception(
f"OpenAI-compatible API error {resp.status}: {response_text}"
)
try:
data = json.loads(response_text)
except json.JSONDecodeError as e:
_LOGGER.error(
"Failed to parse OpenAI-compatible response as JSON: %s", str(e)
)
raise Exception(
f"Invalid JSON response from OpenAI-compatible API: {response_text[:200]}"
)
# Extract text from OpenAI-compatible response
choices = data.get("choices", [])
if choices and "message" in choices[0]:
content = choices[0]["message"].get("content", "")
if not content:
_LOGGER.warning(
"OpenAI-compatible API returned empty content in message"
)
_LOGGER.debug(
"Full OpenAI-compatible API response: %s",
json.dumps(data, indent=2),
)
return content
else:
_LOGGER.warning(
"OpenAI-compatible API response missing expected structure"
)
_LOGGER.debug(
"Full OpenAI-compatible API response: %s",
json.dumps(data, indent=2),
)
return str(data)
class LlamaClient(BaseAIClient):
def __init__(self, token, model="Llama-4-Maverick-17B-128E-Instruct-FP8"):
self.token = token
self.model = model
self.api_url = "https://api.llama.com/v1/chat/completions"
async def get_response(self, messages, **kwargs):
_LOGGER.debug("Making request to Llama API with model: %s", self.model)
headers = {
"Authorization": f"Bearer {self.token}",
"Content-Type": "application/json",
}
payload = {
"model": self.model,
"messages": messages,
"temperature": 0.7,
"top_p": 0.9,
# max_tokens omitted - let Llama use the model's default capacity
}
_LOGGER.debug("Llama request payload: %s", json.dumps(payload, indent=2))
async with aiohttp.ClientSession() as session:
async with session.post(
self.api_url,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
if resp.status != 200:
error_text = await resp.text()
_LOGGER.error("Llama API error %d: %s", resp.status, error_text)
raise Exception(f"Llama API error {resp.status}")
data = await resp.json()
# Extract text from Llama response
completion = data.get("completion_message", {})
content = completion.get("content", {})
return content.get("text", str(data))
async def fetch_openai_models(base_url, api_key, timeout=10):
"""Fetch available OpenAI models dynamically.
Returns a list of model IDs suitable for chat (gpt-*, o*). On failure,
returns a small safe fallback list.
"""
if not base_url:
base_url = "https://api.openai.com/v1"
url = f"{base_url.rstrip('/')}/models"
fallback_models = [
"gpt-4.1-mini",
"gpt-4o-mini",
"o4-mini",
]
try:
async with aiohttp.ClientSession() as session:
async with session.get(
url,
headers={"Authorization": f"Bearer {api_key}"},
timeout=aiohttp.ClientTimeout(total=timeout),
) as resp:
if resp.status != 200:
_LOGGER.warning(
"Failed to fetch OpenAI models (status=%d), using fallback list",
resp.status,
)
return fallback_models
data = await resp.json()
models = data.get("data", [])
if not isinstance(models, list):
return fallback_models
# Filter likely chat-capable models
chat_models = []
for m in models:
mid = m.get("id", "")
if isinstance(mid, str) and (
mid.startswith("gpt-") or mid.startswith("o")
):
chat_models.append(mid)
if not chat_models:
return fallback_models
chat_models.sort()
_LOGGER.debug("Fetched %d OpenAI chat models", len(chat_models))
return chat_models
except Exception as e:
_LOGGER.warning("Error fetching OpenAI models, using fallback list: %s", e)
return fallback_models
async def fetch_gemini_models(api_key, timeout=10):
"""Fetch available Gemini models dynamically.
Returns a list of model IDs suitable for chat. On failure,
returns a small safe fallback list.
"""
if not api_key:
return [
"gemini-2.5-flash",
"gemini-2.5-pro",
]
url = f"https://generativelanguage.googleapis.com/v1beta/models?key={api_key}"
fallback_models = [
"gemini-2.5-flash",
"gemini-2.5-pro",
]
try:
async with aiohttp.ClientSession() as session:
async with session.get(
url,
timeout=aiohttp.ClientTimeout(total=timeout),
) as resp:
if resp.status != 200:
_LOGGER.warning(
"Failed to fetch Gemini models (status=%d), using fallback list",
resp.status,
)
return fallback_models
data = await resp.json()
models = data.get("models", [])
if not isinstance(models, list):
return fallback_models
# Filter likely chat-capable models
chat_models = []
for m in models:
mid = m.get("name", "")
if isinstance(mid, str) and "gemini" in mid.lower():
# Extract model ID from name like "models/gemini-2.5-flash"
model_id = mid.split("/")[-1] if "/" in mid else mid
chat_models.append(model_id)
if not chat_models:
return fallback_models
chat_models.sort()
_LOGGER.debug("Fetched %d Gemini models", len(chat_models))
return chat_models
except Exception as e:
_LOGGER.warning("Error fetching Gemini models, using fallback list: %s", e)
return fallback_models
async def fetch_openai_compatible_models(base_url, api_key=None, timeout=10):
"""Fetch available models from an OpenAI-compatible endpoint.
Many local/self-hosted endpoints (LM Studio, vLLM, etc.) support /v1/models.
If the endpoint does not support it, fall back to ["Custom..."] only.
"""
if not base_url:
return ["Custom..."]
url = f"{base_url.rstrip('/')}/models"
try:
async with aiohttp.ClientSession() as session:
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
async with session.get(
url,
headers=headers,
timeout=aiohttp.ClientTimeout(total=timeout),
) as resp:
if resp.status not in (200, 201):
_LOGGER.debug(
"OpenAI-compatible endpoint did not support /v1/models (status=%d)",
resp.status,
)
return ["Custom..."]
data = await resp.json()
models = data.get("data", [])
if not isinstance(models, list):
return ["Custom..."]
# Collect all model IDs
model_ids = []
for m in models:
mid = m.get("id", "")
if isinstance(mid, str) and mid:
model_ids.append(mid)
if not model_ids:
return ["Custom..."]
model_ids.sort()
_LOGGER.debug(
"Fetched %d models from OpenAI-compatible endpoint", len(model_ids)
)
return model_ids
except Exception as e:
_LOGGER.debug(
"Error fetching models from OpenAI-compatible endpoint, using Custom only: %s",
e,
)
return ["Custom..."]
class OpenAIClient(BaseAIClient):
def __init__(self, token, model="gpt-4.1-mini", base_url=None):
self.token = token
self.model = model
# Default endpoint is OpenAI's Responses API. If the user has pointed Base URL
# at a third-party "OpenAI-compatible" gateway (Open WebUI, LM Studio, vLLM,
# LiteLLM, ...), those servers only implement /chat/completions, so switch
# to Chat Completions when the base URL is not api.openai.com.
if base_url and base_url.strip():
base = base_url.strip().rstrip("/")
if "api.openai.com" in base:
self.api_url = f"{base}/responses"
self.use_chat_completions = False
else:
self.api_url = f"{base}/chat/completions"
self.use_chat_completions = True
else:
self.api_url = "https://api.openai.com/v1/responses"
self.use_chat_completions = False
def _is_restricted_model(self):
"""Check if the model has restricted parameters (no temperature, top_p, etc.)."""
# Models that don't support temperature, top_p and other parameters
restricted_models = ["o3-mini", "o3", "o1-mini", "o1-preview", "o1", "gpt-5"]
model_lower = self.model.lower()
return any(model_id in model_lower for model_id in restricted_models)
async def get_response(self, messages, **kwargs):
_LOGGER.debug("Making request to OpenAI API with model: %s", self.model)
# Validate token
if not self.token or not self.token.startswith("sk-"):
raise Exception("Invalid OpenAI API key format")
headers = {
"Authorization": f"Bearer {self.token}",
"Content-Type": "application/json",
}
if self.use_chat_completions:
# Chat Completions: forward messages as-is.
payload = {
"model": self.model,
"messages": messages,
}
if not self._is_restricted_model():
payload["temperature"] = 0.7
payload["top_p"] = 0.9
else:
# Responses API: flatten messages into a single input string.
parts = []
for msg in messages:
role = (msg.get("role") or "user").lower()
content = msg.get("content") or ""
if role == "system":
parts.append(f"System: {content}")
elif role == "user":
parts.append(f"User: {content}")
elif role == "assistant":
parts.append(f"Assistant: {content}")
else:
parts.append(f"{role.capitalize()}: {content}")
payload = {
"model": self.model,
"input": "\n\n".join(parts),
}
_LOGGER.debug("OpenAI request payload: %s", json.dumps(payload, indent=2))
async with aiohttp.ClientSession() as session:
async with session.post(
self.api_url,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
response_text = await resp.text()
_LOGGER.debug("OpenAI API response status: %d", resp.status)
_LOGGER.debug("OpenAI API response: %s", response_text[:500])
if resp.status != 200:
_LOGGER.error("OpenAI API error %d: %s", resp.status, response_text)
raise Exception(f"OpenAI API error {resp.status}: {response_text}")
try:
data = json.loads(response_text)
except json.JSONDecodeError as e:
_LOGGER.error("Failed to parse OpenAI response as JSON: %s", str(e))
raise Exception(
f"Invalid JSON response from OpenAI: {response_text[:200]}"
)
if self.use_chat_completions:
# Chat Completions response shape
choices = data.get("choices", [])
if choices and "message" in choices[0]:
return choices[0]["message"].get("content", "") or ""
_LOGGER.warning("OpenAI response missing expected structure")
_LOGGER.debug(
"Full OpenAI response: %s", json.dumps(data, indent=2)
)
return str(data)
# Extract text from OpenAI Responses API.
# Primary field: output_text is an SDK-only convenience property
# and is normally absent from the raw HTTP body, so it is only a
# fast path when a gateway happens to include it as a string.
content = data.get("output_text")
if isinstance(content, str) and content:
return content
# Fallback: the Responses API returns
# output: [ { type: "message", content: [ { type: "output_text",
# text: "..." }, ... ] }, ... ]
# plus, for reasoning models, leading items (e.g. type "reasoning")
# that carry no text. Concatenate every output_text block we find.
# NOTE: item["content"] is a LIST here, so it must never be
# returned directly (that caused issue #75: 'list' object has no
# attribute 'strip').
output = data.get("output")
if isinstance(output, list):
text_parts = []
for item in output:
if not isinstance(item, dict):
continue
content_blocks = item.get("content")
if isinstance(content_blocks, list):
for block in content_blocks:
if isinstance(block, dict) and isinstance(
block.get("text"), str
):
text_parts.append(block["text"])
elif isinstance(content_blocks, str) and content_blocks:
text_parts.append(content_blocks)
elif isinstance(item.get("text"), str):
text_parts.append(item["text"])
if text_parts:
return "".join(text_parts)
# Last resort: return full response as string
_LOGGER.warning("OpenAI response missing expected structure")
_LOGGER.debug("Full OpenAI response: %s", json.dumps(data, indent=2))
return str(data)
class GeminiClient(BaseAIClient):
def __init__(self, token, model="gemini-2.5-flash"):
self.token = token.strip() if token else token # Strip whitespace from token
self.model = model
# Use v1beta for all models as per Google's current API documentation
# All Gemini 2.0/2.5 models are available on v1beta endpoint
self.api_url = f"https://generativelanguage.googleapis.com/v1beta/models/{model}:generateContent"
async def get_response(self, messages, **kwargs):
_LOGGER.debug("Making request to Gemini API with model: %s", self.model)
# Validate token
if not self.token:
raise Exception("Missing Gemini API key")
headers = {"Content-Type": "application/json"}
# Convert OpenAI-style messages to Gemini format
gemini_contents = []
for message in messages:
role = message.get("role", "user")
content = message.get("content", "")
if role == "system":
# Gemini doesn't have a system role, so we prepend it to the first user message
if not gemini_contents:
gemini_contents.append(
{"role": "user", "parts": [{"text": f"System: {content}"}]}
)
else:
# Add system message as user message
gemini_contents.append(
{"role": "user", "parts": [{"text": f"System: {content}"}]}
)
elif role == "user":
gemini_contents.append({"role": "user", "parts": [{"text": content}]})
elif role == "assistant":
gemini_contents.append({"role": "model", "parts": [{"text": content}]})
payload = {
"contents": gemini_contents,
"generationConfig": {
"temperature": 0.7,
"topP": 0.9,
# maxOutputTokens omitted - let Gemini use model's maximum capacity
},
}
# Add API key as query parameter (URL encoded)
url_with_key = f"{self.api_url}?key={quote(self.token)}"
_LOGGER.debug("Gemini request payload: %s", json.dumps(payload, indent=2))
async with aiohttp.ClientSession() as session:
async with session.post(
url_with_key,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
response_text = await resp.text()
_LOGGER.debug("Gemini API response status: %d", resp.status)
_LOGGER.debug("Gemini API response: %s", response_text[:500])
if resp.status != 200:
_LOGGER.error("Gemini API error %d: %s", resp.status, response_text)
raise Exception(f"Gemini API error {resp.status}: {response_text}")
try:
data = json.loads(response_text)
except json.JSONDecodeError as e:
_LOGGER.error("Failed to parse Gemini response as JSON: %s", str(e))
raise Exception(
f"Invalid JSON response from Gemini: {response_text[:200]}"
)
# Log token usage for debugging, especially for Gemini 2.5 extended thinking
usage_metadata = data.get("usageMetadata", {})
if usage_metadata:
_LOGGER.debug(
"Gemini token usage - prompt: %d, total: %d, thoughts: %d",
usage_metadata.get("promptTokenCount", 0),
usage_metadata.get("totalTokenCount", 0),
usage_metadata.get("thoughtsTokenCount", 0),
)
# Extract text from Gemini response
candidates = data.get("candidates", [])
if candidates and "content" in candidates[0]:
# Check finish reason for potential issues
finish_reason = candidates[0].get("finishReason", "")
if finish_reason == "MAX_TOKENS":
_LOGGER.warning(
"Gemini response truncated due to MAX_TOKENS limit. "
"Thoughts used: %d tokens. Consider increasing maxOutputTokens.",
usage_metadata.get("thoughtsTokenCount", 0),
)
parts = candidates[0]["content"].get("parts", [])
if parts:
content = parts[0].get("text", "")
if not content:
_LOGGER.warning("Gemini returned empty text content")
_LOGGER.debug(
"Full Gemini response: %s", json.dumps(data, indent=2)
)
return content
else:
_LOGGER.warning("Gemini response missing parts")
_LOGGER.debug(
"Full Gemini response: %s", json.dumps(data, indent=2)
)
else:
_LOGGER.warning("Gemini response missing expected structure")
_LOGGER.debug(
"Full Gemini response: %s", json.dumps(data, indent=2)
)
return str(data)
class AnthropicClient(BaseAIClient):
def __init__(self, token, model="claude-sonnet-4-5-20250929"):
self.token = token
self.model = model
self.api_url = "https://api.anthropic.com/v1/messages"
async def get_response(self, messages, **kwargs):
_LOGGER.debug("Making request to Anthropic API with model: %s", self.model)
headers = {
"x-api-key": self.token,
"Content-Type": "application/json",
"anthropic-version": "2023-06-01",
}
# Convert OpenAI-style messages to Anthropic format
system_message = None
anthropic_messages = []
for message in messages:
role = message.get("role", "user")
content = message.get("content", "")
if role == "system":
# Anthropic uses a separate system parameter
system_message = content
elif role == "user":
anthropic_messages.append({"role": "user", "content": content})
elif role == "assistant":
anthropic_messages.append({"role": "assistant", "content": content})
payload = {
"model": self.model,
"max_tokens": 8192, # Maximum for Anthropic Claude models
"temperature": 0.7,
"messages": anthropic_messages,
}
# Add system message if present
if system_message:
payload["system"] = system_message
_LOGGER.debug("Anthropic request payload: %s", json.dumps(payload, indent=2))
async with aiohttp.ClientSession() as session:
async with session.post(
self.api_url,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
if resp.status != 200:
error_text = await resp.text()
_LOGGER.error("Anthropic API error %d: %s", resp.status, error_text)
# Surface the API's own error message (e.g. "prompt is too
# long: N tokens > 200000 maximum") instead of a bare status
# code so the user can see why the request failed (issue #80).
try:
detail = json.loads(error_text)["error"]["message"]
except (ValueError, KeyError, TypeError):
detail = error_text[:300]
message = f"Anthropic API error {resp.status}: {detail}"
if resp.status == 429:
# Honor the server's retry-after hint so the retry
# loop waits out per-minute token windows.
retry_after = None
try:
header = resp.headers.get("retry-after")
if header is not None:
retry_after = float(header)
except (TypeError, ValueError):
retry_after = None
raise RateLimitedAIError(message, retry_after=retry_after)
if 400 <= resp.status < 500 and resp.status != 408:
# Deterministic client error - retrying the identical
# payload cannot succeed, so don't burn retries on it.
raise NonRetryableAIError(message)
raise Exception(message)
data = await resp.json()
# Extract text from Anthropic response
content_blocks = data.get("content", [])
if content_blocks and isinstance(content_blocks, list):
# Get the text from the first content block
for block in content_blocks:
if block.get("type") == "text":
return block.get("text", str(data))
return str(data)
class OpenRouterClient(BaseAIClient):
def __init__(self, token, model="openai/gpt-4o"):
self.token = token
self.model = model
self.api_url = "https://openrouter.ai/api/v1/chat/completions"
async def get_response(self, messages, **kwargs):
_LOGGER.debug("Making request to OpenRouter API with model: %s", self.model)
headers = {
"Authorization": f"Bearer {self.token}",
"Content-Type": "application/json",
"HTTP-Referer": "https://home-assistant.io", # Optional for OpenRouter rankings
"X-Title": "Home Assistant AI Agent", # Optional for OpenRouter rankings
}
payload = {
"model": self.model,
"messages": messages,
"temperature": 0.7,
"top_p": 0.9,
# max_tokens omitted - let OpenRouter use the model's maximum capacity
}
_LOGGER.debug("OpenRouter request payload: %s", json.dumps(payload, indent=2))
async with aiohttp.ClientSession() as session:
async with session.post(
self.api_url,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
if resp.status != 200:
error_text = await resp.text()
_LOGGER.error(
"OpenRouter API error %d: %s", resp.status, error_text
)
raise Exception(f"OpenRouter API error {resp.status}")
data = await resp.json()
# Extract text from OpenRouter response (OpenAI-compatible format)
choices = data.get("choices", [])
if not choices:
_LOGGER.warning("OpenRouter response missing choices")
_LOGGER.debug(
"Full OpenRouter response: %s", json.dumps(data, indent=2)
)
return str(data)
if choices and "message" in choices[0]:
return choices[0]["message"].get("content", str(data))
return str(data)
class AlterClient(BaseAIClient):
def __init__(self, token, model=""):
self.token = token
self.model = model
self.api_url = "https://alterhq.com/api/v1/chat/completions"
async def get_response(self, messages, **kwargs):
_LOGGER.debug("Making request to Alter API with model: %s", self.model)
headers = {
"Authorization": f"Bearer {self.token}",
"Content-Type": "application/json",
}
payload = {
"model": self.model,
"messages": messages,
"temperature": 0.7,
"top_p": 0.9,
}
_LOGGER.debug("Alter request payload: %s", json.dumps(payload, indent=2))
async with aiohttp.ClientSession() as session:
async with session.post(
self.api_url,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
if resp.status != 200:
error_text = await resp.text()
_LOGGER.error("Alter API error %d: %s", resp.status, error_text)
raise Exception(f"Alter API error {resp.status}")
data = await resp.json()
# Extract text from Alter response (OpenAI-compatible format)
choices = data.get("choices", [])
if not choices:
_LOGGER.warning("Alter response missing choices")
_LOGGER.debug("Full Alter response: %s", json.dumps(data, indent=2))
return str(data)
if choices and "message" in choices[0]:
return choices[0]["message"].get("content", str(data))
return str(data)
class ZaiClient(BaseAIClient):
def __init__(self, token, model="", endpoint_type="general"):
self.token = token
self.model = model
self.endpoint_type = endpoint_type
# General endpoint: https://api.z.ai/api/paas/v4/chat/completions
# Coding endpoint: https://api.z.ai/api/coding/paas/v4/chat/completions
if endpoint_type == "coding":
self.api_url = "https://api.z.ai/api/coding/paas/v4/chat/completions"
else:
self.api_url = "https://api.z.ai/api/paas/v4/chat/completions"
async def get_response(self, messages, **kwargs):
_LOGGER.debug(
"Making request to z.ai API with model: %s, endpoint: %s",
self.model,
self.endpoint_type,
)
headers = {
"Authorization": f"Bearer {self.token}",
"Content-Type": "application/json",
}
payload = {
"model": self.model,
"messages": messages,
"temperature": 0.7,
"top_p": 0.9,
}
_LOGGER.debug("z.ai request payload: %s", json.dumps(payload, indent=2))
async with aiohttp.ClientSession() as session:
async with session.post(
self.api_url,
headers=headers,
json=payload,
timeout=aiohttp.ClientTimeout(total=300),
) as resp:
if resp.status != 200:
error_text = await resp.text()
_LOGGER.error("z.ai API error %d: %s", resp.status, error_text)
raise Exception(f"z.ai API error {resp.status}")
data = await resp.json()
# Extract text from z.ai response (OpenAI-compatible format)
choices = data.get("choices", [])
if not choices:
_LOGGER.warning("z.ai response missing choices")
_LOGGER.debug("Full z.ai response: %s", json.dumps(data, indent=2))
return str(data)
if choices and "message" in choices[0]:
return choices[0]["message"].get("content", str(data))
return str(data)
# === Main Agent ===
class AiAgentHaAgent:
"""Agent for handling queries with dynamic data requests and multiple AI providers."""
SYSTEM_PROMPT = {
"role": "system",
"content": (
"You are an AI assistant integrated with Home Assistant.\n"
"You can request specific data by using only these commands:\n"
"- get_entity_state(entity_id): Get state of a specific entity\n"
"- get_entities_by_domain(domain): Get all entities in a domain\n"
"- get_entities_by_device_class(device_class, domain?): Get entities with specific device_class (e.g., 'temperature', 'humidity', 'motion')\n"
"- get_climate_related_entities(): Get all climate-related entities (climate.* entities + temperature/humidity sensors)\n"
"- get_entities_by_area(area_id): Get all entities in a specific area\n"
"- get_entities(area_id or area_ids): Get entities by area(s) - supports single area_id or list of area_ids\n"
" Use as: get_entities(area_ids=['area1', 'area2']) for multiple areas or get_entities(area_id='single_area')\n"
"- get_calendar_events(entity_id?): Get calendar events\n"
"- get_automations(): Get all automations\n"
"- get_weather_data(): Get current weather and forecast data\n"
"- get_entity_registry(): Get entity registry entries (now includes device_class, state_class, unit_of_measurement)\n"
"- get_device_registry(): Get device registry entries\n"
"- get_area_registry(): Get room/area information\n"
"- get_history(entity_id, hours): Get historical state changes\n"
"- get_person_data(): Get person tracking information\n"
"- get_statistics(entity_id): Get sensor statistics\n"
"- get_scenes(): Get scene configurations\n"
"- get_dashboards(): Get list of all dashboards\n"
"- get_dashboard_config(dashboard_url): Get configuration of a specific dashboard\n"
"- set_entity_state(entity_id, state, attributes?): Set state of an entity (e.g., turn on/off lights, open/close covers)\n"
"- call_service(domain, service, target?, service_data?): Call any Home Assistant service directly\n"
"- create_automation(automation): Create a new automation with the provided configuration\n"
"- create_dashboard(dashboard_config): Create a new dashboard with the provided configuration\n"
"- update_dashboard(dashboard_url, dashboard_config): Update an existing dashboard configuration\n\n"
"IMPORTANT DEVICE_CLASS GUIDANCE:\n"
"- Many sensors have a 'device_class' attribute (temperature, humidity, motion, etc.)\n"
"- Use get_climate_related_entities() for climate dashboards (includes climate.* entities and temperature/humidity sensors)\n"
"- Use get_entities_by_device_class(device_class) to filter by device_class (e.g., 'temperature', 'humidity', 'motion')\n"
"- For climate dashboards, use history-graph and gauge cards for temperature/humidity sensors\n\n"
"DASHBOARD CREATION:\n"
"When a user asks to create a dashboard:\n"
"1. Gather entities using get_climate_related_entities() or other get_* commands\n"
"2. Respond with JSON using request_type: 'dashboard_suggestion' (NEVER use 'final_response'!)\n"
"3. Use Lovelace JSON format (NOT YAML!)\n"
"4. Example response structure:\n"
'{"request_type": "dashboard_suggestion", "message": "Dashboard created", "dashboard": {"title": "...", "views": [...]}}\n'
"5. Do NOT include YAML, markdown, or code blocks - only pure JSON\n\n"
"IMPORTANT AREA/FLOOR GUIDANCE:\n"
"- When users ask for entities from a specific floor, use get_area_registry() first\n"
"- Areas have both 'area_id' and 'floor_id' - these are different concepts\n"
"- Filter areas by their floor_id to find all areas on a specific floor\n"
"- Use get_entities() with area_ids parameter to get entities from multiple areas efficiently\n"
"- Example: get_entities(area_ids=['area1', 'area2', 'area3']) for multiple areas at once\n"
"- This is more efficient than calling get_entities_by_area() multiple times\n\n"
"AUTOMATION CREATION:\n"
"When creating automations, request entities first to know the entity IDs.\n"
"For days, use: ['fri', 'mon', 'sat', 'sun', 'thu', 'tue', 'wed']\n\n"
"RESPONSE FORMATS - You must ALWAYS respond with valid JSON:\n\n"
"For automations:\n"
"{\n"
' "request_type": "automation_suggestion",\n'
' "message": "I\'ve created an automation that might help you. Would you like me to create it?",\n'
' "automation": {\n'
' "alias": "Name of the automation",\n'
' "description": "Description of what the automation does",\n'
' "trigger": [...], // Array of trigger conditions\n'
' "condition": [...], // Optional array of conditions\n'
' "action": [...] // Array of actions to perform\n'
" }\n"
"}\n\n"
"For dashboards (WHEN USER ASKS TO CREATE A DASHBOARD):\n"
"{\n"
' "request_type": "dashboard_suggestion",\n'
' "message": "Description of the dashboard you created",\n'
' "dashboard": {\n'
' "title": "Dashboard Title",\n'
' "url_path": "url-path",\n'
' "icon": "mdi:icon-name",\n'
' "show_in_sidebar": true,\n'
' "views": [{\n'
' "title": "View Title",\n'
' "cards": [...]\n'
" }]\n"
" }\n"
"}\n\n"
"For data requests, use this exact JSON format:\n"
"{\n"
' "request_type": "data_request",\n'
' "request": "command_name",\n'
' "parameters": {...}\n'
"}\n"
'For get_entities with multiple areas: {"request_type": "get_entities", "parameters": {"area_ids": ["area1", "area2"]}}\n'
'For get_entities with single area: {"request_type": "get_entities", "parameters": {"area_id": "single_area"}}\n\n'
"For service calls, use this exact JSON format:\n"
"{\n"
' "request_type": "call_service",\n'
' "domain": "light",\n'
' "service": "turn_on",\n'
' "target": {"entity_id": ["entity1", "entity2"]},\n'
' "service_data": {"brightness": 255}\n'
"}\n\n"
"For answering questions (NOT creating dashboards/automations):\n"
"{\n"
' "request_type": "final_response",\n'
' "response": "your answer to the user"\n'
"}\n\n"
"IMPORTANT: Use 'dashboard_suggestion' when creating dashboards, NOT 'final_response'!\n\n"
"CRITICAL FORMATTING RULES:\n"
"- You must ALWAYS respond with ONLY a valid JSON object\n"
"- DO NOT include any text before the JSON\n"
"- DO NOT include any text after the JSON\n"
"- DO NOT include explanations or descriptions outside the JSON\n"
"- Your entire response must be parseable as JSON\n"
"- Use the 'message' field inside the JSON for user-facing text\n"
"- NEVER mix regular text with JSON in your response\n\n"
"WRONG: 'I'll create this for you. {\"request_type\": ...}'\n"
'CORRECT: \'{"request_type": "dashboard_suggestion", "message": "I\'ll create this for you.", ...}\''
),
}
SYSTEM_PROMPT_LOCAL = {
"role": "system",
"content": (
"You are an AI assistant integrated with Home Assistant.\n"
"You can request specific data by using only these commands:\n"
"- get_entity_state(entity_id): Get state of a specific entity\n"
"- get_entities_by_domain(domain): Get all entities in a domain\n"
"- get_entities_by_device_class(device_class, domain?): Get entities with specific device_class (e.g., 'temperature', 'humidity', 'motion')\n"
"- get_climate_related_entities(): Get all climate-related entities (climate.* entities + temperature/humidity sensors)\n"
"- get_entities_by_area(area_id): Get all entities in a specific area\n"
"- get_entities(area_id or area_ids): Get entities by area(s) - supports single area_id or list of area_ids\n"
" Use as: get_entities(area_ids=['area1', 'area2']) for multiple areas or get_entities(area_id='single_area')\n"
"- get_calendar_events(entity_id?): Get calendar events\n"
"- get_automations(): Get all automations\n"
"- get_weather_data(): Get current weather and forecast data\n"
"- get_entity_registry(): Get entity registry entries (now includes device_class, state_class, unit_of_measurement)\n"
"- get_device_registry(): Get device registry entries\n"
"- get_area_registry(): Get room/area information\n"
"- get_history(entity_id, hours): Get historical state changes\n"
"- get_person_data(): Get person tracking information\n"
"- get_statistics(entity_id): Get sensor statistics\n"
"- get_scenes(): Get scene configurations\n"
"- get_dashboards(): Get list of all dashboards\n"
"- get_dashboard_config(dashboard_url): Get configuration of a specific dashboard\n"
"- set_entity_state(entity_id, state, attributes?): Set state of an entity (e.g., turn on/off lights, open/close covers)\n"
"- call_service(domain, service, target?, service_data?): Call any Home Assistant service directly\n"
"- create_automation(automation): Create a new automation with the provided configuration\n"
"- create_dashboard(dashboard_config): Create a new dashboard with the provided configuration\n"
"- update_dashboard(dashboard_url, dashboard_config): Update an existing dashboard configuration\n\n"
"IMPORTANT DEVICE_CLASS GUIDANCE:\n"
"- Many sensors have a 'device_class' attribute (temperature, humidity, motion, etc.)\n"
"- Use get_climate_related_entities() for climate dashboards (includes climate.* entities and temperature/humidity sensors)\n"
"- Use get_entities_by_device_class(device_class) to filter by device_class (e.g., 'temperature', 'humidity', 'motion')\n"
"- For climate dashboards, use history-graph and gauge cards for temperature/humidity sensors\n\n"
"DASHBOARD CREATION:\n"
"When a user asks to create a dashboard:\n"
"1. Gather entities using get_climate_related_entities() or other get_* commands\n"
"2. Respond with JSON using request_type: 'dashboard_suggestion' (NEVER use 'final_response'!)\n"
"3. Use Lovelace JSON format (NOT YAML!)\n"
"4. Example response structure:\n"
'{"request_type": "dashboard_suggestion", "message": "Dashboard created", "dashboard": {"title": "...", "views": [...]}}\n'
"5. Do NOT include YAML, markdown, or code blocks - only pure JSON\n\n"
"IMPORTANT AREA/FLOOR GUIDANCE:\n"
"- When users ask for entities from a specific floor, use get_area_registry() first\n"
"- Areas have both 'area_id' and 'floor_id' - these are different concepts\n"
"- Filter areas by their floor_id to find all areas on a specific floor\n"
"- Use get_entities() with area_ids parameter to get entities from multiple areas efficiently\n"
"- Example: get_entities(area_ids=['area1', 'area2', 'area3']) for multiple areas at once\n"
"- This is more efficient than calling get_entities_by_area() multiple times\n\n"
"AUTOMATION CREATION:\n"
"When creating automations, request entities first to know the entity IDs.\n"
"For days, use: ['fri', 'mon', 'sat', 'sun', 'thu', 'tue', 'wed']\n\n"
"RESPONSE FORMATS - You must ALWAYS respond with valid JSON:\n\n"
"For automations:\n"
"{\n"
' "request_type": "automation_suggestion",\n'
' "message": "I\'ve created an automation that might help you. Would you like me to create it?",\n'
' "automation": {\n'
' "alias": "Name of the automation",\n'
' "description": "Description of what the automation does",\n'
' "trigger": [...], // Array of trigger conditions\n'
' "condition": [...], // Optional array of conditions\n'
' "action": [...] // Array of actions to perform\n'
" }\n"
"}\n\n"
"For dashboards (WHEN USER ASKS TO CREATE A DASHBOARD):\n"
"{\n"
' "request_type": "dashboard_suggestion",\n'
' "message": "Description of the dashboard you created",\n'
' "dashboard": {\n'
' "title": "Dashboard Title",\n'
' "url_path": "url-path",\n'
' "icon": "mdi:icon-name",\n'
' "show_in_sidebar": true,\n'
' "views": [{\n'
' "title": "View Title",\n'
' "cards": [...]\n'
" }]\n"
" }\n"
"}\n\n"
"For data requests, use this exact JSON format:\n"
"{\n"
' "request_type": "data_request",\n'
' "request": "command_name",\n'
' "parameters": {...}\n'
"}\n"
'For get_entities with multiple areas: {"request_type": "get_entities", "parameters": {"area_ids": ["area1", "area2"]}}\n'
'For get_entities with single area: {"request_type": "get_entities", "parameters": {"area_id": "single_area"}}\n\n'
"For service calls, use this exact JSON format:\n"
"{\n"
' "request_type": "call_service",\n'
' "domain": "light",\n'
' "service": "turn_on",\n'
' "target": {"entity_id": ["entity1", "entity2"]},\n'
' "service_data": {"brightness": 255}\n'
"}\n\n"
"For answering questions (NOT creating dashboards/automations):\n"
"{\n"
' "request_type": "final_response",\n'
' "response": "your answer to the user"\n'
"}\n\n"
"IMPORTANT: Use 'dashboard_suggestion' when creating dashboards, NOT 'final_response'!\n\n"
"CRITICAL FORMATTING RULES:\n"
"- You must ALWAYS respond with ONLY a valid JSON object\n"
"- DO NOT include any text before the JSON\n"
"- DO NOT include any text after the JSON\n"
"- DO NOT include explanations or descriptions outside the JSON\n"
"- Your entire response must be parseable as JSON\n"
"- Use the 'message' field inside the JSON for user-facing text\n"
"- NEVER mix regular text with JSON in your response\n\n"
"WRONG: 'I'll create this for you. {\"request_type\": ...}'\n"
'CORRECT: \'{"request_type": "dashboard_suggestion", "message": "I\'ll create this for you.", ...}\''
),
}
def __init__(self, hass: HomeAssistant, config: Dict[str, Any]):
"""Initialize the agent with provider selection."""
self.hass = hass
self.config = config
self.conversation_history: List[Dict[str, Any]] = []
self._cache: Dict[str, Any] = {}
self.ai_client: BaseAIClient
self._cache_timeout = 300 # 5 minutes
self._max_retries = 10
self._retry_delay = 1 # seconds
self._rate_limit = 60 # requests per minute
self._last_request_time = 0
self._request_count = 0
self._request_window_start = time.time()
provider = config.get("ai_provider", "openai")
models_config = config.get("models", {})
_LOGGER.debug("Initializing AiAgentHaAgent with provider: %s", provider)
_LOGGER.debug("Models config loaded: %s", models_config)
# Set the appropriate system prompt based on provider
if provider in ("local_ollama", "openai_compatible"):
self.system_prompt = self.SYSTEM_PROMPT_LOCAL
_LOGGER.debug(
"Using local-optimized system prompt for provider: %s", provider
)
else:
self.system_prompt = self.SYSTEM_PROMPT
_LOGGER.debug("Using standard system prompt")
# Initialize the appropriate AI client with model selection
if provider == "openai":
model = models_config.get("openai", "gpt-3.5-turbo")
base_url = config.get(CONF_OPENAI_BASE_URL) or ""
self.ai_client = OpenAIClient(
config.get("openai_token"), model, base_url or None
)
elif provider == "gemini":
model = models_config.get("gemini", "gemini-2.5-flash")
self.ai_client = GeminiClient(config.get("gemini_token"), model)
elif provider == "openrouter":
model = models_config.get("openrouter", "openai/gpt-4o")
self.ai_client = OpenRouterClient(config.get("openrouter_token"), model)
elif provider == "anthropic":
model = models_config.get("anthropic", "claude-sonnet-4-5-20250929")
self.ai_client = AnthropicClient(config.get("anthropic_token"), model)
elif provider == "alter":
model = models_config.get("alter", "")
self.ai_client = AlterClient(config.get("alter_token"), model)
elif provider == "zai":
model = models_config.get("zai", "glm-4.7")
endpoint_type = config.get("zai_endpoint", "general")
self.ai_client = ZaiClient(config.get("zai_token"), model, endpoint_type)
elif provider == "local_ollama":
# Support both new local_ollama_url and legacy local_url
url = config.get("local_ollama_url") or config.get("local_url")
model = models_config.get("local_ollama") or models_config.get("local", "")
if not url:
_LOGGER.error("Missing local_ollama_url for local_ollama provider")
raise Exception(
"Missing local_ollama_url configuration for local_ollama provider"
)
self.ai_client = LocalOllamaClient(url, model)
elif provider == "openai_compatible":
url = config.get("openai_compatible_url")
model = models_config.get("openai_compatible", "")
api_key = config.get("openai_compatible_api_key", "") or ""
if not url:
_LOGGER.error(
"Missing openai_compatible_url for openai_compatible provider"
)
raise Exception(
"Missing openai_compatible_url configuration for openai_compatible provider"
)
self.ai_client = OpenaiCompatibleClient(url, model, api_key or None)
else: # default to llama if somehow specified
model = models_config.get("llama", "Llama-4-Maverick-17B-128E-Instruct-FP8")
self.ai_client = LlamaClient(config.get("llama_token"), model)
_LOGGER.debug(
"AiAgentHaAgent initialized successfully with provider: %s, model: %s",
provider,
model,
)
def _validate_api_key(self) -> bool:
"""Validate the API key format."""
provider = self.config.get("ai_provider", "openai")
if provider == "openai":
token = self.config.get("openai_token")
elif provider == "gemini":
token = self.config.get("gemini_token")
elif provider == "openrouter":
token = self.config.get("openrouter_token")
elif provider == "anthropic":
token = self.config.get("anthropic_token")
elif provider == "alter":
token = self.config.get("alter_token")
elif provider == "zai":
token = self.config.get("zai_token")
elif provider == "local_ollama":
# For local_ollama, the “token” is actually the URL; support legacy local_url
token = self.config.get("local_ollama_url") or self.config.get("local_url")
elif provider == "openai_compatible":
# For openai_compatible, validate the URL is present
token = self.config.get("openai_compatible_url")
else:
token = self.config.get("llama_token")
if not token or not isinstance(token, str):
return False
# For local_ollama and openai_compatible, validate URL format
if provider in ("local_ollama", "openai_compatible"):
return bool(token.startswith(("http://", "https://")))
# Add more specific validation based on your API key format
return len(token) >= 32
def _check_rate_limit(self) -> bool:
"""Check if we're within rate limits."""
current_time = time.time()
if current_time - self._request_window_start >= 60:
self._request_count = 0
self._request_window_start = current_time
if self._request_count >= self._rate_limit:
return False
self._request_count += 1
return True
def _get_cached_data(self, key: str) -> Optional[Any]:
"""Get data from cache if it's still valid."""
if key in self._cache:
timestamp, data = self._cache[key]
if time.time() - timestamp < self._cache_timeout:
return data
del self._cache[key]
return None
def _set_cached_data(self, key: str, data: Any) -> None:
"""Store data in cache with timestamp."""
self._cache[key] = (time.time(), data)
# Maximum size (in characters) of a single data message added to the
# conversation. Large installs can return megabytes of entity data, which
# blows past the model's context window (~200k tokens for Claude) and makes
# every request fail with a deterministic 400 (issue #80). 50k chars is
# roughly 13k tokens, which also keeps requests within entry-tier
# per-minute token rate limits (e.g. Anthropic tier 1: 30k input
# tokens/min) even with a data message persisting in the history window.
MAX_DATA_MESSAGE_CHARS = 50_000
# Maximum combined size (in characters) of the message window sent to the
# provider per request. Several capped data messages can still stack past
# context and per-minute token budgets (issue #80); the most recent
# messages are kept. ~100k chars is roughly 25k tokens.
MAX_WINDOW_CHARS = 100_000
def _format_data_message(self, data: Any) -> str:
"""Serialize fetched HA data for the conversation, capping its size.
If the payload is too large, list items are dropped from the end and a
truncation notice is included so the model knows to request more
specific data instead of the full dump.
"""
message = json.dumps({"data": data}, default=str)
if len(message) <= self.MAX_DATA_MESSAGE_CHARS:
return message
note = (
"Data truncated to fit the model context window. "
"Request more specific data (e.g. a specific entity, domain or area) "
"to see the rest."
)
if isinstance(data, list) and data:
truncated = data
while len(message) > self.MAX_DATA_MESSAGE_CHARS and len(truncated) > 1:
# Scale down proportionally to the overshoot, then re-check
# (item sizes vary, so loop until it actually fits).
keep = max(
1,
len(truncated) * self.MAX_DATA_MESSAGE_CHARS // len(message),
)
truncated = truncated[:keep]
message = json.dumps(
{
"data": truncated,
"truncated": True,
"total_items": len(data),
"items_shown": len(truncated),
"note": note,
},
default=str,
)
if len(message) <= self.MAX_DATA_MESSAGE_CHARS:
_LOGGER.warning(
"Data response truncated from %d to %d items to fit context window",
len(data),
len(truncated),
)
return message
# A single item alone exceeds the cap - fall through to the
# hard-truncated preview below so the cap always holds.
# Oversized non-list payloads (or a single oversized list item): keep
# a prefix of the serialized form as a preview.
_LOGGER.warning(
"Oversized data response hard-truncated from %d chars", len(message)
)
preview = message[: self.MAX_DATA_MESSAGE_CHARS]
result = json.dumps({"data_preview": preview, "truncated": True, "note": note})
# JSON-escaping the preview can push the result back over the cap;
# shrink until the final message actually fits.
while len(result) > self.MAX_DATA_MESSAGE_CHARS and preview:
preview = preview[: int(len(preview) * 0.9)]
result = json.dumps(
{"data_preview": preview, "truncated": True, "note": note}
)
return result
def _sanitize_automation_config(self, config: Dict[str, Any]) -> Dict[str, Any]:
"""Sanitize automation configuration to prevent injection attacks."""
sanitized: Dict[str, Any] = {}
for key, value in config.items():
if key in ["alias", "description"]:
# Sanitize strings
sanitized[key] = str(value).strip()[:100] # Limit length
elif key in ["trigger", "condition", "action"]:
# Home Assistant accepts either a single mapping or a list for
# trigger/condition/action. Normalize a single mapping to a
# one-element list instead of silently dropping it (which would
# later raise KeyError and reject a perfectly valid automation).
if isinstance(value, list):
sanitized[key] = value
elif isinstance(value, dict):
sanitized[key] = [value]
elif key == "mode":
# Validate mode
if value in ["single", "restart", "queued", "parallel"]:
sanitized[key] = value
return sanitized
@staticmethod
def _read_automations_file(path: str) -> List[Dict[str, Any]]:
"""Read and parse automations.yaml into a list of automations.
Runs inside an executor thread. Raises FileNotFoundError when the file
does not exist (the caller treats that as "no automations yet").
"""
with open(path, "r", encoding="utf-8") as handle:
data = yaml.safe_load(handle)
if data is None:
return []
if not isinstance(data, list):
# automations.yaml is always a top-level list. If it is anything
# else, refuse to touch it rather than risk clobbering content we
# don't understand.
raise ValueError("automations.yaml does not contain a list of automations")
return data
@staticmethod
def _write_automations_file(path: str, automations: List[Dict[str, Any]]) -> None:
"""Safely persist automations to automations.yaml.
Runs inside an executor thread. Compared to a naive ``yaml.dump`` to an
open file handle, this:
* keeps accented/unicode text readable instead of mangling it into
``\\uXXXX`` escapes (``allow_unicode=True``),
* preserves key order so automations stay diff-friendly
(``sort_keys=False``),
* never line-wraps long Jinja templates (``width``),
* validates that the serialized YAML round-trips back to the same
data before touching disk,
* backs up the previous file to ``<path>.bak`` so the user can roll
back,
* writes atomically (temp file + ``os.replace``) so a crash or a bad
write can never leave a half-written / corrupted file in place.
"""
# Use safe_dump (the SafeDumper) so the writer mirrors the safe_load
# reader: only plain YAML types are ever emitted, never opaque
# ``!!python/object`` tags.
content = yaml.safe_dump(
automations,
default_flow_style=False,
allow_unicode=True,
sort_keys=False,
width=4096,
)
# Guard against ever writing YAML that does not decode back to the
# exact same data structure.
if yaml.safe_load(content) != automations:
raise ValueError(
"Refusing to write automations.yaml: serialized YAML did not "
"round-trip cleanly"
)
path_exists = os.path.exists(path)
# Back up the existing file before overwriting it.
if path_exists:
shutil.copy2(path, f"{path}.bak")
# Atomic write: write to a temp file in the same directory, fsync, then
# atomically replace the target.
directory = os.path.dirname(path) or "."
fd, tmp_path = tempfile.mkstemp(
prefix=".automations.", suffix=".yaml.tmp", dir=directory
)
try:
with os.fdopen(fd, "w", encoding="utf-8") as handle:
handle.write(content)
handle.flush()
os.fsync(handle.fileno())
# Preserve the original file's permission bits. mkstemp creates the
# temp file as 0600, so without this an atomic replace would strip
# any group/world bits the user or an external editor relied on.
if path_exists:
shutil.copymode(path, tmp_path)
os.replace(tmp_path, path)
except BaseException:
# Never leave a stray temp file behind on failure.
if os.path.exists(tmp_path):
os.remove(tmp_path)
raise
async def get_entity_state(self, entity_id: str) -> Dict[str, Any]:
"""Get the state of a specific entity."""
try:
_LOGGER.debug("Requesting entity state for: %s", entity_id)
state = self.hass.states.get(entity_id)
if not state:
_LOGGER.warning("Entity not found: %s", entity_id)
return {"error": f"Entity {entity_id} not found"}
# Get area information from entity/device registry
# Wrapped in try-except to handle cases where registries aren't available (e.g., in tests)
area_id = None
area_name = None
try:
from homeassistant.helpers import area_registry as ar
from homeassistant.helpers import device_registry as dr
from homeassistant.helpers import entity_registry as er
entity_registry = er.async_get(self.hass)
device_registry = dr.async_get(self.hass)
area_registry = ar.async_get(self.hass)
if entity_registry and hasattr(entity_registry, "async_get"):
# Try to find the entity in the registry
entity_entry = entity_registry.async_get(entity_id)
if entity_entry:
_LOGGER.debug("Entity %s found in registry", entity_id)
# Check if entity has a direct area assignment
if hasattr(entity_entry, "area_id") and entity_entry.area_id:
area_id = entity_entry.area_id
_LOGGER.debug(
"Entity %s has direct area assignment: %s",
entity_id,
area_id,
)
# Otherwise check if the entity's device has an area
elif (
hasattr(entity_entry, "device_id")
and entity_entry.device_id
and device_registry
and hasattr(device_registry, "async_get")
):
_LOGGER.debug(
"Entity %s has device_id: %s, checking device area",
entity_id,
entity_entry.device_id,
)
device_entry = device_registry.async_get(
entity_entry.device_id
)
if device_entry:
if (
hasattr(device_entry, "area_id")
and device_entry.area_id
):
area_id = device_entry.area_id
_LOGGER.debug(
"Device %s has area: %s",
entity_entry.device_id,
area_id,
)
else:
_LOGGER.debug(
"Device %s has no area assigned",
entity_entry.device_id,
)
else:
_LOGGER.debug(
"Device %s not found in registry",
entity_entry.device_id,
)
else:
_LOGGER.debug(
"Entity %s has no area_id and no device_id", entity_id
)
else:
_LOGGER.debug(
"Entity %s not found in entity registry", entity_id
)
else:
_LOGGER.debug("Entity registry not available for %s", entity_id)
# Get area name from area_id
if (
area_id
and area_registry
and hasattr(area_registry, "async_get_area")
):
area_entry = area_registry.async_get_area(area_id)
if area_entry and hasattr(area_entry, "name"):
area_name = area_entry.name
_LOGGER.debug(
"Resolved area_id %s to area_name: %s", area_id, area_name
)
else:
_LOGGER.debug("Could not resolve area_id %s to name", area_id)
elif area_id:
_LOGGER.debug(
"Have area_id %s but area_registry not available", area_id
)
except Exception as e:
# Registries not available (likely in test environment) - skip area information
_LOGGER.warning(
"Exception retrieving area information for %s: %s",
entity_id,
str(e),
)
result = {
"entity_id": state.entity_id,
"state": state.state,
"last_changed": (
state.last_changed.isoformat() if state.last_changed else None
),
"friendly_name": state.attributes.get("friendly_name"),
"area_id": area_id,
"area_name": area_name,
"attributes": {
k: (v.isoformat() if hasattr(v, "isoformat") else v)
for k, v in state.attributes.items()
},
}
_LOGGER.debug(
"Retrieved entity state for %s: area_id=%s, area_name=%s",
entity_id,
area_id,
area_name,
)
return result
except Exception as e:
_LOGGER.exception("Error getting entity state: %s", str(e))
return {"error": f"Error getting entity state: {str(e)}"}
async def get_entities_by_domain(self, domain: str) -> List[Dict[str, Any]]:
"""Get all entities for a specific domain."""
try:
_LOGGER.debug("Requesting all entities for domain: %s", domain)
states = [
state
for state in self.hass.states.async_all()
if state.entity_id.startswith(f"{domain}.")
]
_LOGGER.debug("Found %d entities in domain %s", len(states), domain)
return [await self.get_entity_state(state.entity_id) for state in states]
except Exception as e:
_LOGGER.exception("Error getting entities by domain: %s", str(e))
return [{"error": f"Error getting entities for domain {domain}: {str(e)}"}]
async def get_entities_by_device_class(
self, device_class: str, domain: str = None
) -> List[Dict[str, Any]]:
"""Get all entities with a specific device_class.
Args:
device_class: The device class to filter by (e.g., 'temperature', 'humidity', 'motion')
domain: Optional domain to restrict search (e.g., 'sensor', 'binary_sensor')
Returns:
List of entity state dictionaries that match the device_class
"""
try:
_LOGGER.debug(
"Requesting all entities with device_class: %s (domain: %s)",
device_class,
domain or "all",
)
matching_entities = []
for state in self.hass.states.async_all():
# Filter by domain if specified
if domain and not state.entity_id.startswith(f"{domain}."):
continue
# Check if this entity has the matching device_class
entity_device_class = state.attributes.get("device_class")
if entity_device_class == device_class:
matching_entities.append(state.entity_id)
_LOGGER.debug(
"Found %d entities with device_class %s",
len(matching_entities),
device_class,
)
# Get full state information for each matching entity
return [
await self.get_entity_state(entity_id)
for entity_id in matching_entities
]
except Exception as e:
_LOGGER.exception("Error getting entities by device_class: %s", str(e))
return [
{
"error": f"Error getting entities with device_class {device_class}: {str(e)}"
}
]
async def get_climate_related_entities(self) -> List[Dict[str, Any]]:
"""Get all climate-related entities including climate domain and temperature/humidity sensors.
Returns:
List of entity state dictionaries for:
- All climate.* entities (thermostats, HVAC systems)
- All sensor.* entities with device_class: temperature
- All sensor.* entities with device_class: humidity
"""
try:
_LOGGER.debug("Requesting all climate-related entities")
climate_entities = []
# Get all climate domain entities (thermostats, HVAC)
climate_domain = await self.get_entities_by_domain("climate")
climate_entities.extend(climate_domain)
# Get temperature sensors
temp_sensors = await self.get_entities_by_device_class(
"temperature", "sensor"
)
climate_entities.extend(temp_sensors)
# Get humidity sensors
humidity_sensors = await self.get_entities_by_device_class(
"humidity", "sensor"
)
climate_entities.extend(humidity_sensors)
# Deduplicate by entity_id (edge case: if an entity appears in multiple categories)
seen_entity_ids = set()
unique_entities = []
for entity in climate_entities:
entity_id = entity.get("entity_id")
if entity_id and entity_id not in seen_entity_ids:
seen_entity_ids.add(entity_id)
unique_entities.append(entity)
_LOGGER.debug(
"Found %d total climate-related entities (deduplicated from %d)",
len(unique_entities),
len(climate_entities),
)
return unique_entities
except Exception as e:
_LOGGER.exception("Error getting climate-related entities: %s", str(e))
return [{"error": f"Error getting climate-related entities: {str(e)}"}]
async def get_entities_by_area(self, area_id: str) -> List[Dict[str, Any]]:
"""Get all entities for a specific area."""
try:
_LOGGER.debug("Requesting all entities for area: %s", area_id)
# Get entity registry to find entities assigned to the area
from homeassistant.helpers import device_registry as dr
from homeassistant.helpers import entity_registry as er
entity_registry = er.async_get(self.hass)
device_registry = dr.async_get(self.hass)
entities_in_area = []
# Find entities assigned to the area (directly or through their device)
for entity in entity_registry.entities.values():
# Check if entity is directly assigned to the area
if entity.area_id == area_id:
entities_in_area.append(entity.entity_id)
# Check if entity's device is assigned to the area
elif entity.device_id:
device = device_registry.devices.get(entity.device_id)
if device and device.area_id == area_id:
entities_in_area.append(entity.entity_id)
_LOGGER.debug(
"Found %d entities in area %s", len(entities_in_area), area_id
)
# Get state information for each entity
result = []
for entity_id in entities_in_area:
state_info = await self.get_entity_state(entity_id)
if not state_info.get("error"): # Only include entities that exist
result.append(state_info)
return result
except Exception as e:
_LOGGER.exception("Error getting entities by area: %s", str(e))
return [{"error": f"Error getting entities for area {area_id}: {str(e)}"}]
async def get_entities(self, area_id=None, area_ids=None) -> List[Dict[str, Any]]:
"""Get entities by area(s) - flexible method that supports single area or multiple areas."""
try:
# Handle different parameter formats
areas_to_process = []
if area_ids:
# Multiple areas provided
if isinstance(area_ids, list):
areas_to_process = area_ids
else:
areas_to_process = [area_ids]
elif area_id:
# Single area provided
if isinstance(area_id, list):
areas_to_process = area_id
else:
areas_to_process = [area_id]
else:
return [{"error": "No area_id or area_ids provided"}]
_LOGGER.debug("Requesting entities for areas: %s", areas_to_process)
all_entities = []
for area in areas_to_process:
entities_in_area = await self.get_entities_by_area(area)
all_entities.extend(entities_in_area)
# Remove duplicates based on entity_id
seen_entities = set()
unique_entities = []
for entity in all_entities:
if isinstance(entity, dict) and "entity_id" in entity:
if entity["entity_id"] not in seen_entities:
seen_entities.add(entity["entity_id"])
unique_entities.append(entity)
else:
unique_entities.append(entity) # Keep error messages
_LOGGER.debug(
"Found %d unique entities across %d areas",
len(unique_entities),
len(areas_to_process),
)
return unique_entities
except Exception as e:
_LOGGER.exception("Error getting entities: %s", str(e))
return [{"error": f"Error getting entities: {str(e)}"}]
async def get_calendar_events(
self, entity_id: Optional[str] = None
) -> List[Dict[str, Any]]:
"""Get calendar events, optionally filtered by entity_id."""
try:
if entity_id:
_LOGGER.debug(
"Requesting calendar events for specific entity: %s", entity_id
)
return [await self.get_entity_state(entity_id)]
_LOGGER.debug("Requesting all calendar events")
return await self.get_entities_by_domain("calendar")
except Exception as e:
_LOGGER.exception("Error getting calendar events: %s", str(e))
return [{"error": f"Error getting calendar events: {str(e)}"}]
async def get_automations(self) -> List[Dict[str, Any]]:
"""Get all automations."""
try:
_LOGGER.debug("Requesting all automations")
return await self.get_entities_by_domain("automation")
except Exception as e:
_LOGGER.exception("Error getting automations: %s", str(e))
return [{"error": f"Error getting automations: {str(e)}"}]
async def get_entity_registry(self) -> List[Dict]:
"""Get entity registry entries with device_class and other metadata.
Area information is resolved from the entity or its device.
"""
_LOGGER.debug("Requesting all entity registry entries")
try:
from homeassistant.helpers import area_registry as ar
from homeassistant.helpers import device_registry as dr
from homeassistant.helpers import entity_registry as er
entity_registry = er.async_get(self.hass)
if not entity_registry:
return []
device_registry = dr.async_get(self.hass)
area_registry = ar.async_get(self.hass)
result = []
for entry in entity_registry.entities.values():
# Get the current state to access device_class and other attributes
state = self.hass.states.get(entry.entity_id)
device_class = state.attributes.get("device_class") if state else None
state_class = state.attributes.get("state_class") if state else None
unit_of_measurement = (
state.attributes.get("unit_of_measurement") if state else None
)
# Resolve area_id and area_name
# First check entity's direct area assignment
area_id = entry.area_id
area_name = None
# If entity doesn't have area, check device's area
if not area_id and entry.device_id and device_registry:
device_entry = device_registry.async_get(entry.device_id)
if device_entry and hasattr(device_entry, "area_id"):
area_id = device_entry.area_id
# Resolve area_name from area_id
if area_id and area_registry:
area_entry = area_registry.async_get_area(area_id)
if area_entry and hasattr(area_entry, "name"):
area_name = area_entry.name
result.append(
{
"entity_id": entry.entity_id,
"device_id": entry.device_id,
"platform": entry.platform,
"disabled": entry.disabled,
"area_id": area_id,
"area_name": area_name,
"original_name": entry.original_name,
"unique_id": entry.unique_id,
"device_class": device_class,
"state_class": state_class,
"unit_of_measurement": unit_of_measurement,
}
)
return result
except Exception as e:
_LOGGER.exception("Error getting entity registry entries: %s", str(e))
return [{"error": f"Error getting entity registry entries: {str(e)}"}]
async def get_device_registry(self) -> List[Dict]:
"""Get device registry entries"""
_LOGGER.debug("Requesting all device registry entries")
try:
from homeassistant.helpers import device_registry as dr
registry = dr.async_get(self.hass)
if not registry:
return []
return [
{
"id": device.id,
"name": device.name,
"model": device.model,
"manufacturer": device.manufacturer,
"sw_version": device.sw_version,
"hw_version": device.hw_version,
"connections": (
list(device.connections) if device.connections else []
),
"identifiers": (
list(device.identifiers) if device.identifiers else []
),
"area_id": device.area_id,
"disabled": device.disabled_by is not None,
"entry_type": (
device.entry_type.value if device.entry_type else None
),
"name_by_user": device.name_by_user,
}
for device in registry.devices.values()
]
except Exception as e:
_LOGGER.exception("Error getting device registry entries: %s", str(e))
return [{"error": f"Error getting device registry entries: {str(e)}"}]
async def get_history(self, entity_id: str, hours: int = 24) -> List[Dict]:
"""Get historical state changes for an entity"""
_LOGGER.debug("Requesting historical state changes for entity: %s", entity_id)
try:
from homeassistant.components.recorder.history import get_significant_states
now = dt_util.utcnow()
start = now - timedelta(hours=hours)
# Get history using the recorder history module
history_data = await self.hass.async_add_executor_job(
get_significant_states,
self.hass,
start,
now,
[entity_id],
)
# Convert to serializable format
result = []
for entity_id_key, states in history_data.items():
for state in states:
# Skip if it's a dict (mypy type narrowing)
if isinstance(state, dict):
continue
result.append(
{
"entity_id": state.entity_id,
"state": state.state,
"last_changed": state.last_changed.isoformat(),
"last_updated": state.last_updated.isoformat(),
"attributes": dict(state.attributes),
}
)
return result
except Exception as e:
_LOGGER.exception("Error getting history: %s", str(e))
return [{"error": f"Error getting history: {str(e)}"}]
async def get_area_registry(self) -> Dict[str, Any]:
"""Get area registry information"""
_LOGGER.debug("Get area registry information")
try:
from homeassistant.helpers import area_registry as ar
registry = ar.async_get(self.hass)
if not registry:
return {}
result = {}
for area in registry.areas.values():
result[area.id] = {
"name": area.name,
"normalized_name": area.normalized_name,
"picture": area.picture,
"icon": area.icon,
"floor_id": area.floor_id,
"labels": list(area.labels) if area.labels else [],
}
return result
except Exception as e:
_LOGGER.exception("Error getting area registry: %s", str(e))
return {"error": f"Error getting area registry: {str(e)}"}
async def get_person_data(self) -> List[Dict]:
"""Get person tracking information"""
_LOGGER.debug("Requesting person tracking information")
try:
result = []
for state in self.hass.states.async_all("person"):
result.append(
{
"entity_id": state.entity_id,
"name": state.attributes.get("friendly_name", state.entity_id),
"state": state.state,
"latitude": state.attributes.get("latitude"),
"longitude": state.attributes.get("longitude"),
"source": state.attributes.get("source"),
"gps_accuracy": state.attributes.get("gps_accuracy"),
"last_changed": (
state.last_changed.isoformat()
if state.last_changed
else None
),
}
)
return result
except Exception as e:
_LOGGER.exception("Error getting person tracking information: %s", str(e))
return [{"error": f"Error getting person tracking information: {str(e)}"}]
async def get_statistics(self, entity_id: str) -> Dict:
"""Get statistics for an entity"""
_LOGGER.debug("Requesting statistics for entity: %s", entity_id)
try:
from homeassistant.components import recorder
# Check if recorder is available
if not self.hass.data.get(recorder.DATA_INSTANCE):
return {"error": "Recorder component is not available"}
# from homeassistant.components.recorder.statistics import get_latest_short_term_statistics
import homeassistant.components.recorder.statistics as stats_module
# Get latest statistics
stats = await self.hass.async_add_executor_job(
# get_latest_short_term_statistics,
stats_module.get_last_short_term_statistics,
self.hass,
1,
entity_id,
True,
set(),
)
if entity_id in stats:
stat_data = stats[entity_id][0] if stats[entity_id] else {}
return {
"entity_id": entity_id,
"start": stat_data.get("start"),
"mean": stat_data.get("mean"),
"min": stat_data.get("min"),
"max": stat_data.get("max"),
"last_reset": stat_data.get("last_reset"),
"state": stat_data.get("state"),
"sum": stat_data.get("sum"),
}
else:
return {"error": f"No statistics available for entity {entity_id}"}
except Exception as e:
_LOGGER.exception("Error getting statistics: %s", str(e))
return {"error": f"Error getting statistics: {str(e)}"}
async def get_scenes(self) -> List[Dict]:
"""Get scene configurations"""
_LOGGER.debug("Requesting scene configurations")
try:
result = []
for state in self.hass.states.async_all("scene"):
result.append(
{
"entity_id": state.entity_id,
"name": state.attributes.get("friendly_name", state.entity_id),
"last_activated": state.attributes.get("last_activated"),
"icon": state.attributes.get("icon"),
"last_changed": (
state.last_changed.isoformat()
if state.last_changed
else None
),
}
)
return result
except Exception as e:
_LOGGER.exception("Error getting scene configurations: %s", str(e))
return [{"error": f"Error getting scene configurations: {str(e)}"}]
async def get_weather_data(self) -> Dict[str, Any]:
"""Get weather data from any available weather entity in the system."""
try:
# Find all weather entities
weather_entities = [
state
for state in self.hass.states.async_all()
if state.domain == "weather"
]
if not weather_entities:
return {
"error": "No weather entities found in the system. Please add a weather integration."
}
# Use the first available weather entity
state = weather_entities[0]
_LOGGER.debug("Using weather entity: %s", state.entity_id)
# Get all available attributes
all_attributes = state.attributes
_LOGGER.debug(
"Available weather attributes: %s", json.dumps(all_attributes)
)
# Get forecast data
forecast = all_attributes.get("forecast", [])
# Process forecast data
processed_forecast = []
for day in forecast:
forecast_entry = {
"datetime": day.get("datetime"),
"temperature": day.get("temperature"),
"condition": day.get("condition"),
"precipitation": day.get("precipitation"),
"precipitation_probability": day.get("precipitation_probability"),
"humidity": day.get("humidity"),
"wind_speed": day.get("wind_speed"),
"wind_bearing": day.get("wind_bearing"),
}
# Only add entries that have at least some data
if any(v is not None for v in forecast_entry.values()):
processed_forecast.append(forecast_entry)
# Get current weather data
current = {
"entity_id": state.entity_id,
"temperature": all_attributes.get("temperature"),
"humidity": all_attributes.get("humidity"),
"pressure": all_attributes.get("pressure"),
"wind_speed": all_attributes.get("wind_speed"),
"wind_bearing": all_attributes.get("wind_bearing"),
"condition": state.state,
"forecast_available": len(processed_forecast) > 0,
}
# Log the processed data for debugging
_LOGGER.debug(
"Processed weather data: %s",
json.dumps(
{"current": current, "forecast_count": len(processed_forecast)}
),
)
return {"current": current, "forecast": processed_forecast}
except Exception as e:
_LOGGER.exception("Error getting weather data: %s", str(e))
return {"error": f"Error getting weather data: {str(e)}"}
async def create_automation(
self, automation_config: Dict[str, Any]
) -> Dict[str, Any]:
"""Create a new automation with validation and sanitization."""
try:
_LOGGER.debug(
"Creating automation with config: %s", json.dumps(automation_config)
)
# Validate required fields
if not all(
key in automation_config for key in ["alias", "trigger", "action"]
):
return {"error": "Missing required fields in automation configuration"}
# Sanitize configuration
sanitized_config = self._sanitize_automation_config(automation_config)
# Make sure the core building blocks survived sanitization. A
# malformed trigger/action (e.g. not a list or mapping) is dropped
# by the sanitizer, so validate here and fail with a clear message
# instead of raising KeyError further down.
if not sanitized_config.get("alias"):
return {"error": "Automation must include a non-empty alias"}
if not sanitized_config.get("trigger"):
return {"error": "Automation must include at least one trigger"}
if not sanitized_config.get("action"):
return {"error": "Automation must include at least one action"}
# Generate a unique ID for the automation
automation_id = f"ai_agent_auto_{int(time.time() * 1000)}"
# Create the automation entry
automation_entry = {
"id": automation_id,
"alias": sanitized_config["alias"],
"description": sanitized_config.get("description", ""),
"trigger": sanitized_config["trigger"],
"condition": sanitized_config.get("condition", []),
"action": sanitized_config["action"],
"mode": sanitized_config.get("mode", "single"),
}
# Read current automations.yaml using async executor
automations_path = self.hass.config.path("automations.yaml")
try:
current_automations = await self.hass.async_add_executor_job(
self._read_automations_file, automations_path
)
except FileNotFoundError:
current_automations = []
# Check for duplicate automation names
if any(
auto.get("alias") == automation_entry["alias"]
for auto in current_automations
):
return {
"error": f"An automation with the name '{automation_entry['alias']}' already exists"
}
# Append new automation
current_automations.append(automation_entry)
# Write back to file safely: backs up the previous file, validates
# the YAML round-trips, preserves unicode/key order, and replaces
# the file atomically so a bad write can never corrupt it.
await self.hass.async_add_executor_job(
self._write_automations_file,
automations_path,
current_automations,
)
# Reload automations
await self.hass.services.async_call("automation", "reload")
# Clear automation-related caches
self._cache.clear()
return {
"success": True,
"message": f"Automation '{automation_entry['alias']}' created successfully",
}
except Exception as e:
_LOGGER.exception("Error creating automation: %s", str(e))
return {"error": f"Error creating automation: {str(e)}"}
async def get_dashboards(self) -> List[Dict[str, Any]]:
"""Get list of all dashboards."""
try:
_LOGGER.debug("Requesting all dashboards")
# Get dashboards via WebSocket API
ws_api = self.hass.data.get("websocket_api")
if not ws_api:
return [{"error": "WebSocket API not available"}]
# Use the lovelace service to get dashboards
try:
from homeassistant.components.lovelace import DOMAIN as LOVELACE_DOMAIN
# Get lovelace data using property access (required for HA 2026.2+)
# lovelace_data is a LovelaceData dataclass with a 'dashboards' attribute
lovelace_data = self.hass.data.get(LOVELACE_DOMAIN)
if lovelace_data is None:
return [{"error": "Lovelace not available"}]
# Safety check for dashboards attribute (backward compatibility)
if not hasattr(lovelace_data, "dashboards"):
return [{"error": "Lovelace dashboards not available"}]
# Use property access instead of dictionary access
dashboards = lovelace_data.dashboards
# Get YAML dashboard configs for metadata (title, icon, etc.)
# yaml_dashboards contains the configuration with metadata
yaml_configs = getattr(lovelace_data, "yaml_dashboards", {}) or {}
dashboard_list = []
# Iterate over all dashboards (None key = default dashboard)
for url_path, dashboard_obj in dashboards.items():
# Try to get metadata from yaml_dashboards first
yaml_config = yaml_configs.get(url_path, {}) or {}
# Get title - check yaml config, then use defaults
title = yaml_config.get("title")
if not title:
title = (
"Overview"
if url_path is None
else (url_path or "Dashboard")
)
# Get icon - check yaml config, then use defaults
icon = yaml_config.get("icon")
if not icon:
icon = "mdi:home" if url_path is None else "mdi:view-dashboard"
# Get sidebar/admin settings from yaml config or defaults
show_in_sidebar = yaml_config.get("show_in_sidebar", True)
require_admin = yaml_config.get("require_admin", False)
dashboard_list.append(
{
"url_path": url_path,
"title": title,
"icon": icon,
"show_in_sidebar": show_in_sidebar,
"require_admin": require_admin,
}
)
_LOGGER.debug("Found %d dashboards", len(dashboard_list))
return dashboard_list
except Exception as e:
_LOGGER.warning("Could not get dashboards via lovelace: %s", str(e))
return [{"error": f"Could not retrieve dashboards: {str(e)}"}]
except Exception as e:
_LOGGER.exception("Error getting dashboards: %s", str(e))
return [{"error": f"Error getting dashboards: {str(e)}"}]
async def get_dashboard_config(
self, dashboard_url: Optional[str] = None
) -> Dict[str, Any]:
"""Get configuration of a specific dashboard."""
try:
_LOGGER.debug(
"Requesting dashboard config for: %s", dashboard_url or "default"
)
# Get dashboard configuration
try:
from homeassistant.components.lovelace import DOMAIN as LOVELACE_DOMAIN
# Get lovelace data using property access (required for HA 2026.2+)
lovelace_data = self.hass.data.get(LOVELACE_DOMAIN)
if lovelace_data is None:
return {"error": "Lovelace not available"}
# Safety check for dashboards attribute (backward compatibility)
if not hasattr(lovelace_data, "dashboards"):
return {"error": "Lovelace dashboards not available"}
# Use property access instead of dictionary access
# The dashboards dict uses None as key for the default dashboard
dashboards = lovelace_data.dashboards
# Get the dashboard (None key = default dashboard)
dashboard_key = None if dashboard_url is None else dashboard_url
if dashboard_key in dashboards:
dashboard = dashboards[dashboard_key]
config = await dashboard.async_get_info()
return dict(config) if config else {"error": "No dashboard config"}
else:
if dashboard_url is None:
return {"error": "Default dashboard not found"}
else:
return {"error": f"Dashboard '{dashboard_url}' not found"}
except Exception as e:
_LOGGER.warning("Could not get dashboard config: %s", str(e))
return {"error": f"Could not retrieve dashboard config: {str(e)}"}
except Exception as e:
_LOGGER.exception("Error getting dashboard config: %s", str(e))
return {"error": f"Error getting dashboard config: {str(e)}"}
async def create_dashboard(
self, dashboard_config: Dict[str, Any]
) -> Dict[str, Any]:
"""Create a new dashboard using Home Assistant's Lovelace WebSocket API."""
try:
_LOGGER.debug(
"Creating dashboard with config: %s",
json.dumps(dashboard_config, default=str),
)
# Validate required fields
if not dashboard_config.get("title"):
return {"error": "Dashboard title is required"}
if not dashboard_config.get("url_path"):
return {"error": "Dashboard URL path is required"}
# Sanitize the URL path
url_path = (
dashboard_config["url_path"].lower().replace(" ", "-").replace("_", "-")
)
# Prepare dashboard configuration for Lovelace
dashboard_data = {
"title": dashboard_config["title"],
"icon": dashboard_config.get("icon", "mdi:view-dashboard"),
"show_in_sidebar": dashboard_config.get("show_in_sidebar", True),
"require_admin": dashboard_config.get("require_admin", False),
"views": dashboard_config.get("views", []),
}
try:
# Create dashboard file directly - this is the most reliable method
import os
import yaml
# Create the dashboard YAML file
lovelace_config_file = self.hass.config.path(
f"ui-lovelace-{url_path}.yaml"
)
# Use async_add_executor_job to perform file I/O asynchronously
def write_dashboard_file():
with open(lovelace_config_file, "w") as f:
yaml.dump(
dashboard_data,
f,
default_flow_style=False,
allow_unicode=True,
)
await self.hass.async_add_executor_job(write_dashboard_file)
_LOGGER.info(
"Successfully created dashboard file: %s", lovelace_config_file
)
# Now update configuration.yaml
try:
config_file = self.hass.config.path("configuration.yaml")
dashboard_config_entry = {
url_path: {
"mode": "yaml",
"title": dashboard_config["title"],
"icon": dashboard_config.get("icon", "mdi:view-dashboard"),
"show_in_sidebar": dashboard_config.get(
"show_in_sidebar", True
),
"filename": f"ui-lovelace-{url_path}.yaml",
}
}
def update_config_file():
try:
with open(config_file, "r") as f:
content = f.read()
# Dashboard configuration to add
dashboard_yaml = f""" {url_path}:
mode: yaml
title: {dashboard_config['title']}
icon: {dashboard_config.get('icon', 'mdi:view-dashboard')}
show_in_sidebar: {str(dashboard_config.get('show_in_sidebar', True)).lower()}
filename: ui-lovelace-{url_path}.yaml"""
# Check if lovelace section exists
if "lovelace:" not in content:
# Add complete lovelace section at the end
lovelace_section = f"""
# Lovelace dashboards configuration added by AI Agent
lovelace:
dashboards:
{dashboard_yaml}
"""
with open(config_file, "a") as f:
f.write(lovelace_section)
return True
# If lovelace exists, check for dashboards section
lines = content.split("\n")
new_lines = []
dashboard_added = False
in_lovelace = False
lovelace_indent = 0
for i, line in enumerate(lines):
new_lines.append(line)
# Detect lovelace section
if (
line.strip() == "lovelace:"
or line.strip().startswith("lovelace:")
):
in_lovelace = True
lovelace_indent = len(line) - len(line.lstrip())
continue
# If we're in lovelace section
if in_lovelace:
current_indent = (
len(line) - len(line.lstrip())
if line.strip()
else 0
)
# If we hit another top-level section, we're out of lovelace
if (
line.strip()
and current_indent <= lovelace_indent
and not line.startswith(" ")
):
if line.strip() != "lovelace:":
in_lovelace = False
# Look for dashboards section
if in_lovelace and "dashboards:" in line:
# Add our dashboard after the dashboards: line
new_lines.append(dashboard_yaml)
dashboard_added = True
in_lovelace = False # We're done
break
# If we found lovelace but no dashboards section, add it
if not dashboard_added and "lovelace:" in content:
# Find lovelace section and add dashboards
new_lines = []
for line in lines:
new_lines.append(line)
if (
line.strip() == "lovelace:"
or line.strip().startswith("lovelace:")
):
# Add dashboards section right after lovelace
new_lines.append(" dashboards:")
new_lines.append(dashboard_yaml)
dashboard_added = True
break
if dashboard_added:
with open(config_file, "w") as f:
f.write("\n".join(new_lines))
return True
else:
# Last resort: append to end of file
with open(config_file, "a") as f:
f.write(f"\n dashboards:\n{dashboard_yaml}\n")
return True
except Exception as e:
_LOGGER.error(
"Failed to update configuration.yaml: %s", str(e)
)
# Fallback to simple append method
try:
with open(config_file, "r") as f:
content = f.read()
# Check if lovelace section exists
if "lovelace:" not in content:
# Add lovelace section
lovelace_config = f"""
# Lovelace dashboards
lovelace:
dashboards:
{url_path}:
mode: yaml
title: {dashboard_config['title']}
icon: {dashboard_config.get('icon', 'mdi:view-dashboard')}
show_in_sidebar: {str(dashboard_config.get('show_in_sidebar', True)).lower()}
filename: ui-lovelace-{url_path}.yaml
"""
with open(config_file, "a") as f:
f.write(lovelace_config)
else:
# Add to existing lovelace section (simple approach)
dashboard_entry = f""" {url_path}:
mode: yaml
title: {dashboard_config['title']}
icon: {dashboard_config.get('icon', 'mdi:view-dashboard')}
show_in_sidebar: {str(dashboard_config.get('show_in_sidebar', True)).lower()}
filename: ui-lovelace-{url_path}.yaml
"""
# Find the dashboards section and add to it
lines = content.split("\n")
new_lines = []
in_dashboards = False
dashboards_indented = False
for line in lines:
new_lines.append(line)
if (
"dashboards:" in line
and "lovelace"
in content[: content.find(line)]
):
in_dashboards = True
# Add our dashboard entry after dashboards:
new_lines.append(dashboard_entry.rstrip())
in_dashboards = False
# If we couldn't find dashboards section, add it under lovelace
if not any("dashboards:" in line for line in lines):
for i, line in enumerate(new_lines):
if line.strip() == "lovelace:":
new_lines.insert(i + 1, " dashboards:")
new_lines.insert(
i + 2, dashboard_entry.rstrip()
)
break
with open(config_file, "w") as f:
f.write("\n".join(new_lines))
return True
except Exception as fallback_error:
_LOGGER.error(
"Fallback config update also failed: %s",
str(fallback_error),
)
return False
config_updated = await self.hass.async_add_executor_job(
update_config_file
)
if config_updated:
success_message = f"""Dashboard '{dashboard_config['title']}' created successfully!
✅ Dashboard file created: ui-lovelace-{url_path}.yaml
✅ Configuration.yaml updated automatically
🔄 Please restart Home Assistant to see your new dashboard in the sidebar."""
return {
"success": True,
"message": success_message,
"url_path": url_path,
"restart_required": True,
}
else:
# Config update failed, provide manual instructions
config_instructions = f"""Dashboard '{dashboard_config['title']}' created successfully!
✅ Dashboard file created: ui-lovelace-{url_path}.yaml
⚠️ Could not automatically update configuration.yaml
Please manually add this to your configuration.yaml:
lovelace:
dashboards:
{url_path}:
mode: yaml
title: {dashboard_config['title']}
icon: {dashboard_config.get('icon', 'mdi:view-dashboard')}
show_in_sidebar: {str(dashboard_config.get('show_in_sidebar', True)).lower()}
filename: ui-lovelace-{url_path}.yaml
Then restart Home Assistant to see your new dashboard in the sidebar."""
return {
"success": True,
"message": config_instructions,
"url_path": url_path,
"restart_required": True,
}
except Exception as config_error:
_LOGGER.error(
"Error updating configuration.yaml: %s", str(config_error)
)
# Provide manual instructions as fallback
config_instructions = f"""Dashboard '{dashboard_config['title']}' created successfully!
✅ Dashboard file created: ui-lovelace-{url_path}.yaml
⚠️ Could not automatically update configuration.yaml
Please manually add this to your configuration.yaml:
lovelace:
dashboards:
{url_path}:
mode: yaml
title: {dashboard_config['title']}
icon: {dashboard_config.get('icon', 'mdi:view-dashboard')}
show_in_sidebar: {str(dashboard_config.get('show_in_sidebar', True)).lower()}
filename: ui-lovelace-{url_path}.yaml
Then restart Home Assistant to see your new dashboard in the sidebar."""
return {
"success": True,
"message": config_instructions,
"url_path": url_path,
"restart_required": True,
}
except Exception as e:
_LOGGER.error("Failed to create dashboard file: %s", str(e))
return {"error": f"Failed to create dashboard file: {str(e)}"}
except Exception as e:
_LOGGER.exception("Error creating dashboard: %s", str(e))
return {"error": f"Error creating dashboard: {str(e)}"}
async def update_dashboard(
self, dashboard_url: str, dashboard_config: Dict[str, Any]
) -> Dict[str, Any]:
"""Update an existing dashboard using Home Assistant's Lovelace WebSocket API."""
try:
_LOGGER.debug(
"Updating dashboard %s with config: %s",
dashboard_url,
json.dumps(dashboard_config, default=str),
)
# Prepare updated dashboard configuration
dashboard_data = {
"title": dashboard_config.get("title", "Updated Dashboard"),
"icon": dashboard_config.get("icon", "mdi:view-dashboard"),
"show_in_sidebar": dashboard_config.get("show_in_sidebar", True),
"require_admin": dashboard_config.get("require_admin", False),
"views": dashboard_config.get("views", []),
}
try:
# Update dashboard file directly
import os
import yaml
# Try updating the YAML file
dashboard_file = self.hass.config.path(
f"ui-lovelace-{dashboard_url}.yaml"
)
# Check if file exists asynchronously
def check_file_exists():
return os.path.exists(dashboard_file)
file_exists = await self.hass.async_add_executor_job(check_file_exists)
if not file_exists:
dashboard_file = self.hass.config.path(
f"dashboards/{dashboard_url}.yaml"
)
file_exists = await self.hass.async_add_executor_job(
lambda: os.path.exists(dashboard_file)
)
if file_exists:
# Use async_add_executor_job to perform file I/O asynchronously
def update_dashboard_file():
with open(dashboard_file, "w") as f:
yaml.dump(
dashboard_data,
f,
default_flow_style=False,
allow_unicode=True,
)
await self.hass.async_add_executor_job(update_dashboard_file)
_LOGGER.info(
"Successfully updated dashboard file: %s", dashboard_file
)
return {
"success": True,
"message": f"Dashboard '{dashboard_url}' updated successfully!",
}
else:
return {"error": f"Dashboard file for '{dashboard_url}' not found"}
except Exception as e:
_LOGGER.error("Failed to update dashboard file: %s", str(e))
return {"error": f"Failed to update dashboard file: {str(e)}"}
except Exception as e:
_LOGGER.exception("Error updating dashboard: %s", str(e))
return {"error": f"Error updating dashboard: {str(e)}"}
async def process_query(
self, user_query: str, provider: Optional[str] = None, debug: bool = False
) -> Dict[str, Any]:
"""Process a user query with input validation and rate limiting."""
try:
if not user_query or not isinstance(user_query, str):
return {"success": False, "error": "Invalid query format"}
# Get the correct configuration for the requested provider
if provider and provider in self.hass.data[DOMAIN]["configs"]:
config = self.hass.data[DOMAIN]["configs"][provider]
else:
config = self.config
_LOGGER.debug(f"Processing query with provider: {provider}")
# Log sanitized config (masks all tokens/keys for security)
_LOGGER.debug(
f"Using config: {json.dumps(sanitize_for_logging(config), default=str)}"
)
selected_provider = provider or config.get("ai_provider", "llama")
models_config = config.get("models", {})
provider_config = {
"openai": {
"token_key": "openai_token", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get("openai", "gpt-3.5-turbo"),
"client_class": OpenAIClient,
},
"gemini": {
"token_key": "gemini_token", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get("gemini", "gemini-1.5-flash"),
"client_class": GeminiClient,
},
"openrouter": {
"token_key": "openrouter_token", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get("openrouter", "openai/gpt-4o"),
"client_class": OpenRouterClient,
},
"llama": {
"token_key": "llama_token", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get(
"llama", "Llama-4-Maverick-17B-128E-Instruct-FP8"
),
"client_class": LlamaClient,
},
"anthropic": {
"token_key": "anthropic_token", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get(
"anthropic", "claude-sonnet-4-5-20250929"
),
"client_class": AnthropicClient,
},
"alter": {
"token_key": "alter_token", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get("alter", ""),
"client_class": AlterClient,
},
"zai": {
"token_key": "zai_token", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get("zai", ""),
"client_class": ZaiClient,
},
"local_ollama": {
"token_key": "local_ollama_url", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get("local_ollama", ""),
"client_class": LocalOllamaClient,
},
"openai_compatible": {
"token_key": "openai_compatible_url", # nosec B105 - dict-key field name, not a credential value (false positive)
"model": models_config.get("openai_compatible", ""),
"client_class": OpenaiCompatibleClient,
},
}
# Validate provider and get configuration
if selected_provider not in provider_config:
_LOGGER.warning(
f"Invalid provider {selected_provider}, falling back to llama"
)
selected_provider = "llama"
provider_settings = provider_config[selected_provider]
token = self.config.get(provider_settings["token_key"])
def _with_debug(result: Dict[str, Any]) -> Dict[str, Any]:
"""Attach a sanitized trace when UI requests debug info."""
if debug and "debug" not in result:
result["debug"] = self._build_debug_trace(
selected_provider,
provider_settings,
config.get("zai_endpoint", "general"),
)
return result
# Validate token/URL
if not token:
is_url_provider = selected_provider in (
"local_ollama",
"openai_compatible",
)
error_msg = f"No {'URL' if is_url_provider else 'token'} configured for provider {selected_provider}"
_LOGGER.error(error_msg)
return _with_debug({"success": False, "error": error_msg})
# Initialize client
try:
if selected_provider == "zai":
# ZaiClient takes (token, model, endpoint_type)
endpoint_type = config.get("zai_endpoint", "general")
self.ai_client = provider_settings["client_class"](
token=token,
model=provider_settings["model"],
endpoint_type=endpoint_type,
)
_LOGGER.debug(
f"Initialized {selected_provider} client with model {provider_settings['model']}, endpoint_type {endpoint_type}"
)
elif selected_provider in ("local_ollama", "openai_compatible"):
# LocalOllamaClient and OpenaiCompatibleClient take (url, model)
if selected_provider == "local_ollama":
# Support legacy local_url
url = token or config.get("local_url")
self.ai_client = provider_settings["client_class"](
url, provider_settings["model"]
)
else:
url = token
api_key = config.get("openai_compatible_api_key", "") or ""
self.ai_client = provider_settings["client_class"](
url, provider_settings["model"], api_key or None
)
_LOGGER.debug(
f"Initialized {selected_provider} client with model {provider_settings['model']}"
)
else:
# Other clients take (token, model)
self.ai_client = provider_settings["client_class"](
token=token, model=provider_settings["model"]
)
_LOGGER.debug(
f"Initialized {selected_provider} client with model {provider_settings['model']}"
)
except Exception as e:
error_msg = f"Error initializing {selected_provider} client: {str(e)}"
_LOGGER.error(error_msg)
return _with_debug({"success": False, "error": error_msg})
# Process the query with rate limiting and retries
if not self._check_rate_limit():
return _with_debug(
{
"success": False,
"error": "Rate limit exceeded. Please wait before trying again.",
}
)
# Sanitize user input
user_query = user_query.strip()[:1000] # Limit length and trim whitespace
_LOGGER.debug("Processing new query: %s", user_query)
# Check cache for identical query
cache_key = f"query_{hash(user_query)}_{provider}_{debug}"
cached_result = self._get_cached_data(cache_key)
if cached_result:
return (
dict(cached_result)
if isinstance(cached_result, dict)
else {"error": "Invalid cached result"}
)
# Add system message to conversation if it's the first message
if not self.conversation_history:
_LOGGER.debug("Adding system message to new conversation")
self.conversation_history.append(self.system_prompt)
# Remember where this query started so a failure can be rolled
# back instead of leaving dangling/oversized messages that poison
# every subsequent query (issue #80).
history_checkpoint = len(self.conversation_history)
# Add user query to conversation
self.conversation_history.append({"role": "user", "content": user_query})
_LOGGER.debug("Added user query to conversation history")
# Prevent infinite loops while leaving room for multi-step
# discovery: with data responses capped (issue #80) the model may
# need several narrower data requests instead of one big dump.
max_iterations = 8
iteration = 0
while iteration < max_iterations:
iteration += 1
_LOGGER.debug(f"Processing iteration {iteration} of {max_iterations}")
try:
# Get AI response
_LOGGER.debug("Requesting response from AI provider")
response = await self._get_ai_response()
_LOGGER.debug("Received response from AI provider: %s", response)
try:
# Try to parse the response as JSON with simplified approach
response_clean = response.strip()
# Remove potential BOM and other invisible characters
import codecs
if response_clean.startswith(codecs.BOM_UTF8.decode("utf-8")):
response_clean = response_clean[1:]
# Remove other common invisible characters
invisible_chars = [
"\ufeff",
"\u200b",
"\u200c",
"\u200d",
"\u2060",
]
for char in invisible_chars:
response_clean = response_clean.replace(char, "")
_LOGGER.debug(
"Cleaned response length: %d", len(response_clean)
)
_LOGGER.debug(
"Cleaned response first 100 chars: %s", response_clean[:100]
)
_LOGGER.debug(
"Cleaned response last 100 chars: %s", response_clean[-100:]
)
# Simple strategy: try to parse the cleaned response directly
response_data = None
try:
_LOGGER.debug("Attempting basic JSON parse...")
response_data = json.loads(response_clean)
_LOGGER.debug("Basic JSON parse succeeded!")
except json.JSONDecodeError as e:
_LOGGER.warning("Basic JSON parse failed: %s", str(e))
_LOGGER.debug("JSON error position: %d", e.pos)
if e.pos < len(response_clean):
_LOGGER.debug(
"Character at error position: %s (ord: %d)",
repr(response_clean[e.pos]),
ord(response_clean[e.pos]),
)
_LOGGER.debug(
"Context around error: %s",
repr(
response_clean[max(0, e.pos - 10) : e.pos + 10]
),
)
# Fallback: try to extract JSON by finding the first { and last }
json_start = response_clean.find("{")
json_end = response_clean.rfind("}")
if (
json_start != -1
and json_end != -1
and json_end > json_start
):
json_part = response_clean[json_start : json_end + 1]
_LOGGER.debug(
"Trying fallback extraction from pos %d to %d",
json_start,
json_end,
)
_LOGGER.debug("Extracted JSON: %s", json_part[:200])
try:
response_data = json.loads(json_part)
_LOGGER.debug("Fallback JSON extraction succeeded!")
except json.JSONDecodeError as e2:
_LOGGER.warning(
"Fallback JSON extraction also failed: %s",
str(e2),
)
raise e # Re-raise the original error
else:
_LOGGER.warning(
"Could not find JSON boundaries in response"
)
raise e # Re-raise the original error
if response_data is None:
raise json.JSONDecodeError(
"All parsing strategies failed", response_clean, 0
)
_LOGGER.debug("Successfully parsed JSON response")
_LOGGER.debug(
"Parsed response type: %s",
response_data.get("request_type", "unknown"),
)
# Check if this is a data request (either format)
data_request_types = [
"get_entity_state",
"get_entities_by_domain",
"get_entities_by_device_class",
"get_climate_related_entities",
"get_entities_by_area",
"get_entities",
"get_calendar_events",
"get_automations",
"get_entity_registry",
"get_device_registry",
"get_weather_data",
"get_area_registry",
"get_history",
"get_person_data",
"get_statistics",
"get_scenes",
"get_dashboards",
"get_dashboard_config",
"set_entity_state",
"create_automation",
"create_dashboard",
"update_dashboard",
]
if (
response_data.get("request_type") == "data_request"
or response_data.get("request_type") in data_request_types
):
# Handle data request (both standard format and direct request type)
if response_data.get("request_type") == "data_request":
request_type = response_data.get("request")
else:
request_type = response_data.get("request_type")
parameters = response_data.get("parameters", {})
_LOGGER.debug(
"Processing data request: %s with parameters: %s",
request_type,
json.dumps(parameters),
)
# Add AI's response to conversation history
self.conversation_history.append(
{
"role": "assistant",
"content": json.dumps(
response_data
), # Store clean JSON
}
)
# Get requested data
data: Union[Dict[str, Any], List[Dict[str, Any]]]
if request_type == "get_entity_state":
data = await self.get_entity_state(
parameters.get("entity_id")
)
elif request_type == "get_entities_by_domain":
data = await self.get_entities_by_domain(
parameters.get("domain")
)
elif request_type == "get_entities_by_area":
data = await self.get_entities_by_area(
parameters.get("area_id")
)
elif request_type == "get_entities":
data = await self.get_entities(
area_id=parameters.get("area_id"),
area_ids=parameters.get("area_ids"),
)
elif request_type == "get_entities_by_device_class":
data = await self.get_entities_by_device_class(
parameters.get("device_class"),
parameters.get("domain"),
)
elif request_type == "get_climate_related_entities":
data = await self.get_climate_related_entities()
elif request_type == "get_calendar_events":
data = await self.get_calendar_events(
parameters.get("entity_id")
)
elif request_type == "get_automations":
data = await self.get_automations()
elif request_type == "get_entity_registry":
data = await self.get_entity_registry()
elif request_type == "get_device_registry":
data = await self.get_device_registry()
elif request_type == "get_weather_data":
data = await self.get_weather_data()
elif request_type == "get_area_registry":
data = await self.get_area_registry()
elif request_type == "get_history":
data = await self.get_history(
parameters.get("entity_id"),
parameters.get("hours", 24),
)
elif request_type == "get_person_data":
data = await self.get_person_data()
elif request_type == "get_statistics":
data = await self.get_statistics(
parameters.get("entity_id")
)
elif request_type == "get_scenes":
data = await self.get_scenes()
elif request_type == "get_dashboards":
data = await self.get_dashboards()
elif request_type == "get_dashboard_config":
data = await self.get_dashboard_config(
parameters.get("dashboard_url")
)
elif request_type == "set_entity_state":
data = await self.set_entity_state(
parameters.get("entity_id"),
parameters.get("state"),
parameters.get("attributes"),
)
elif request_type == "create_automation":
data = await self.create_automation(
parameters.get("automation")
)
elif request_type == "create_dashboard":
data = await self.create_dashboard(
parameters.get("dashboard_config")
)
elif request_type == "update_dashboard":
data = await self.update_dashboard(
parameters.get("dashboard_url"),
parameters.get("dashboard_config"),
)
else:
data = {
"error": f"Unknown request type: {request_type}"
}
_LOGGER.warning(
"Unknown request type: %s", request_type
)
# Check if any data request resulted in an error
if isinstance(data, dict) and "error" in data:
return _with_debug(
{"success": False, "error": data["error"]}
)
elif isinstance(data, list) and any(
"error" in item
for item in data
if isinstance(item, dict)
):
errors = [
item["error"]
for item in data
if isinstance(item, dict) and "error" in item
]
return _with_debug(
{"success": False, "error": "; ".join(errors)}
)
_LOGGER.debug(
"Retrieved data for request: %s",
json.dumps(data, default=str),
)
# Add data to conversation as a user message (not system to avoid overwriting system prompt in Anthropic API)
self.conversation_history.append(
{
"role": "user",
"content": self._format_data_message(data),
}
)
continue
elif response_data.get("request_type") == "final_response":
# Add final response to conversation history
self.conversation_history.append(
{
"role": "assistant",
"content": json.dumps(
response_data
), # Store clean JSON
}
)
# Return final response
_LOGGER.debug(
"Received final response: %s",
response_data.get("response"),
)
result = {
"success": True,
"answer": response_data.get("response", ""),
}
result = _with_debug(result)
self._set_cached_data(cache_key, result)
return result
elif (
response_data.get("request_type") == "automation_suggestion"
):
# Add automation suggestion to conversation history
self.conversation_history.append(
{
"role": "assistant",
"content": json.dumps(
response_data
), # Store clean JSON
}
)
# Return automation suggestion
_LOGGER.debug(
"Received automation suggestion: %s",
json.dumps(response_data.get("automation")),
)
result = {
"success": True,
"answer": json.dumps(response_data),
}
result = _with_debug(result)
self._set_cached_data(cache_key, result)
return result
elif (
response_data.get("request_type") == "dashboard_suggestion"
):
# Add dashboard suggestion to conversation history
self.conversation_history.append(
{
"role": "assistant",
"content": json.dumps(
response_data
), # Store clean JSON
}
)
# Return dashboard suggestion
_LOGGER.debug(
"Received dashboard suggestion: %s",
json.dumps(response_data.get("dashboard")),
)
result = {
"success": True,
"answer": json.dumps(response_data),
}
result = _with_debug(result)
self._set_cached_data(cache_key, result)
return result
elif response_data.get("request_type") in [
"get_entities",
"get_entities_by_area",
]:
# Handle direct get_entities request (for backward compatibility)
parameters = response_data.get("parameters", {})
_LOGGER.debug(
"Processing direct get_entities request with parameters: %s",
json.dumps(parameters),
)
# Add AI's response to conversation history
self.conversation_history.append(
{
"role": "assistant",
"content": json.dumps(
response_data
), # Store clean JSON
}
)
# Get entities data
if response_data.get("request_type") == "get_entities":
data = await self.get_entities(
area_id=parameters.get("area_id"),
area_ids=parameters.get("area_ids"),
)
else: # get_entities_by_area
data = await self.get_entities_by_area(
parameters.get("area_id")
)
_LOGGER.debug(
"Retrieved %d entities",
len(data) if isinstance(data, list) else 1,
)
# Add data to conversation as a user message (not system to avoid overwriting system prompt in Anthropic API)
self.conversation_history.append(
{
"role": "user",
"content": self._format_data_message(data),
}
)
continue
elif response_data.get("request_type") == "call_service":
# Handle service call request
domain = response_data.get("domain")
service = response_data.get("service")
target = response_data.get("target", {})
service_data = response_data.get("service_data", {})
# Resolve nested requests in target
if target and "entity_id" in target:
entity_id_value = target["entity_id"]
if (
isinstance(entity_id_value, dict)
and "request_type" in entity_id_value
):
# This is a nested request, resolve it
nested_request_type = entity_id_value.get(
"request_type"
)
nested_parameters = entity_id_value.get(
"parameters", {}
)
_LOGGER.debug(
"Resolving nested request: %s with parameters: %s",
nested_request_type,
json.dumps(nested_parameters),
)
# Resolve the nested request
if nested_request_type == "get_entities":
entities_data = await self.get_entities(
area_id=nested_parameters.get("area_id"),
area_ids=nested_parameters.get("area_ids"),
)
elif nested_request_type == "get_entities_by_area":
entities_data = await self.get_entities_by_area(
nested_parameters.get("area_id")
)
elif (
nested_request_type == "get_entities_by_domain"
):
entities_data = (
await self.get_entities_by_domain(
nested_parameters.get("domain")
)
)
else:
_LOGGER.error(
"Unsupported nested request type: %s",
nested_request_type,
)
return {
"success": False,
"error": f"Unsupported nested request type: {nested_request_type}",
}
# Extract entity IDs from the resolved data
if isinstance(entities_data, list):
entity_ids = [
entity.get("entity_id")
for entity in entities_data
if entity.get("entity_id")
]
target["entity_id"] = entity_ids
_LOGGER.debug(
"Resolved nested request to entity IDs: %s",
entity_ids,
)
else:
_LOGGER.error(
"Nested request returned unexpected data format"
)
return _with_debug(
{
"success": False,
"error": "Nested request returned unexpected data format",
}
)
# Handle backward compatibility with old format
if not domain or not service:
request = response_data.get("request")
parameters = response_data.get("parameters", {})
if request and "entity_id" in parameters:
entity_id = parameters["entity_id"]
# Infer domain from entity_id
if "." in entity_id:
domain = entity_id.split(".")[0]
service = request
target = {"entity_id": entity_id}
# Remove entity_id from parameters to avoid duplication
service_data = {
k: v
for k, v in parameters.items()
if k != "entity_id"
}
_LOGGER.debug(
"Converted old format: domain=%s, service=%s",
domain,
service,
)
_LOGGER.debug(
"Processing service call: %s.%s with target: %s and data: %s",
domain,
service,
json.dumps(target),
json.dumps(service_data),
)
# Add AI's response to conversation history
self.conversation_history.append(
{
"role": "assistant",
"content": json.dumps(
response_data
), # Store clean JSON
}
)
# Call the service
data = await self.call_service(
domain, service, target, service_data
)
# Check if service call resulted in an error
if isinstance(data, dict) and "error" in data:
return _with_debug(
{"success": False, "error": data["error"]}
)
_LOGGER.debug(
"Service call completed: %s",
json.dumps(data, default=str),
)
# Add data to conversation as a user message (not system to avoid overwriting system prompt in Anthropic API)
self.conversation_history.append(
{
"role": "user",
"content": self._format_data_message(data),
}
)
# Go to next iteration to continue the loop
continue
# Unknown request type
_LOGGER.warning(
"Unknown response type: %s",
response_data.get("request_type"),
)
return _with_debug(
{
"success": False,
"error": f"Unknown response type: {response_data.get('request_type')}",
}
)
except json.JSONDecodeError as e:
# Check if this is a local provider that might have already wrapped the response
provider = self.config.get("ai_provider", "unknown")
if provider in ("local_ollama", "openai_compatible"):
_LOGGER.debug(
"Local provider returned non-JSON response (this is normal and handled): %s",
response[:200],
)
else:
# Log more of the response to help with debugging for non-local providers
response_preview = (
response[:1000] if len(response) > 1000 else response
)
_LOGGER.warning(
"Failed to parse response as JSON: %s. Response length: %d. Response preview: %s",
str(e),
len(response),
response_preview,
)
# Log additional debugging information
_LOGGER.debug(
"First 50 characters as bytes: %s",
response[:50].encode("utf-8") if response else b"",
)
_LOGGER.debug(
"Response starts with: %s",
repr(response[:10]) if response else "None",
)
# Also log the response to a separate debug file for detailed analysis (non-local providers only)
if provider not in ("local_ollama", "openai_compatible"):
try:
import os
debug_dir = "/config/ai_agent_ha_debug"
def write_debug_file():
if not os.path.exists(debug_dir):
os.makedirs(debug_dir)
import datetime
timestamp = datetime.datetime.now().strftime(
"%Y%m%d_%H%M%S"
)
debug_file = os.path.join(
debug_dir, f"failed_response_{timestamp}.txt"
)
with open(debug_file, "w", encoding="utf-8") as f:
f.write(f"Timestamp: {timestamp}\n")
f.write(f"Provider: {provider}\n")
f.write(f"Error: {str(e)}\n")
f.write(f"Response length: {len(response)}\n")
f.write(
f"Response bytes: {response.encode('utf-8') if response else b''}\n"
)
f.write(f"Response repr: {repr(response)}\n")
f.write(f"Full response:\n{response}\n")
return debug_file
# Run file operations in executor to avoid blocking
debug_file = await self.hass.async_add_executor_job(
write_debug_file
)
_LOGGER.info(
"Failed response saved to debug file: %s",
debug_file,
)
except Exception as debug_error:
_LOGGER.debug(
"Could not save debug file: %s", str(debug_error)
)
# Check if this looks like a corrupted automation suggestion
if (
response.strip().startswith(
'{"request_type": "automation_suggestion'
)
and len(response) > 10000
and response.count("for its use in various fields") > 50
):
_LOGGER.warning(
"Detected corrupted automation suggestion response with repetitive text"
)
result = _with_debug(
{
"success": False,
"error": "AI generated corrupted automation response. Please try again with a more specific automation request.",
}
)
self._set_cached_data(cache_key, result)
return result
# If response is not valid JSON, try to wrap it as a final response
try:
# Truncate extremely long responses to prevent memory issues
response_to_wrap = response
if len(response) > 50000:
response_to_wrap = (
response[:5000]
+ "... [Response truncated due to excessive length]"
)
_LOGGER.warning(
"Truncated extremely long response from %d to 5000 characters",
len(response),
)
wrapped_response = {
"request_type": "final_response",
"response": response_to_wrap,
}
# Keep the conversation paired: record the assistant
# reply so history doesn't end with a dangling user
# message (issue #80).
self.conversation_history.append(
{"role": "assistant", "content": response_to_wrap}
)
result = {
"success": True,
"answer": json.dumps(wrapped_response),
}
_LOGGER.debug("Wrapped non-JSON response as final_response")
except Exception as wrap_error:
_LOGGER.error(
"Failed to wrap response: %s", str(wrap_error)
)
result = {
"success": False,
"error": f"Invalid response format: {str(e)}",
}
result = _with_debug(result)
self._set_cached_data(cache_key, result)
return result
except Exception as e:
_LOGGER.exception("Error processing AI response: %s", str(e))
# Roll back this query's messages so the failure doesn't
# poison subsequent queries (issue #80).
del self.conversation_history[history_checkpoint:]
return _with_debug(
{
"success": False,
"error": f"Error processing AI response: {str(e)}",
}
)
# If we've reached max iterations without a final response
_LOGGER.warning("Reached maximum iterations without final response")
# Roll back this query's messages so the failure doesn't poison
# subsequent queries (issue #80).
del self.conversation_history[history_checkpoint:]
result = {
"success": False,
"error": "Maximum iterations reached without final response",
}
result = _with_debug(result)
self._set_cached_data(cache_key, result)
return result
except Exception as e:
_LOGGER.exception("Error in process_query: %s", str(e))
return _with_debug(
{"success": False, "error": f"Error in process_query: {str(e)}"}
)
def _build_debug_trace(
self,
provider: Optional[str],
provider_settings: Optional[Dict[str, Any]],
endpoint_type: Optional[str],
) -> Dict[str, Any]:
"""Return a sanitized snapshot of the HA↔AI conversation for UI display."""
history_tail = (
self.conversation_history[-20:] if self.conversation_history else []
)
return {
"provider": provider,
"model": provider_settings.get("model") if provider_settings else None,
"endpoint_type": endpoint_type,
"conversation": history_tail,
}
async def _get_ai_response(self) -> str:
"""Get response from the selected AI provider with retries and rate limiting."""
if not self._check_rate_limit():
raise Exception("Rate limit exceeded. Please try again later.")
retry_count = 0
last_error = None
# Limit conversation history to the last 10 messages to prevent token
# overflow, and cap the window's total size as well: even
# individually-capped data messages can stack up past context and
# per-minute token budgets (issue #80). Most recent messages win.
candidates = [
m for m in self.conversation_history[-10:] if m.get("role") != "system"
]
recent_messages: List[Dict[str, Any]] = []
total_chars = 0
for message in reversed(candidates):
size = len(str(message.get("content") or ""))
if recent_messages and total_chars + size > self.MAX_WINDOW_CHARS:
break
recent_messages.insert(0, message)
total_chars += size
# Dropping/slicing can cut mid-turn; ensure the window starts with a
# user turn (issue #80).
while recent_messages and recent_messages[0].get("role") == "assistant":
recent_messages.pop(0)
# System prompt is always the first message
recent_messages = [self.system_prompt] + recent_messages
_LOGGER.debug("Sending %d messages to AI provider", len(recent_messages))
_LOGGER.debug("AI provider: %s", self.config.get("ai_provider", "unknown"))
while retry_count < self._max_retries:
try:
_LOGGER.debug(
"Attempt %d/%d: Calling AI client",
retry_count + 1,
self._max_retries,
)
response = await self.ai_client.get_response(recent_messages)
# Every client is expected to return a string. Guard against a
# client handing back a non-string (e.g. a raw list/dict from an
# unexpected provider response shape) so the downstream string
# operations below don't crash with an unhelpful AttributeError
# and burn all retries (see issue #75).
if response is not None and not isinstance(response, str):
_LOGGER.warning(
"AI client returned non-string response of type %s; coercing to str",
type(response).__name__,
)
response = (
json.dumps(response)
if isinstance(response, (list, dict))
else str(response)
)
_LOGGER.debug(
"AI client returned response of length: %d", len(response or "")
)
_LOGGER.debug("AI response preview: %s", (response or "")[:200])
# Check for extremely long responses that might indicate model issues
if response and len(response) > 50000:
_LOGGER.warning(
"AI returned extremely long response (%d characters), this may indicate a model issue",
len(response),
)
# Check for repetitive patterns that indicate a corrupted response
if response.count("for its use in various fields") > 50:
_LOGGER.error(
"Detected corrupted repetitive response, aborting this iteration"
)
raise Exception(
"AI generated corrupted response with repetitive text. Please try again with a clearer request."
)
# Check if response is empty
if not response or response.strip() == "":
_LOGGER.warning(
"AI client returned empty response on attempt %d",
retry_count + 1,
)
if retry_count + 1 >= self._max_retries:
raise Exception(
"AI provider returned empty response after all retries"
)
else:
retry_count += 1
await asyncio.sleep(self._retry_delay * retry_count)
continue
return str(response)
except NonRetryableAIError:
# Deterministic client error (e.g. 400 "prompt is too long") -
# retrying the same payload cannot succeed (issue #80).
raise
except Exception as e:
_LOGGER.error(
"AI client error on attempt %d: %s", retry_count + 1, str(e)
)
last_error = e
retry_count += 1
if retry_count < self._max_retries:
delay: float = self._retry_delay * retry_count
if isinstance(e, RateLimitedAIError) and e.retry_after:
# Wait at least as long as the provider asked for
# (capped at 60s) so per-minute token windows can
# actually reset (issue #80).
delay = max(delay, min(e.retry_after, 60))
await asyncio.sleep(delay)
continue
raise Exception(
f"Failed after {retry_count} retries. Last error: {str(last_error)}"
)
def clear_conversation_history(self) -> None:
"""Clear the conversation history and cache."""
self.conversation_history = []
self._cache.clear()
_LOGGER.debug("Conversation history and cache cleared")
async def set_entity_state(
self, entity_id: str, state: str, attributes: Optional[Dict[str, Any]] = None
) -> Dict[str, Any]:
"""Set the state of an entity."""
try:
_LOGGER.debug(
"Setting state for entity %s to %s with attributes: %s",
entity_id,
state,
json.dumps(attributes or {}),
)
# Validate entity exists
if not self.hass.states.get(entity_id):
return {"error": f"Entity {entity_id} not found"}
# Call the appropriate service based on the domain
domain = entity_id.split(".")[0]
if domain == "light":
service = (
"turn_on" if state.lower() in ["on", "true", "1"] else "turn_off"
)
service_data = {"entity_id": entity_id}
if attributes and service == "turn_on":
service_data.update(attributes)
await self.hass.services.async_call("light", service, service_data)
elif domain == "switch":
service = (
"turn_on" if state.lower() in ["on", "true", "1"] else "turn_off"
)
await self.hass.services.async_call(
"switch", service, {"entity_id": entity_id}
)
elif domain == "cover":
if state.lower() in ["open", "up"]:
service = "open_cover"
elif state.lower() in ["close", "down"]:
service = "close_cover"
elif state.lower() == "stop":
service = "stop_cover"
else:
return {"error": f"Invalid state {state} for cover entity"}
await self.hass.services.async_call(
"cover", service, {"entity_id": entity_id}
)
elif domain == "climate":
service_data = {"entity_id": entity_id}
if state.lower() in ["on", "true", "1"]:
service = "turn_on"
elif state.lower() in ["off", "false", "0"]:
service = "turn_off"
elif state.lower() in ["heat", "cool", "dry", "fan_only", "auto"]:
service = "set_hvac_mode"
service_data["hvac_mode"] = state.lower()
else:
return {"error": f"Invalid state {state} for climate entity"}
await self.hass.services.async_call("climate", service, service_data)
elif domain == "fan":
service = (
"turn_on" if state.lower() in ["on", "true", "1"] else "turn_off"
)
service_data = {"entity_id": entity_id}
if attributes and service == "turn_on":
service_data.update(attributes)
await self.hass.services.async_call("fan", service, service_data)
else:
# For other domains, try to set the state directly
self.hass.states.async_set(entity_id, state, attributes or {})
# Get the new state to confirm the change
new_state = self.hass.states.get(entity_id)
return {
"success": True,
"entity_id": entity_id,
"new_state": new_state.state,
"new_attributes": new_state.attributes,
}
except Exception as e:
_LOGGER.exception("Error setting entity state: %s", str(e))
return {"error": f"Error setting entity state: {str(e)}"}
async def call_service(
self,
domain: str,
service: str,
target: Optional[Dict[str, Any]] = None,
service_data: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""Call a Home Assistant service."""
try:
_LOGGER.debug(
"Calling service %s.%s with target: %s and data: %s",
domain,
service,
json.dumps(target or {}),
json.dumps(service_data or {}),
)
# Prepare the service call data
call_data = {}
# Add target entities if provided
if target:
if "entity_id" in target:
entity_ids = target["entity_id"]
if isinstance(entity_ids, list):
call_data["entity_id"] = entity_ids
else:
call_data["entity_id"] = [entity_ids]
# Add other target properties
for key, value in target.items():
if key != "entity_id":
call_data[key] = value
# Add service data if provided
if service_data:
call_data.update(service_data)
_LOGGER.debug("Final service call data: %s", json.dumps(call_data))
# Call the service
await self.hass.services.async_call(domain, service, call_data)
# Get the updated states of affected entities
result_entities = []
if "entity_id" in call_data:
for entity_id in call_data["entity_id"]:
state = self.hass.states.get(entity_id)
if state:
result_entities.append(
{
"entity_id": entity_id,
"state": state.state,
"attributes": dict(state.attributes),
}
)
return {
"success": True,
"service": f"{domain}.{service}",
"entities_affected": result_entities,
"message": f"Successfully called {domain}.{service}",
}
except Exception as e:
_LOGGER.exception(
"Error calling service %s.%s: %s", domain, service, str(e)
)
return {"error": f"Error calling service {domain}.{service}: {str(e)}"}
async def save_user_prompt_history(
self, user_id: str, history: List[str]
) -> Dict[str, Any]:
"""Save user's prompt history to HA storage."""
try:
store: Store = Store(self.hass, 1, f"ai_agent_ha_history_{user_id}")
await store.async_save({"history": history})
return {"success": True}
except Exception as e:
_LOGGER.exception("Error saving prompt history: %s", str(e))
return {"error": f"Error saving prompt history: {str(e)}"}
async def load_user_prompt_history(self, user_id: str) -> Dict[str, Any]:
"""Load user's prompt history from HA storage."""
try:
store: Store = Store(self.hass, 1, f"ai_agent_ha_history_{user_id}")
data = await store.async_load()
history = data.get("history", []) if data else []
return {"success": True, "history": history}
except Exception as e:
_LOGGER.exception("Error loading prompt history: %s", str(e))
return {"error": f"Error loading prompt history: {str(e)}", "history": []}