From f3b76c0bc1e7909c3188ef7f7a65108507bec721 Mon Sep 17 00:00:00 2001 From: SMKRV Date: Thu, 21 Nov 2024 19:04:28 +0300 Subject: [PATCH] Release v2.0.0 --- custom_components/ha_text_ai/config_flow.py | 119 ++++++++++++++++---- 1 file changed, 94 insertions(+), 25 deletions(-) diff --git a/custom_components/ha_text_ai/config_flow.py b/custom_components/ha_text_ai/config_flow.py index a94d21a..dfc4622 100644 --- a/custom_components/ha_text_ai/config_flow.py +++ b/custom_components/ha_text_ai/config_flow.py @@ -47,33 +47,102 @@ _LOGGER = logging.getLogger(__name__) STEP_USER_DATA_SCHEMA = vol.Schema({ vol.Required(CONF_API_KEY): str, vol.Optional(CONF_MODEL, default=DEFAULT_MODEL): str, - vol.Optional( - CONF_TEMPERATURE, - default=DEFAULT_TEMPERATURE - ): vol.All( - vol.Coerce(float), - vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE) - ), - vol.Optional( - CONF_MAX_TOKENS, - default=DEFAULT_MAX_TOKENS - ): vol.All( - vol.Coerce(int), - vol.Range(min=MIN_MAX_TOKENS, max=MAX_MAX_TOKENS) - ), - vol.Optional( - CONF_API_ENDPOINT, - default=DEFAULT_API_ENDPOINT - ): str, - vol.Optional( - CONF_REQUEST_INTERVAL, - default=DEFAULT_REQUEST_INTERVAL - ): vol.All( - vol.Coerce(float), - vol.Range(min=MIN_REQUEST_INTERVAL) - ), + vol.Optional(CONF_TEMPERATURE, default=str(DEFAULT_TEMPERATURE)): str, + vol.Optional(CONF_MAX_TOKENS, default=str(DEFAULT_MAX_TOKENS)): str, + vol.Optional(CONF_API_ENDPOINT, default=DEFAULT_API_ENDPOINT): str, + vol.Optional(CONF_REQUEST_INTERVAL, default=str(DEFAULT_REQUEST_INTERVAL)): str, }) +async def async_step_user( + self, + user_input: Optional[Dict[str, Any]] = None +) -> FlowResult: + """Handle the initial step.""" + errors: Dict[str, str] = {} + + if user_input is not None: + try: + _LOGGER.debug("Received user input: %s", user_input) + + # Преобразование строковых значений в числовые + try: + temperature = float(user_input[CONF_TEMPERATURE]) + max_tokens = int(user_input[CONF_MAX_TOKENS]) + request_interval = float(user_input[CONF_REQUEST_INTERVAL]) + + # Обновляем значения в user_input + validated_input = { + CONF_API_KEY: user_input[CONF_API_KEY], + CONF_MODEL: user_input[CONF_MODEL], + CONF_TEMPERATURE: temperature, + CONF_MAX_TOKENS: max_tokens, + CONF_API_ENDPOINT: user_input[CONF_API_ENDPOINT], + CONF_REQUEST_INTERVAL: request_interval + } + + _LOGGER.debug("Converted values: %s", validated_input) + + # Проверка диапазонов + if not MIN_TEMPERATURE <= temperature <= MAX_TEMPERATURE: + _LOGGER.error("Invalid temperature: %f", temperature) + errors["base"] = "invalid_temperature" + elif not MIN_MAX_TOKENS <= max_tokens <= MAX_MAX_TOKENS: + _LOGGER.error("Invalid max_tokens: %d", max_tokens) + errors["base"] = "invalid_max_tokens" + elif request_interval < MIN_REQUEST_INTERVAL: + _LOGGER.error("Invalid request_interval: %f", request_interval) + errors["base"] = "invalid_request_interval" + else: + # Проверка URL + endpoint = validated_input[CONF_API_ENDPOINT] + try: + result = urlparse(endpoint) + if not all([result.scheme, result.netloc]): + errors["base"] = "invalid_url_format" + else: + is_valid, error_code, available_models = await validate_api_connection( + self.hass, + validated_input[CONF_API_KEY], + endpoint, + validated_input[CONF_MODEL] + ) + + if is_valid: + await self.async_set_unique_id(validated_input[CONF_API_KEY]) + self._abort_if_unique_id_configured() + return self.async_create_entry( + title="HA Text AI", + data=validated_input + ) + errors["base"] = error_code + if error_code == ERROR_INVALID_MODEL: + _LOGGER.warning( + "Selected model %s not found in available models: %s", + validated_input[CONF_MODEL], + ", ".join(available_models) + ) + except Exception as e: + _LOGGER.error("URL parsing error: %s", str(e)) + errors["base"] = "invalid_url_format" + + except ValueError as ve: + _LOGGER.error("Value conversion error: %s", str(ve)) + errors["base"] = "invalid_input_format" + + except Exception as err: + _LOGGER.error("Unexpected error: %s", str(err)) + errors["base"] = "unknown" + + return self.async_show_form( + step_id="user", + data_schema=STEP_USER_DATA_SCHEMA, + errors=errors, + description_placeholders={ + "default_model": DEFAULT_MODEL, + "default_endpoint": DEFAULT_API_ENDPOINT, + } + ) + async def validate_api_connection( hass, api_key: str,