mirror of
https://github.com/smkrv/ha-text-ai.git
synced 2026-07-30 23:13:57 +08:00
Update config_flow.py
This commit is contained in:
@@ -5,7 +5,9 @@ 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
|
||||||
@@ -65,16 +67,22 @@ class TextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
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:
|
|
||||||
|
except AuthenticationError as err:
|
||||||
raise ApiKeyError from err
|
raise ApiKeyError from err
|
||||||
except openai.error.APIConnectionError as 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
|
raise ApiConnectionError from err
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
Reference in New Issue
Block a user