diff --git a/custom_components/ha_text_ai/config_flow.py b/custom_components/ha_text_ai/config_flow.py index beec4b6..022dfb3 100644 --- a/custom_components/ha_text_ai/config_flow.py +++ b/custom_components/ha_text_ai/config_flow.py @@ -62,24 +62,6 @@ STEP_USER_DATA_SCHEMA = vol.Schema({ ), }) -async def validate_endpoint(endpoint: str) -> Tuple[bool, str]: - """Validate API endpoint accessibility.""" - try: - parsed_url = urlparse(endpoint) - if parsed_url.scheme not in ('http', 'https'): - return False, "invalid_endpoint_scheme" - - connector = aiohttp.TCPConnector(ssl=SSL_CONTEXT) - async with timeout(5): - async with aiohttp.ClientSession(connector=connector) as session: - async with session.get(endpoint) as response: - if response.status != 200: - return False, "endpoint_not_available" - return True, "" - except Exception as e: - _LOGGER.error("Error validating endpoint: %s", str(e)) - return False, "endpoint_error" - async def validate_api_connection( api_key: str, endpoint: str, @@ -87,64 +69,56 @@ async def validate_api_connection( retry_count: int = 3, retry_delay: float = 1.0 ) -> Tuple[bool, str, list]: - """Validate API connection with improved retry logic.""" - # Validate endpoint first - endpoint_valid, endpoint_error = await validate_endpoint(endpoint) - if not endpoint_valid: - return False, endpoint_error, [] - - connector = aiohttp.TCPConnector(ssl=SSL_CONTEXT) - async with aiohttp.ClientSession(connector=connector) as session: - for attempt in range(retry_count): - try: - async with timeout(10): - client = AsyncOpenAI( - api_key=api_key, - base_url=endpoint, - http_client=session - ) - - models = await client.models.list() - model_ids = [model.id for model in models.data] - - if model not in model_ids: - _LOGGER.warning( - "Model %s not found in available models: %s", - model, - ", ".join(model_ids) - ) - return False, "invalid_model", model_ids - return True, "", model_ids - - except asyncio.TimeoutError: - _LOGGER.warning( - "Timeout during API validation (attempt %d/%d)", - attempt + 1, - retry_count + """Validate API connection with retry logic.""" + for attempt in range(retry_count): + try: + async with timeout(10): + client = AsyncOpenAI( + api_key=api_key, + base_url=endpoint, ) - if attempt == retry_count - 1: - return False, "timeout", [] - await asyncio.sleep(retry_delay) - except AuthenticationError as err: - _LOGGER.error("Authentication error: %s", str(err)) - return False, "invalid_auth", [] + models = await client.models.list() + model_ids = [model.id for model in models.data] - except RateLimitError as err: - _LOGGER.error("Rate limit exceeded: %s", str(err)) - return False, "rate_limit", [] + if model not in model_ids: + _LOGGER.warning( + "Model %s not found in available models: %s", + model, + ", ".join(model_ids) + ) + return False, "invalid_model", model_ids + return True, "", model_ids - except APIConnectionError as err: - _LOGGER.error("API connection error: %s", str(err)) - return False, "cannot_connect", [] + except asyncio.TimeoutError: + _LOGGER.warning( + "Timeout during API validation (attempt %d/%d)", + attempt + 1, + retry_count + ) + if attempt == retry_count - 1: + return False, "timeout", [] + await asyncio.sleep(retry_delay) - except APIError as err: - _LOGGER.error("API error: %s", str(err)) - return False, "api_error", [] + except AuthenticationError as err: + _LOGGER.error("Authentication error: %s", str(err)) + return False, "invalid_auth", [] - except Exception as err: - _LOGGER.exception("Unexpected error during validation: %s", str(err)) - return False, "unknown", [] + except RateLimitError as err: + _LOGGER.error("Rate limit exceeded: %s", str(err)) + return False, "rate_limit", [] + + except APIConnectionError as err: + _LOGGER.error("API connection error: %s", str(err)) + return False, "cannot_connect", [] + + except APIError as err: + _LOGGER.error("API error: %s", str(err)) + return False, "api_error", [] + + except Exception as err: + _LOGGER.exception("Unexpected error during validation: %s", str(err)) + return False, "unknown", [] class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): """Handle a config flow for HA text AI."""