This commit is contained in:
SMKRV
2024-11-19 19:39:05 +03:00
parent 675975d951
commit 929d916d41
+17 -10
View File
@@ -7,6 +7,7 @@ from typing import Any, Dict, Optional
from openai import AsyncOpenAI, APIError, AuthenticationError, RateLimitError from openai import AsyncOpenAI, APIError, AuthenticationError, RateLimitError
from homeassistant.core import HomeAssistant from homeassistant.core import HomeAssistant
from homeassistant.helpers.update_coordinator import DataUpdateCoordinator from homeassistant.helpers.update_coordinator import DataUpdateCoordinator
from homeassistant.util import dt as dt_util
import async_timeout import async_timeout
from .const import DOMAIN from .const import DOMAIN
@@ -76,7 +77,7 @@ class HATextAICoordinator(DataUpdateCoordinator):
params = question_data["params"] params = question_data["params"]
try: try:
response_content = await self._make_api_call( response_data = await self._make_api_call(
question, question,
model=params.get("model"), model=params.get("model"),
temperature=params.get("temperature"), temperature=params.get("temperature"),
@@ -86,12 +87,13 @@ class HATextAICoordinator(DataUpdateCoordinator):
self._responses[question] = { self._responses[question] = {
"question": question, "question": question,
"response": response_content, "response": response_data["response"],
"error": None, "error": None,
"timestamp": self.hass.loop.time(), "timestamp": dt_util.utcnow(),
"model": params.get("model", self.model), "model": response_data["model"],
"temperature": params.get("temperature", self.temperature), "temperature": params.get("temperature", self.temperature),
"max_tokens": params.get("max_tokens", self.max_tokens) "max_tokens": params.get("max_tokens", self.max_tokens),
"response_time": response_data.get("response_time")
} }
self._error_count = 0 self._error_count = 0
self._is_ready = True self._is_ready = True
@@ -126,7 +128,7 @@ class HATextAICoordinator(DataUpdateCoordinator):
"question": question, "question": question,
"response": None, "response": None,
"error": error_msg, "error": error_msg,
"timestamp": self.hass.loop.time(), "timestamp": dt_util.utcnow(),
"model": self.model, "model": self.model,
"temperature": self.temperature, "temperature": self.temperature,
"max_tokens": self.max_tokens "max_tokens": self.max_tokens
@@ -145,7 +147,6 @@ class HATextAICoordinator(DataUpdateCoordinator):
self._error_count += 1 self._error_count += 1
if not self._question_queue.empty(): if not self._question_queue.empty():
try: try:
# Clear the queue if we have timeout issues
while not self._question_queue.empty(): while not self._question_queue.empty():
self._question_queue.get_nowait() self._question_queue.get_nowait()
self._question_queue.task_done() self._question_queue.task_done()
@@ -159,7 +160,7 @@ class HATextAICoordinator(DataUpdateCoordinator):
temperature: Optional[float] = None, temperature: Optional[float] = None,
max_tokens: Optional[int] = None, max_tokens: Optional[int] = None,
system_prompt: Optional[str] = None system_prompt: Optional[str] = None
) -> str: ) -> Dict[str, Any]:
"""Make API call to OpenAI.""" """Make API call to OpenAI."""
try: try:
messages = [] messages = []
@@ -168,13 +169,20 @@ class HATextAICoordinator(DataUpdateCoordinator):
messages.append({"role": "system", "content": current_system_prompt}) messages.append({"role": "system", "content": current_system_prompt})
messages.append({"role": "user", "content": question}) messages.append({"role": "user", "content": question})
start_time = dt_util.utcnow()
completion = await self.client.chat.completions.create( completion = await self.client.chat.completions.create(
model=model or self.model, model=model or self.model,
messages=messages, messages=messages,
temperature=temperature if temperature is not None else self.temperature, temperature=temperature if temperature is not None else self.temperature,
max_tokens=max_tokens if max_tokens is not None else self.max_tokens, max_tokens=max_tokens if max_tokens is not None else self.max_tokens,
) )
return completion.choices[0].message.content response_time = (dt_util.utcnow() - start_time).total_seconds()
return {
"response": completion.choices[0].message.content,
"model": completion.model,
"response_time": response_time
}
except Exception as err: except Exception as err:
_LOGGER.error("Error in API call: %s", err) _LOGGER.error("Error in API call: %s", err)
@@ -209,7 +217,6 @@ class HATextAICoordinator(DataUpdateCoordinator):
async def async_shutdown(self) -> None: async def async_shutdown(self) -> None:
"""Shutdown the coordinator.""" """Shutdown the coordinator."""
try: try:
# Clear the queue
while not self._question_queue.empty(): while not self._question_queue.empty():
self._question_queue.get_nowait() self._question_queue.get_nowait()
self._question_queue.task_done() self._question_queue.task_done()