mirror of
https://github.com/smkrv/ha-text-ai.git
synced 2026-07-21 22:54:00 +08:00
Update __init__.py
This commit is contained in:
@@ -5,7 +5,7 @@ import asyncio
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import openai
|
||||
from openai import AsyncOpenAI
|
||||
import voluptuous as vol
|
||||
|
||||
from homeassistant.config_entries import ConfigEntry
|
||||
@@ -37,9 +37,13 @@ async def async_setup(hass: HomeAssistant, config: dict) -> bool:
|
||||
|
||||
async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
|
||||
"""Set up HA Text AI from a config entry."""
|
||||
client = AsyncOpenAI(
|
||||
api_key=entry.data[CONF_API_KEY],
|
||||
base_url=entry.data.get(CONF_API_BASE, DEFAULT_API_BASE)
|
||||
)
|
||||
|
||||
hass.data[DOMAIN][entry.entry_id] = {
|
||||
CONF_API_KEY: entry.data[CONF_API_KEY],
|
||||
CONF_API_BASE: entry.data.get(CONF_API_BASE, DEFAULT_API_BASE),
|
||||
"client": client,
|
||||
CONF_REQUEST_INTERVAL: entry.data.get(CONF_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL),
|
||||
"queue": [],
|
||||
"processing": False,
|
||||
@@ -95,16 +99,13 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
|
||||
return
|
||||
|
||||
hass.data[DOMAIN][entry_id]["processing"] = True
|
||||
client = hass.data[DOMAIN][entry_id]["client"]
|
||||
|
||||
while hass.data[DOMAIN][entry_id]["queue"]:
|
||||
request = hass.data[DOMAIN][entry_id]["queue"].pop(0)
|
||||
|
||||
try:
|
||||
openai.api_key = hass.data[DOMAIN][entry_id][CONF_API_KEY]
|
||||
openai.api_base = hass.data[DOMAIN][entry_id][CONF_API_BASE]
|
||||
|
||||
response = await hass.async_add_executor_job(
|
||||
lambda: openai.ChatCompletion.create(
|
||||
response = await client.chat.completions.create(
|
||||
model=request["model"],
|
||||
messages=[{"role": "user", "content": request["prompt"]}],
|
||||
temperature=request["temperature"],
|
||||
@@ -113,7 +114,6 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
|
||||
frequency_penalty=request["frequency_penalty"],
|
||||
presence_penalty=request["presence_penalty"],
|
||||
)
|
||||
)
|
||||
|
||||
response_text = response.choices[0].message.content
|
||||
await hass.services.async_call(
|
||||
|
||||
Reference in New Issue
Block a user