mirror of
https://github.com/smkrv/ha-text-ai.git
synced 2026-07-29 12:13:56 +08:00
Release v2.0.0
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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"]
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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."""
|
||||||
|
|||||||
Reference in New Issue
Block a user