Files
HomeAssistance/custom_components/ai_automation_suggester/config_flow.py
T
2026-06-05 22:34:31 -04:00

691 lines
36 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.
# custom_components/ai_automation_suggester/config_flow.py
"""Config flow for AI Automation Suggester."""
from __future__ import annotations
import logging
from typing import Any, Dict, Optional
import aiohttp
import voluptuous as vol
from homeassistant import config_entries
from homeassistant.core import callback
from homeassistant.helpers.aiohttp_client import async_get_clientsession
from homeassistant.helpers.selector import TextSelector, TextSelectorConfig
from .const import *
from .endpoint_utils import (
bearer_auth_headers,
ensure_http_url,
ollama_api_candidates,
ollama_base_url,
openai_model_endpoint_candidates,
)
_LOGGER = logging.getLogger(__name__)
# ─────────────────────────────────────────────────────────────
# Lightweight provider validators (unchanged)
# ─────────────────────────────────────────────────────────────
class ProviderValidator:
"""Ping each provider with a dummy request to validate credentials."""
def __init__(self, hass, request_timeout: int | None = None):
self.session = async_get_clientsession(hass)
self.timeout = aiohttp.ClientTimeout(total=max(10, int(request_timeout or DEFAULT_REQUEST_TIMEOUT)))
async def validate_openai(self, api_key: str) -> Optional[str]:
hdr = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
try:
resp = await self.session.get("https://api.openai.com/v1/models", headers=hdr, timeout=self.timeout)
return None if resp.status == 200 else await resp.text()
except Exception as err: # noqa: BLE001
return str(err)
async def validate_anthropic(self, api_key: str, model: str) -> Optional[str]:
hdr = {
"x-api-key": api_key,
"anthropic-version": VERSION_ANTHROPIC,
"content-type": "application/json",
}
try:
resp = await self.session.get("https://api.anthropic.com/v1/models", headers=hdr, timeout=self.timeout)
return None if resp.status == 200 else await resp.text()
except Exception as err:
return str(err)
async def validate_google(self, api_key: str, model: str) -> Optional[str]:
try:
resp = await self.session.get(
f"https://generativelanguage.googleapis.com/v1beta/models/{model}?key={api_key}",
timeout=self.timeout,
)
return None if resp.status == 200 else await resp.text()
except Exception as err:
return str(err)
async def validate_groq(self, api_key: str) -> Optional[str]:
hdr = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
try:
resp = await self.session.get("https://api.groq.com/openai/v1/models", headers=hdr, timeout=self.timeout)
return None if resp.status == 200 else await resp.text()
except Exception as err:
return str(err)
async def validate_localai(self, ip: str, port: int, https: bool) -> Optional[str]:
proto = "https" if https else "http"
try:
resp = await self.session.get(f"{proto}://{ip}:{port}/v1/models", timeout=self.timeout)
return None if resp.status == 200 else await resp.text()
except Exception as err:
return str(err)
async def validate_ollama(
self,
ip: str | None,
port: int | None,
https: bool,
base_url: str | None = None,
api_key: str | None = None,
) -> Optional[str]:
base = ollama_base_url(base_url=base_url, ip_address=ip, port=port, https=https)
if not base:
return "Ollama host/port or base URL is required"
last_error = None
headers = bearer_auth_headers(api_key)
try:
for endpoint in ollama_api_candidates(base, "api/tags"):
resp = await self.session.get(endpoint, headers=headers, timeout=self.timeout)
if resp.status == 200:
return None
last_error = f"{endpoint}: {resp.status} {await resp.text()}"
return last_error
except Exception as err:
return str(err)
async def validate_custom_openai(self, endpoint: str, api_key: str | None) -> Optional[str]:
hdr = {"Content-Type": "application/json"}
if api_key:
hdr["Authorization"] = f"Bearer {api_key}"
last_error = None
try:
for model_endpoint in openai_model_endpoint_candidates(endpoint):
resp = await self.session.get(model_endpoint, headers=hdr, timeout=self.timeout)
if resp.status == 200:
return None
last_error = f"{model_endpoint}: {resp.status} {await resp.text()}"
return last_error or "No valid model endpoint could be built"
except Exception as err:
return str(err)
async def validate_perplexity(self, api_key: str, model: str) -> Optional[str]:
hdr = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
# Perplexity 'sonar' models require max_tokens >= 16; a smaller value is
# rejected with a 400 during validation (issue #171).
payload = {"model": model, "messages": [{"role": "user", "content": "ping"}], "max_tokens": 16}
try:
resp = await self.session.post(ENDPOINT_PERPLEXITY, headers=hdr, json=payload, timeout=self.timeout)
return None if resp.status == 200 else await resp.text()
except Exception as err:
return str(err)
async def validate_openrouter(self, api_key: str, model: str) -> Optional[str]:
hdr = {"content-type": "application/json"}
if api_key:
hdr["Authorization"] = f"Bearer {api_key}"
try:
resp = await self.session.get(
"https://openrouter.ai/api/v1/models", headers=hdr, timeout=self.timeout
)
return None if resp.status == 200 else await resp.text()
except Exception as err:
return str(err)
async def validate_generic_openai(self, endpoint: str, api_key: str) -> Optional[str]:
hdr = {"Content-Type": "application/json"}
if api_key:
hdr["Authorization"] = f"Bearer {api_key}"
try:
resp = await self.session.get(ensure_http_url(endpoint), headers=hdr, timeout=self.timeout)
return None if resp.status == 200 else await resp.text()
except Exception as err:
return str(err)
# ─────────────────────────────────────────────────────────────
# Configflow main class
# ─────────────────────────────────────────────────────────────
class AIAutomationConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
"""Handle integration setup via the UI."""
VERSION = CONFIG_VERSION
def __init__(self) -> None:
self.provider: str | None = None
self.data: Dict[str, Any] = {}
self.validator: ProviderValidator | None = None
# ───────── Initial provider choice ─────────
async def async_step_user(self, user_input: Dict[str, Any] | None = None):
errors: Dict[str, str] = {}
if user_input:
self.provider = user_input[CONF_PROVIDER]
self.data.update(user_input)
return await {
"OpenAI": self.async_step_openai,
"Anthropic": self.async_step_anthropic,
"Google": self.async_step_google,
"Groq": self.async_step_groq,
"LocalAI": self.async_step_localai,
"Ollama": self.async_step_ollama,
"Custom OpenAI": self.async_step_custom_openai,
"Mistral AI": self.async_step_mistral,
"Perplexity AI": self.async_step_perplexity,
"OpenRouter": self.async_step_openrouter,
"OpenAI Azure": self.async_step_openai_azure,
"Generic OpenAI": self.async_step_generic_openai,
"LiteLLM": self.async_step_litellm,
}[self.provider]()
return self.async_show_form(
step_id="user",
data_schema=vol.Schema(
{
vol.Required(CONF_PROVIDER): vol.In(
[
"Anthropic",
"Custom OpenAI",
"Generic OpenAI",
"Google",
"Groq",
"LiteLLM",
"LocalAI",
"Mistral AI",
"Ollama",
"OpenAI Azure",
"OpenAI",
"OpenRouter",
"Perplexity AI",
]
)
}
),
errors=errors,
)
# ───────── helper to reduce repetition ─────────
async def _provider_form(
self,
step_id: str,
schema: vol.Schema,
validate_fn,
title: str,
errors: Dict[str, str],
placeholders: Dict[str, str],
user_input: Dict[str, Any] | None,
):
if user_input:
self.validator = ProviderValidator(self.hass, user_input.get(CONF_REQUEST_TIMEOUT))
err = await validate_fn(user_input)
if err is None:
self.data.update({
**user_input,
CONF_MAX_INPUT_TOKENS: user_input.get(CONF_MAX_INPUT_TOKENS, DEFAULT_MAX_INPUT_TOKENS),
CONF_MAX_OUTPUT_TOKENS: user_input.get(CONF_MAX_OUTPUT_TOKENS, DEFAULT_MAX_OUTPUT_TOKENS),
})
return self.async_create_entry(title=title, data=self.data)
errors["base"] = "api_error"
placeholders["error_message"] = err
return self.async_show_form(step_id=step_id, data_schema=schema, errors=errors, description_placeholders=placeholders)
# ───────── providerspecific steps (OpenAI shown; others similar) ─────────
def _add_token_fields(self, base: Dict[Any, Any]) -> Dict[Any, Any]:
"""Append common tuning fields to the schema."""
base[vol.Optional(CONF_MAX_INPUT_TOKENS, default=DEFAULT_MAX_INPUT_TOKENS)] = vol.All(
vol.Coerce(int), vol.Range(min=100)
)
base[vol.Optional(CONF_MAX_OUTPUT_TOKENS, default=DEFAULT_MAX_OUTPUT_TOKENS)] = vol.All(
vol.Coerce(int), vol.Range(min=100)
)
base[vol.Optional(CONF_CUSTOM_SYSTEM_PROMPT, default="")] = str
base[vol.Optional(CONF_EXCLUDED_DOMAINS, default="")] = str
base[vol.Optional(CONF_EXCLUDED_ENTITIES, default="")] = str
base[vol.Optional(CONF_EXCLUDED_AREAS, default="")] = str
base[vol.Optional(CONF_HISTORY_RETENTION, default=DEFAULT_HISTORY_RETENTION)] = vol.All(
vol.Coerce(int), vol.Range(min=1, max=250)
)
base[vol.Optional(CONF_REQUEST_TIMEOUT, default=DEFAULT_REQUEST_TIMEOUT)] = vol.All(
vol.Coerce(int), vol.Range(min=10, max=1800)
)
return base
async def async_step_openai(self, user_input=None):
schema = {
vol.Required(CONF_OPENAI_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_OPENAI_MODEL, default=DEFAULT_MODELS["OpenAI"]): str,
vol.Optional(CONF_OPENAI_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0)),
vol.Optional(CONF_OPENAI_REASONING_EFFORT, default=DEFAULT_OPENAI_REASONING_EFFORT): vol.In(["minimal", "low", "medium", "high"]),
}
self._add_token_fields(schema)
return await self._provider_form(
"openai",
vol.Schema(schema),
lambda ui: self.validator.validate_openai(ui[CONF_OPENAI_API_KEY]),
"AI Automation Suggester (OpenAI)",
{},
{},
user_input,
)
async def async_step_anthropic(self, user_input=None):
async def _v(ui):
return await self.validator.validate_anthropic(
ui[CONF_ANTHROPIC_API_KEY], ui.get(CONF_ANTHROPIC_MODEL, DEFAULT_MODELS["Anthropic"])
)
schema = {
vol.Required(CONF_ANTHROPIC_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_ANTHROPIC_MODEL, default=DEFAULT_MODELS["Anthropic"]): str,
vol.Optional(CONF_ANTHROPIC_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0)),
}
self._add_token_fields(schema)
return await self._provider_form(
"anthropic",
vol.Schema(schema),
_v,
"AI Automation Suggester (Anthropic)",
{},
{},
user_input,
)
async def async_step_google(self, user_input=None):
async def _v(ui):
return await self.validator.validate_google(
ui[CONF_GOOGLE_API_KEY], ui.get(CONF_GOOGLE_MODEL, DEFAULT_MODELS["Google"])
)
schema = {
vol.Required(CONF_GOOGLE_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_GOOGLE_MODEL, default=DEFAULT_MODELS["Google"]): str,
vol.Optional(CONF_GOOGLE_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0)),
}
self._add_token_fields(schema)
return await self._provider_form(
"google",
vol.Schema(schema),
_v,
"AI Automation Suggester (Google)",
{},
{},
user_input,
)
async def async_step_groq(self, user_input=None):
schema = {
vol.Required(CONF_GROQ_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_GROQ_MODEL, default=DEFAULT_MODELS["Groq"]): str,
vol.Optional(CONF_GROQ_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(
vol.Coerce(float), vol.Range(min=0.0, max=2.0)
),
}
self._add_token_fields(schema)
return await self._provider_form(
"groq",
vol.Schema(schema),
lambda ui: self.validator.validate_groq(ui[CONF_GROQ_API_KEY]),
"AI Automation Suggester (Groq)",
{},
{},
user_input,
)
async def async_step_localai(self, user_input=None):
async def _v(ui):
return await self.validator.validate_localai(ui[CONF_LOCALAI_IP_ADDRESS], ui[CONF_LOCALAI_PORT], ui[CONF_LOCALAI_HTTPS])
schema = {
vol.Required(CONF_LOCALAI_IP_ADDRESS): str,
vol.Required(CONF_LOCALAI_PORT, default=8080): int,
vol.Required(CONF_LOCALAI_HTTPS, default=False): bool,
vol.Optional(CONF_LOCALAI_MODEL, default=DEFAULT_MODELS["LocalAI"]): str,
vol.Optional(CONF_LOCALAI_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(
vol.Coerce(float), vol.Range(min=0.0, max=2.0)
),
}
self._add_token_fields(schema)
return await self._provider_form(
"localai",
vol.Schema(schema),
_v,
"AI Automation Suggester (LocalAI)",
{},
{},
user_input,
)
async def async_step_ollama(self, user_input=None):
async def _v(ui):
return await self.validator.validate_ollama(
ui.get(CONF_OLLAMA_IP_ADDRESS),
ui.get(CONF_OLLAMA_PORT),
ui.get(CONF_OLLAMA_HTTPS, False),
ui.get(CONF_OLLAMA_BASE_URL),
ui.get(CONF_OLLAMA_API_KEY),
)
schema = {
vol.Optional(CONF_OLLAMA_BASE_URL, default=""): str,
vol.Optional(CONF_OLLAMA_API_KEY, default=""): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_OLLAMA_IP_ADDRESS, default="localhost"): str,
vol.Optional(CONF_OLLAMA_PORT, default=11434): int,
vol.Optional(CONF_OLLAMA_HTTPS, default=False): bool,
vol.Optional(CONF_OLLAMA_MODEL, default=DEFAULT_MODELS["Ollama"]): str,
vol.Optional(CONF_OLLAMA_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(
vol.Coerce(float), vol.Range(min=0.0, max=2.0)
),
vol.Optional(CONF_OLLAMA_DISABLE_THINK, default=False): bool,
}
self._add_token_fields(schema)
return await self._provider_form(
"ollama",
vol.Schema(schema),
_v,
"AI Automation Suggester (Ollama)",
{},
{},
user_input,
)
async def async_step_custom_openai(self, user_input=None):
async def _v(ui):
return await self.validator.validate_custom_openai(ui[CONF_CUSTOM_OPENAI_ENDPOINT], ui.get(CONF_CUSTOM_OPENAI_API_KEY))
schema = {
vol.Required(CONF_CUSTOM_OPENAI_ENDPOINT): str,
vol.Optional(CONF_CUSTOM_OPENAI_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_CUSTOM_OPENAI_MODEL, default=DEFAULT_MODELS["Custom OpenAI"]): str,
vol.Optional(CONF_CUSTOM_OPENAI_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0)),
}
self._add_token_fields(schema)
return await self._provider_form(
"custom_openai",
vol.Schema(schema),
_v,
"AI Automation Suggester (Custom OpenAI)",
{},
{},
user_input,
)
# Mistral: no live validation needed
async def async_step_mistral(self, user_input=None):
if user_input:
self.data.update(user_input)
return self.async_create_entry(title="AI Automation Suggester (Mistral AI)", data=self.data)
schema = {
vol.Required(CONF_MISTRAL_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_MISTRAL_MODEL, default=DEFAULT_MODELS["Mistral AI"]): str,
vol.Optional(CONF_MISTRAL_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(
vol.Coerce(float), vol.Range(min=0.0, max=2.0)
),
}
self._add_token_fields(schema)
return self.async_show_form(step_id="mistral", data_schema=vol.Schema(schema))
async def async_step_perplexity(self, user_input=None):
async def _v(ui):
return await self.validator.validate_perplexity(
ui[CONF_PERPLEXITY_API_KEY], ui.get(CONF_PERPLEXITY_MODEL, DEFAULT_MODELS["Perplexity AI"])
)
schema = {
vol.Required(CONF_PERPLEXITY_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_PERPLEXITY_MODEL, default=DEFAULT_MODELS["Perplexity AI"]): str,
vol.Optional(CONF_PERPLEXITY_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(
vol.Coerce(float), vol.Range(min=0.0, max=2.0)
),
}
self._add_token_fields(schema)
return await self._provider_form(
"perplexity",
vol.Schema(schema),
_v,
"AI Automation Suggester (Perplexity)",
{},
{},
user_input,
)
async def async_step_openrouter(self, user_input=None):
async def _v(ui):
return await self.validator.validate_openrouter(
ui[CONF_OPENROUTER_API_KEY],
ui.get(CONF_OPENROUTER_MODEL, DEFAULT_MODELS["OpenRouter"]),
)
schema = {
vol.Required(CONF_OPENROUTER_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(
CONF_OPENROUTER_MODEL, default=DEFAULT_MODELS["OpenRouter"]
): str,
vol.Optional(CONF_OPENROUTER_REASONING_MAX_TOKENS, default=0): vol.All(
vol.Coerce(int), vol.Range(min=0)
),
vol.Optional(
CONF_OPENROUTER_TEMPERATURE, default=DEFAULT_TEMPERATURE
): vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0)),
}
self._add_token_fields(schema)
return await self._provider_form(
"openrouter",
vol.Schema(schema),
_v,
"AI Automation Suggester (OpenRouter)",
{},
{},
user_input,
)
async def async_step_openai_azure(self, user_input=None):
async def _v(ui):
if not ui.get(CONF_OPENAI_AZURE_API_KEY) or not ui.get(CONF_OPENAI_AZURE_DEPLOYMENT_ID) or not ui.get(CONF_OPENAI_AZURE_API_VERSION):
return "All fields are required"
return None
schema = {
vol.Required(CONF_OPENAI_AZURE_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Required(CONF_OPENAI_AZURE_DEPLOYMENT_ID, default=DEFAULT_MODELS["OpenAI Azure"]): str,
vol.Required(CONF_OPENAI_AZURE_ENDPOINT, default="{your-resource-name}.openai.azure.com"): str,
vol.Required(CONF_OPENAI_AZURE_API_VERSION, default="2025-01-01-preview"): str,
vol.Optional(CONF_OPENAI_AZURE_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0)),
}
self._add_token_fields(schema)
return await self._provider_form(
"openai_azure",
vol.Schema(schema),
_v,
"AI Automation Suggester (OpenAI Azure)",
{},
{},
user_input,
)
async def async_step_generic_openai(self, user_input=None):
"""Handle the Generic OpenAI API configuration."""
async def _v(ui):
if not ui.get(CONF_GENERIC_OPENAI_ENDPOINT):
return "API URL is required"
if ui.get(CONF_GENERIC_OPENAI_ENABLE_VALIDATION, False):
if not ui.get(CONF_GENERIC_OPENAI_VALIDATION_ENDPOINT):
return "Validation endpoint is required when validation is enabled"
return await self.validator.validate_generic_openai(ui[CONF_GENERIC_OPENAI_VALIDATION_ENDPOINT], ui.get(CONF_GENERIC_OPENAI_API_KEY))
schema = {
vol.Required(CONF_GENERIC_OPENAI_ENDPOINT): str,
vol.Required(CONF_GENERIC_OPENAI_API_KEY): TextSelector(TextSelectorConfig(type="password")),
vol.Required(CONF_GENERIC_OPENAI_MODEL, default=DEFAULT_MODELS["Generic OpenAI"]): str,
vol.Optional(CONF_GENERIC_OPENAI_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0)),
vol.Optional(CONF_GENERIC_OPENAI_VALIDATION_ENDPOINT, default=""): str,
vol.Optional(CONF_GENERIC_OPENAI_ENABLE_VALIDATION, default=False): bool,
}
self._add_token_fields(schema)
return await self._provider_form(
"generic_openai",
vol.Schema(schema),
_v,
"AI Automation Suggester (Generic OpenAI)",
{},
{},
user_input,
)
async def async_step_litellm(self, user_input=None):
"""Handle the LiteLLM configuration."""
async def _v(ui):
if not ui.get(CONF_LITELLM_MODEL):
return "Model is required (e.g. openai/gpt-4o, anthropic/claude-sonnet-4-6)"
schema = {
vol.Required(CONF_LITELLM_MODEL, default=DEFAULT_MODELS["LiteLLM"]): str,
vol.Optional(CONF_LITELLM_API_KEY, default=""): TextSelector(TextSelectorConfig(type="password")),
vol.Optional(CONF_LITELLM_API_BASE, default=""): str,
vol.Optional(CONF_LITELLM_TEMPERATURE, default=DEFAULT_TEMPERATURE): vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0)),
}
self._add_token_fields(schema)
return await self._provider_form(
"litellm",
vol.Schema(schema),
_v,
"AI Automation Suggester (LiteLLM)",
{},
{},
user_input,
)
# ───────── Options flow (edit after setup) ─────────
@staticmethod
@callback
def async_get_options_flow(config_entry):
return AIAutomationOptionsFlowHandler(config_entry)
class AIAutomationOptionsFlowHandler(config_entries.OptionsFlow):
"""Allow postsetup tweaking of models, keys, token budgets."""
def __init__(self, config_entry):
super().__init__()
self._config_entry = config_entry
def _get_option(self, key, default=None):
"""Get value from options, then data, then default."""
if key in self._config_entry.options:
return self._config_entry.options.get(key)
if key in self._config_entry.data:
return self._config_entry.data.get(key)
return default
async def async_step_init(self, user_input=None):
if user_input:
new_data = {
**self._config_entry.options,
**user_input,
CONF_MAX_INPUT_TOKENS: user_input.get(
CONF_MAX_INPUT_TOKENS,
self._get_option(CONF_MAX_INPUT_TOKENS, DEFAULT_MAX_INPUT_TOKENS)
),
CONF_MAX_OUTPUT_TOKENS: user_input.get(
CONF_MAX_OUTPUT_TOKENS,
self._get_option(CONF_MAX_OUTPUT_TOKENS, DEFAULT_MAX_OUTPUT_TOKENS)
),
}
return self.async_create_entry(title="", data=new_data)
provider = self._config_entry.data.get(CONF_PROVIDER)
schema: Dict[Any, Any] = {
vol.Optional(CONF_MAX_INPUT_TOKENS, default=self._get_option(CONF_MAX_INPUT_TOKENS, DEFAULT_MAX_INPUT_TOKENS)
): vol.All(vol.Coerce(int), vol.Range(min=100)),
vol.Optional(CONF_MAX_OUTPUT_TOKENS, default=self._get_option(CONF_MAX_OUTPUT_TOKENS, DEFAULT_MAX_OUTPUT_TOKENS)
): vol.All(vol.Coerce(int), vol.Range(min=100)),
vol.Optional(CONF_CUSTOM_SYSTEM_PROMPT, default=self._get_option(CONF_CUSTOM_SYSTEM_PROMPT, "")): str,
vol.Optional(CONF_EXCLUDED_DOMAINS, default=self._get_option(CONF_EXCLUDED_DOMAINS, "")): str,
vol.Optional(CONF_EXCLUDED_ENTITIES, default=self._get_option(CONF_EXCLUDED_ENTITIES, "")): str,
vol.Optional(CONF_EXCLUDED_AREAS, default=self._get_option(CONF_EXCLUDED_AREAS, "")): str,
vol.Optional(CONF_HISTORY_RETENTION, default=self._get_option(CONF_HISTORY_RETENTION, DEFAULT_HISTORY_RETENTION)): vol.All(vol.Coerce(int), vol.Range(min=1, max=250)),
vol.Optional(CONF_REQUEST_TIMEOUT, default=self._get_option(CONF_REQUEST_TIMEOUT, DEFAULT_REQUEST_TIMEOUT)): vol.All(vol.Coerce(int), vol.Range(min=10, max=1800)),
}
# providerspecific editable fields
if provider == "OpenAI":
schema[vol.Optional(CONF_OPENAI_API_KEY, default=self._get_option(CONF_OPENAI_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_OPENAI_MODEL, default=self._get_option(CONF_OPENAI_MODEL, DEFAULT_MODELS["OpenAI"]))] = str
schema[vol.Optional(CONF_OPENAI_TEMPERATURE, default=self._get_option(CONF_OPENAI_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
schema[vol.Optional(CONF_OPENAI_REASONING_EFFORT, default=self._get_option(CONF_OPENAI_REASONING_EFFORT, DEFAULT_OPENAI_REASONING_EFFORT))] = vol.In(["minimal", "low", "medium", "high"])
elif provider == "Anthropic":
schema[vol.Optional(CONF_ANTHROPIC_API_KEY, default=self._get_option(CONF_ANTHROPIC_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_ANTHROPIC_MODEL, default=self._get_option(CONF_ANTHROPIC_MODEL, DEFAULT_MODELS["Anthropic"]))] = str
schema[vol.Optional(CONF_ANTHROPIC_TEMPERATURE, default=self._get_option(CONF_ANTHROPIC_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
elif provider == "Google":
schema[vol.Optional(CONF_GOOGLE_API_KEY, default=self._get_option(CONF_GOOGLE_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_GOOGLE_MODEL, default=self._get_option(CONF_GOOGLE_MODEL, DEFAULT_MODELS["Google"]))] = str
schema[vol.Optional(CONF_GOOGLE_TEMPERATURE, default=self._get_option(CONF_GOOGLE_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
elif provider == "Groq":
schema[vol.Optional(CONF_GROQ_API_KEY, default=self._get_option(CONF_GROQ_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_GROQ_MODEL, default=self._get_option(CONF_GROQ_MODEL, DEFAULT_MODELS["Groq"]))] = str
schema[vol.Optional(CONF_GROQ_TEMPERATURE, default=self._get_option(CONF_GROQ_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
elif provider == "LocalAI":
schema[vol.Optional(CONF_LOCALAI_HTTPS, default=self._get_option(CONF_LOCALAI_HTTPS, False))] = bool
schema[vol.Optional(CONF_LOCALAI_MODEL, default=self._get_option(CONF_LOCALAI_MODEL, DEFAULT_MODELS["LocalAI"]))] = str
schema[vol.Optional(CONF_LOCALAI_TEMPERATURE, default=self._get_option(CONF_LOCALAI_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
schema[vol.Optional(CONF_LOCALAI_IP_ADDRESS, default=self._get_option(CONF_LOCALAI_IP_ADDRESS, "localhost"))] = str
schema[vol.Optional(CONF_LOCALAI_PORT, default=self._get_option(CONF_LOCALAI_PORT, 8080))] = int
elif provider == "Ollama":
schema[vol.Optional(CONF_OLLAMA_BASE_URL, default=self._get_option(CONF_OLLAMA_BASE_URL, ""))] = str
schema[vol.Optional(CONF_OLLAMA_API_KEY, default=self._get_option(CONF_OLLAMA_API_KEY, ""))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_OLLAMA_IP_ADDRESS, default=self._get_option(CONF_OLLAMA_IP_ADDRESS, "localhost"))] = str
schema[vol.Optional(CONF_OLLAMA_PORT, default=self._get_option(CONF_OLLAMA_PORT, 11434))] = int
schema[vol.Optional(CONF_OLLAMA_HTTPS, default=self._get_option(CONF_OLLAMA_HTTPS, False))] = bool
schema[vol.Optional(CONF_OLLAMA_MODEL, default=self._get_option(CONF_OLLAMA_MODEL, DEFAULT_MODELS["Ollama"]))] = str
schema[vol.Optional(CONF_OLLAMA_TEMPERATURE, default=self._get_option(CONF_OLLAMA_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
schema[vol.Optional(CONF_OLLAMA_DISABLE_THINK, default=self._get_option(CONF_OLLAMA_DISABLE_THINK, False))] = bool
elif provider == "Custom OpenAI":
schema[vol.Optional(CONF_CUSTOM_OPENAI_ENDPOINT, default=self._get_option(CONF_CUSTOM_OPENAI_ENDPOINT))] = str
schema[vol.Optional(CONF_CUSTOM_OPENAI_API_KEY, default=self._get_option(CONF_CUSTOM_OPENAI_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_CUSTOM_OPENAI_MODEL, default=self._get_option(CONF_CUSTOM_OPENAI_MODEL, DEFAULT_MODELS["Custom OpenAI"]))] = str
schema[vol.Optional(CONF_CUSTOM_OPENAI_TEMPERATURE, default=self._get_option(CONF_CUSTOM_OPENAI_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
elif provider == "Mistral AI":
schema[vol.Optional(CONF_MISTRAL_API_KEY, default=self._get_option(CONF_MISTRAL_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_MISTRAL_MODEL, default=self._get_option(CONF_MISTRAL_MODEL, DEFAULT_MODELS["Mistral AI"]))] = str
schema[vol.Optional(CONF_MISTRAL_TEMPERATURE, default=self._get_option(CONF_MISTRAL_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
elif provider == "Perplexity AI":
schema[vol.Optional(CONF_PERPLEXITY_API_KEY, default=self._get_option(CONF_PERPLEXITY_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_PERPLEXITY_MODEL, default=self._get_option(CONF_PERPLEXITY_MODEL, DEFAULT_MODELS["Perplexity AI"]))] = str
schema[vol.Optional(CONF_PERPLEXITY_TEMPERATURE, default=self._get_option(CONF_PERPLEXITY_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
elif provider == "OpenRouter":
schema[vol.Optional(CONF_OPENROUTER_API_KEY, default=self._get_option(CONF_OPENROUTER_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_OPENROUTER_MODEL, default=self._get_option(CONF_OPENROUTER_MODEL, DEFAULT_MODELS["OpenRouter"]))] = str
schema[vol.Optional(CONF_OPENROUTER_REASONING_MAX_TOKENS, default=self._get_option(CONF_OPENROUTER_REASONING_MAX_TOKENS, 0))] = vol.All(vol.Coerce(int), vol.Range(min=0))
schema[vol.Optional(CONF_OPENROUTER_TEMPERATURE, default=self._get_option(CONF_OPENROUTER_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
elif provider == "OpenAI Azure":
schema[vol.Optional(CONF_OPENAI_AZURE_API_KEY, default=self._get_option(CONF_OPENAI_AZURE_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_OPENAI_AZURE_ENDPOINT, default=self._get_option(CONF_OPENAI_AZURE_ENDPOINT))] = str
schema[vol.Optional(CONF_OPENAI_AZURE_DEPLOYMENT_ID, default=self._get_option(CONF_OPENAI_AZURE_DEPLOYMENT_ID, DEFAULT_MODELS["OpenAI Azure"]))] = str
schema[vol.Optional(CONF_OPENAI_AZURE_API_VERSION, default=self._get_option(CONF_OPENAI_AZURE_API_VERSION, "2025-01-01-preview"))] = str
schema[vol.Optional(CONF_OPENAI_AZURE_TEMPERATURE, default=self._get_option(CONF_OPENAI_AZURE_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
elif provider == "Generic OpenAI":
schema[vol.Optional(CONF_GENERIC_OPENAI_API_KEY, default=self._get_option(CONF_GENERIC_OPENAI_API_KEY))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_GENERIC_OPENAI_ENDPOINT, default=self._get_option(CONF_GENERIC_OPENAI_ENDPOINT))] = str
schema[vol.Optional(CONF_GENERIC_OPENAI_MODEL, default=self._get_option(CONF_GENERIC_OPENAI_MODEL, DEFAULT_MODELS["Generic OpenAI"]))] = str
schema[vol.Optional(CONF_GENERIC_OPENAI_TEMPERATURE, default=self._get_option(CONF_GENERIC_OPENAI_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
schema[vol.Optional(CONF_GENERIC_OPENAI_VALIDATION_ENDPOINT, default=self._get_option(CONF_GENERIC_OPENAI_VALIDATION_ENDPOINT, ""))] = str
schema[vol.Optional(CONF_GENERIC_OPENAI_ENABLE_VALIDATION, default=self._get_option(CONF_GENERIC_OPENAI_ENABLE_VALIDATION, False))] = bool
elif provider == "LiteLLM":
schema[vol.Optional(CONF_LITELLM_API_KEY, default=self._get_option(CONF_LITELLM_API_KEY, ""))] = TextSelector(TextSelectorConfig(type="password"))
schema[vol.Optional(CONF_LITELLM_MODEL, default=self._get_option(CONF_LITELLM_MODEL, DEFAULT_MODELS["LiteLLM"]))] = str
schema[vol.Optional(CONF_LITELLM_API_BASE, default=self._get_option(CONF_LITELLM_API_BASE, ""))] = str
schema[vol.Optional(CONF_LITELLM_TEMPERATURE, default=self._get_option(CONF_LITELLM_TEMPERATURE, DEFAULT_TEMPERATURE))] = vol.All(vol.Coerce(float), vol.Range(min=0.0, max=2.0))
return self.async_show_form(step_id="init", data_schema=vol.Schema(schema))