Release v2.0.0

This commit is contained in:
SMKRV
2024-11-24 23:43:36 +03:00
parent 28248ac3c4
commit a4925fc943
3 changed files with 100 additions and 76 deletions
-3
View File
@@ -43,8 +43,6 @@ from .const import (
SERVICE_SET_SYSTEM_PROMPT, SERVICE_SET_SYSTEM_PROMPT,
) )
DOMAIN = "ha_text_ai"
_LOGGER = logging.getLogger(__name__) _LOGGER = logging.getLogger(__name__)
CONFIG_SCHEMA = cv.config_entry_only_config_schema(DOMAIN) CONFIG_SCHEMA = cv.config_entry_only_config_schema(DOMAIN)
@@ -228,7 +226,6 @@ async def async_setup_entry(hass: HomeAssistant, entry: ConfigEntry) -> bool:
_LOGGER.debug("Creating API client for %s with endpoint %s", api_provider, endpoint) _LOGGER.debug("Creating API client for %s with endpoint %s", api_provider, endpoint)
# Создаем API клиент
api_client = APIClient( api_client = APIClient(
session=session, session=session,
endpoint=endpoint, endpoint=endpoint,
+89 -37
View File
@@ -1,9 +1,21 @@
"""API Client for HA Text AI.""" """API Client for HA Text AI."""
import logging import logging
import asyncio
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
from aiohttp import ClientSession, ClientTimeout
from async_timeout import timeout
from homeassistant.core import HomeAssistant from homeassistant.core import HomeAssistant
from homeassistant.exceptions import HomeAssistantError from homeassistant.exceptions import HomeAssistantError
from .const import (
API_TIMEOUT,
API_RETRY_COUNT,
API_PROVIDER_ANTHROPIC,
MIN_TEMPERATURE,
MAX_TEMPERATURE,
MIN_MAX_TOKENS,
MAX_MAX_TOKENS,
)
_LOGGER = logging.getLogger(__name__) _LOGGER = logging.getLogger(__name__)
@@ -12,7 +24,7 @@ class APIClient:
def __init__( def __init__(
self, self,
session: Any, session: ClientSession,
endpoint: str, endpoint: str,
headers: Dict[str, str], headers: Dict[str, str],
api_provider: str, api_provider: str,
@@ -24,6 +36,51 @@ class APIClient:
self.headers = headers self.headers = headers
self.api_provider = api_provider self.api_provider = api_provider
self.model = model self.model = model
self.timeout = ClientTimeout(total=API_TIMEOUT)
def _validate_parameters(
self,
temperature: float,
max_tokens: int,
) -> None:
"""Validate API parameters."""
if not MIN_TEMPERATURE <= temperature <= MAX_TEMPERATURE:
raise ValueError(
f"Temperature must be between {MIN_TEMPERATURE} and {MAX_TEMPERATURE}"
)
if not MIN_MAX_TOKENS <= max_tokens <= MAX_MAX_TOKENS:
raise ValueError(
f"Max tokens must be between {MIN_MAX_TOKENS} and {MAX_MAX_TOKENS}"
)
async def _make_request(
self,
url: str,
payload: Dict[str, Any],
) -> Dict[str, Any]:
"""Make API request with retry logic."""
for attempt in range(API_RETRY_COUNT):
try:
async with timeout(API_TIMEOUT):
async with self.session.post(
url,
json=payload,
headers=self.headers,
timeout=self.timeout
) as response:
if response.status != 200:
error_data = await response.json()
raise HomeAssistantError(f"API error: {error_data}")
return await response.json()
except asyncio.TimeoutError:
if attempt == API_RETRY_COUNT - 1:
raise HomeAssistantError("API request timed out")
await asyncio.sleep(1 * (attempt + 1))
except Exception as e:
if attempt == API_RETRY_COUNT - 1:
raise
_LOGGER.warning("API request failed, retrying: %s", str(e))
await asyncio.sleep(1 * (attempt + 1))
async def create( async def create(
self, self,
@@ -34,7 +91,9 @@ class APIClient:
) -> Dict[str, Any]: ) -> Dict[str, Any]:
"""Create completion using appropriate API.""" """Create completion using appropriate API."""
try: try:
if self.api_provider == "anthropic": self._validate_parameters(temperature, max_tokens)
if self.api_provider == API_PROVIDER_ANTHROPIC:
return await self._create_anthropic_completion( return await self._create_anthropic_completion(
model, messages, temperature, max_tokens model, messages, temperature, max_tokens
) )
@@ -62,26 +121,21 @@ class APIClient:
"max_tokens": max_tokens, "max_tokens": max_tokens,
} }
async with self.session.post(url, json=payload, headers=self.headers) as response: data = await self._make_request(url, payload)
if response.status != 200: return {
error_data = await response.json() "choices": [
raise HomeAssistantError(f"OpenAI API error: {error_data}") {
"message": {
data = await response.json() "content": data["choices"][0]["message"]["content"]
return {
"choices": [
{
"message": {
"content": data["choices"][0]["message"]["content"]
}
} }
],
"usage": {
"prompt_tokens": data["usage"]["prompt_tokens"],
"completion_tokens": data["usage"]["completion_tokens"],
"total_tokens": data["usage"]["total_tokens"]
} }
],
"usage": {
"prompt_tokens": data["usage"]["prompt_tokens"],
"completion_tokens": data["usage"]["completion_tokens"],
"total_tokens": data["usage"]["total_tokens"]
} }
}
async def _create_anthropic_completion( async def _create_anthropic_completion(
self, self,
@@ -94,7 +148,10 @@ class APIClient:
url = f"{self.endpoint}/v1/messages" url = f"{self.endpoint}/v1/messages"
# Convert messages to Anthropic format # Convert messages to Anthropic format
system_prompt = next((msg["content"] for msg in messages if msg["role"] == "system"), None) system_prompt = next(
(msg["content"] for msg in messages if msg["role"] == "system"),
None
)
conversation = [msg for msg in messages if msg["role"] != "system"] conversation = [msg for msg in messages if msg["role"] != "system"]
payload = { payload = {
@@ -106,23 +163,18 @@ class APIClient:
if system_prompt: if system_prompt:
payload["system"] = system_prompt payload["system"] = system_prompt
async with self.session.post(url, json=payload, headers=self.headers) as response: data = await self._make_request(url, payload)
if response.status != 200: return {
error_data = await response.json() "choices": [
raise HomeAssistantError(f"Anthropic API error: {error_data}") {
"message": {
data = await response.json() "content": data["content"][0]["text"]
return {
"choices": [
{
"message": {
"content": data["content"][0]["text"]
}
} }
],
"usage": {
"prompt_tokens": data["usage"]["input_tokens"],
"completion_tokens": data["usage"]["output_tokens"],
"total_tokens": data["usage"]["input_tokens"] + data["usage"]["output_tokens"]
} }
],
"usage": {
"prompt_tokens": data["usage"]["input_tokens"],
"completion_tokens": data["usage"]["output_tokens"],
"total_tokens": data["usage"]["input_tokens"] + data["usage"]["output_tokens"]
} }
}
+11 -36
View File
@@ -1,49 +1,24 @@
"""The HA Text AI integration.""" """The HA Text AI coordinator."""
from __future__ import annotations from __future__ import annotations
import logging import logging
import os
import shutil
from datetime import datetime, timedelta from datetime import datetime, timedelta
from typing import Any, Dict from typing import Any, Dict, List, Optional
import voluptuous as vol from homeassistant.core import HomeAssistant
from async_timeout import timeout from homeassistant.helpers.update_coordinator import DataUpdateCoordinator
from homeassistant.util import dt as dt_util
from homeassistant.config_entries import ConfigEntry
from homeassistant.const import CONF_API_KEY, CONF_NAME, Platform
from homeassistant.core import HomeAssistant, ServiceCall
from homeassistant.exceptions import ConfigEntryNotReady, HomeAssistantError
from homeassistant.helpers import config_validation as cv
from homeassistant.helpers import aiohttp_client
from .coordinator import HATextAICoordinator
from .api_client import APIClient
from .const import ( from .const import (
DOMAIN, DOMAIN,
PLATFORMS, STATE_READY,
CONF_MODEL, STATE_PROCESSING,
CONF_TEMPERATURE, STATE_ERROR,
CONF_MAX_TOKENS, STATE_RATE_LIMITED,
CONF_API_ENDPOINT, STATE_MAINTENANCE,
CONF_REQUEST_INTERVAL,
CONF_API_PROVIDER,
API_PROVIDER_OPENAI,
API_PROVIDER_ANTHROPIC,
DEFAULT_MODEL,
DEFAULT_TEMPERATURE,
DEFAULT_MAX_TOKENS,
DEFAULT_OPENAI_ENDPOINT,
DEFAULT_ANTHROPIC_ENDPOINT,
DEFAULT_REQUEST_INTERVAL,
API_TIMEOUT,
SERVICE_ASK_QUESTION,
SERVICE_CLEAR_HISTORY,
SERVICE_GET_HISTORY,
SERVICE_SET_SYSTEM_PROMPT,
) )
_LOGGER = logging.getLogger(__name__) _LOGGER = logging.getLogger(__name__)
class HATextAICoordinator(DataUpdateCoordinator): class HATextAICoordinator(DataUpdateCoordinator):
"""Class to manage fetching data from the API.""" """Class to manage fetching data from the API."""