Update config_flow.py

This commit is contained in:
smkrv
2024-11-14 23:23:01 +03:00
committed by GitHub
parent 9a8c676570
commit 0789f8c386
+108 -100
View File
@@ -1,123 +1,131 @@
"""Config flow for HA Text AI Integration.""" """Config flow for HA Text AI Integration."""
from __future__ import annotations from __future__ import annotations
import logging import logging
from typing import Any from typing import Any
import voluptuous as vol import voluptuous as vol
import openai from openai import AsyncOpenAI
from openai.types.error import APIError
from openai.exceptions import AuthenticationError, APIConnectionError
from homeassistant import config_entries from homeassistant import config_entries
from homeassistant.core import HomeAssistant, callback from homeassistant.core import HomeAssistant, callback
from homeassistant.data_entry_flow import FlowResult from homeassistant.data_entry_flow import FlowResult
from homeassistant.exceptions import HomeAssistantError from homeassistant.exceptions import HomeAssistantError
from .const import ( from .const import (
DOMAIN, DOMAIN,
CONF_API_KEY, CONF_API_KEY,
CONF_API_BASE, CONF_API_BASE,
CONF_REQUEST_INTERVAL, CONF_REQUEST_INTERVAL,
DEFAULT_API_BASE, DEFAULT_API_BASE,
DEFAULT_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL,
) )
_LOGGER = logging.getLogger(__name__) _LOGGER = logging.getLogger(__name__)
class TextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): class TextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
"""Handle a config flow for HA Text AI Integration.""" """Handle a config flow for HA Text AI Integration."""
VERSION = 1 VERSION = 1
async def async_step_user( async def async_step_user(
self, user_input: dict[str, Any] | None = None self, user_input: dict[str, Any] | None = None
) -> FlowResult: ) -> FlowResult:
"""Handle the initial step.""" """Handle the initial step."""
errors = {} errors = {}
if user_input is not None: if user_input is not None:
try: try:
await self._test_api_key(user_input[CONF_API_KEY], user_input.get(CONF_API_BASE)) await self._test_api_key(user_input[CONF_API_KEY], user_input.get(CONF_API_BASE))
return self.async_create_entry( return self.async_create_entry(
title="HA Text AI", title="HA Text AI",
data=user_input, data=user_input,
) )
except ApiKeyError: except ApiKeyError:
errors["base"] = "invalid_api_key" errors["base"] = "invalid_api_key"
except ApiConnectionError: except ApiConnectionError:
errors["base"] = "cannot_connect" errors["base"] = "cannot_connect"
except Exception: # pylint: disable=broad-except except Exception: # pylint: disable=broad-except
_LOGGER.exception("Unexpected exception") _LOGGER.exception("Unexpected exception")
errors["base"] = "unknown" errors["base"] = "unknown"
return self.async_show_form( return self.async_show_form(
step_id="user", step_id="user",
data_schema=vol.Schema( data_schema=vol.Schema(
{ {
vol.Required(CONF_API_KEY): str, vol.Required(CONF_API_KEY): str,
vol.Optional(CONF_API_BASE, default=DEFAULT_API_BASE): str, vol.Optional(CONF_API_BASE, default=DEFAULT_API_BASE): str,
vol.Optional(CONF_REQUEST_INTERVAL, default=DEFAULT_REQUEST_INTERVAL): int, vol.Optional(CONF_REQUEST_INTERVAL, default=DEFAULT_REQUEST_INTERVAL): int,
} }
), ),
errors=errors, errors=errors,
) )
@staticmethod @staticmethod
async def _test_api_key(api_key: str, api_base: str | None) -> None: async def _test_api_key(api_key: str, api_base: str | None) -> None:
"""Test if the API key is valid.""" """Test if the API key is valid."""
try: try:
openai.api_key = api_key client = AsyncOpenAI(
if api_base: api_key=api_key,
openai.api_base = api_base base_url=api_base if api_base else DEFAULT_API_BASE
)
models = await openai.Model.alist() models = await client.models.list()
if not models: if not models.data:
raise ApiKeyError raise ApiKeyError
except openai.error.AuthenticationError as err:
raise ApiKeyError from err except AuthenticationError as err:
except openai.error.APIConnectionError as err: raise ApiKeyError from err
raise ApiConnectionError from err except APIConnectionError as err:
raise ApiConnectionError from err
except APIError as err:
if err.status_code == 401:
raise ApiKeyError from err
raise ApiConnectionError from err
@staticmethod @staticmethod
@callback @callback
def async_get_options_flow( def async_get_options_flow(
config_entry: config_entries.ConfigEntry, config_entry: config_entries.ConfigEntry,
) -> TextAIOptionsFlow: ) -> TextAIOptionsFlow:
"""Get the options flow for this handler.""" """Get the options flow for this handler."""
return TextAIOptionsFlow(config_entry) return TextAIOptionsFlow(config_entry)
class TextAIOptionsFlow(config_entries.OptionsFlow): class TextAIOptionsFlow(config_entries.OptionsFlow):
"""Handle options flow for HA Text AI Integration.""" """Handle options flow for HA Text AI Integration."""
def __init__(self, config_entry: config_entries.ConfigEntry) -> None: def __init__(self, config_entry: config_entries.ConfigEntry) -> None:
"""Initialize options flow.""" """Initialize options flow."""
self.config_entry = config_entry self.config_entry = config_entry
async def async_step_init( async def async_step_init(
self, user_input: dict[str, Any] | None = None self, user_input: dict[str, Any] | None = None
) -> FlowResult: ) -> FlowResult:
"""Manage options.""" """Manage options."""
if user_input is not None: if user_input is not None:
return self.async_create_entry(title="", data=user_input) return self.async_create_entry(title="", data=user_input)
return self.async_show_form( return self.async_show_form(
step_id="init", step_id="init",
data_schema=vol.Schema( data_schema=vol.Schema(
{ {
vol.Optional( vol.Optional(
CONF_REQUEST_INTERVAL, CONF_REQUEST_INTERVAL,
default=self.config_entry.options.get( default=self.config_entry.options.get(
CONF_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL CONF_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL
), ),
): int, ): int,
} }
), ),
) )
class ApiKeyError(HomeAssistantError): class ApiKeyError(HomeAssistantError):
"""Error to indicate there is an invalid API key.""" """Error to indicate there is an invalid API key."""
class ApiConnectionError(HomeAssistantError): class ApiConnectionError(HomeAssistantError):
"""Error to indicate there is a connection error.""" """Error to indicate there is a connection error."""