mirror of
https://github.com/smkrv/ha-text-ai.git
synced 2026-07-22 07:03:58 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
76c5629fa0 | ||
|
|
7958bd010b | ||
|
|
8cd876195a | ||
|
|
376753e001 | ||
|
|
b6e73e847d | ||
|
|
440c734214 | ||
|
|
73788373cd | ||
|
|
4bfc96019b | ||
|
|
2138fc7654 | ||
|
|
95bd2ebb41 | ||
|
|
cad0fd7031 | ||
|
|
c003b258f6 | ||
|
|
65a10c77f4 | ||
|
|
e1463828c9 | ||
|
|
5ebb9c9c66 | ||
|
|
f17c631a79 | ||
|
|
0e06794384 | ||
|
|
d8a924909b | ||
|
|
29f1659a02 | ||
|
|
5b7905de80 | ||
|
|
cf9ac6dcea | ||
|
|
568eb3e16c | ||
|
|
53fb150389 | ||
|
|
acbb53d2af | ||
|
|
e19db29441 |
-41
@@ -1,41 +0,0 @@
|
||||
# Python
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
*.so
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
|
||||
# Home Assistant
|
||||
.storage
|
||||
.cloud
|
||||
.google.token
|
||||
|
||||
# IDE
|
||||
.idea/
|
||||
.vscode/
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
|
||||
# OS
|
||||
.DS_Store
|
||||
Thumbs.db
|
||||
*.psd
|
||||
*.zip
|
||||
*.txt
|
||||
*.pdf
|
||||
@@ -2,8 +2,8 @@
|
||||
|
||||
<div align="center">
|
||||
|
||||
  [](https://creativecommons.org/licenses/by-nc-sa/4.0/) [](https://github.com/hacs/integration)
|
||||
       
|
||||
  [](https://creativecommons.org/licenses/by-nc-sa/4.0/) [](https://github.com/hacs/integration)
|
||||
       
|
||||
|
||||
|
||||
<img src="https://github.com/smkrv/ha-text-ai/blob/15c717fcb0204bf4a0d4b4b4c6f0bb93e9f6c9a9/custom_components/ha_text_ai/icons/logo%402x.png" alt="HA Text AI" style="width: 50%; max-width: 256px; max-height: 128px; aspect-ratio: 2/1; object-fit: contain;"/>
|
||||
@@ -12,22 +12,23 @@
|
||||
</div>
|
||||
|
||||
<p align="center">
|
||||
Transform your smart home experience with powerful AI assistance powered by multiple AI providers including OpenAI GPT and Anthropic Claude models. Get intelligent responses, automate complex scenarios, and enhance your home automation with advanced natural language processing.
|
||||
Transform your smart home experience with powerful AI assistance powered by multiple AI providers including OpenAI GPT, DeepSeek and Anthropic Claude models. Get intelligent responses, automate complex scenarios, and enhance your home automation with advanced natural language processing.
|
||||
|
||||
</p>
|
||||
|
||||
---
|
||||
|
||||
> [!IMPORTANT]
|
||||
> 🤝 Community Driven
|
||||
> 🤝 Community Driven: for more details on the integration,
|
||||
> check out the discussion on the **[Home Assistant Community forum](https://community.home-assistant.io/t/ha-text-ai-transforming-home-automation-through-multi-llm-integration/799741)**
|
||||
>
|
||||
> <a href="https://community.home-assistant.io/t/ha-text-ai-transforming-home-automation-with-multi-provider-language-models/799741"><img src="https://img.shields.io/badge/Community-blue?style=for-the-badge&logo=homeassistant&logoColor=white&color=03a9f4"/></a>
|
||||
> <a href="https://my.home-assistant.io/redirect/hacs_repository/?owner=smkrv&repository=ha-text-ai&category=Integration"><img src="https://my.home-assistant.io/badges/hacs_repository.svg" width="210" height="auto"></a>
|
||||
>
|
||||
> [Screenshots](assets/images/screenshots/screenshot.jpg)
|
||||
|
||||
## 🌟 Features
|
||||
|
||||
- 🧠 **Multi-Provider AI Integration**: Support for OpenAI GPT and Anthropic Claude models
|
||||
- 🧠 **Multi-Provider AI Integration**: Support for OpenAI GPT, DeepSeek and Anthropic Claude models
|
||||
- 💬 **Advanced Language Processing**: Context-aware, multi-turn conversations
|
||||
- 📝 **Enhanced Memory Management**: Secure file-based history storage
|
||||
- ⚡ **Performance Optimization**: Efficient token usage and smart rate limiting
|
||||
@@ -42,6 +43,7 @@ Transform your smart home experience with powerful AI assistance powered by mult
|
||||
### 🧠 **Multi-Provider AI Integration**
|
||||
- Support for OpenAI GPT models
|
||||
- Anthropic Claude integration
|
||||
- DeepSeek integration
|
||||
- Custom API endpoints
|
||||
- Flexible model selection
|
||||
|
||||
@@ -104,12 +106,13 @@ Transform your smart home experience with powerful AI assistance powered by mult
|
||||
|
||||
## 📋 Prerequisites
|
||||
|
||||
- Home Assistant 2024.11 or later
|
||||
- Home Assistant 2024.2.2 or later
|
||||
- Active API key from:
|
||||
- OpenAI ([Get key](https://platform.openai.com/account/api-keys))
|
||||
- Anthropic ([Get key](https://console.anthropic.com/))
|
||||
- DeepSeek 🆕 ([Get key](https://platform.deepseek.com/api_keys))
|
||||
- DeepSeek ([Get key](https://platform.deepseek.com/api_keys))
|
||||
- OpenRouter ([Get key](https://openrouter.ai/keys))
|
||||
- Google Gemini 🆕 ([Get key](https://ai.google.dev/gemini-api/docs/api-key)) thanks to ([@Azzedde](https://github.com/Azzedde))
|
||||
- Any OpenAI-compatible API provider
|
||||
- Python 3.9 or newer
|
||||
- Stable internet connection
|
||||
@@ -117,7 +120,7 @@ Transform your smart home experience with powerful AI assistance powered by mult
|
||||
## Configuration Options
|
||||
|
||||
### 🔧 **Core Configuration Settings**
|
||||
- 🌐 **API Provider**: OpenAI/Anthropic
|
||||
- 🌐 **API Provider**: OpenAI/Anthropic/DeepSeek
|
||||
- 🔑 **API Key**: Provider-specific authentication
|
||||
- 🤖 **Model Selection**: Flexible, provider-specific models
|
||||
- 🌡️ **Temperature**: Creativity control (0.0-2.0)
|
||||
@@ -157,6 +160,9 @@ To be compatible, a provider should support:
|
||||
## ⚡ Installation
|
||||
|
||||
### HACS Installation (Recommended)
|
||||
>[!TIP]
|
||||
>HA Text AI is available in the default HACS repository. You can install it directly through HACS or click the button below to open it there.
|
||||
|
||||
<a href="https://my.home-assistant.io/redirect/hacs_repository/?owner=smkrv&repository=ha-text-ai&category=Integration"><img src="https://my.home-assistant.io/badges/hacs_repository.svg" width="170" height="auto"></a>
|
||||
1. Open HACS in Home Assistant
|
||||
2. Click on "Integrations"
|
||||
@@ -167,8 +173,6 @@ To be compatible, a provider should support:
|
||||
7. Click "Download"
|
||||
8. Restart Home Assistant
|
||||
|
||||
Note: Also Integration has been submitted to HACS store and is currently pending review in [pull request #2896](https://github.com/hacs/default/pull/2896).
|
||||
|
||||
### Manual Installation
|
||||
1. Download the latest release
|
||||
2. Extract and copy `custom_components/ha_text_ai` to your `custom_components` directory
|
||||
|
||||
@@ -40,13 +40,16 @@ from .const import (
|
||||
API_PROVIDER_OPENAI,
|
||||
API_PROVIDER_ANTHROPIC,
|
||||
API_PROVIDER_DEEPSEEK,
|
||||
API_PROVIDER_GEMINI,
|
||||
DEFAULT_MODEL,
|
||||
DEFAULT_DEEPSEEK_MODEL,
|
||||
DEFAULT_GEMINI_MODEL,
|
||||
DEFAULT_TEMPERATURE,
|
||||
DEFAULT_MAX_TOKENS,
|
||||
DEFAULT_OPENAI_ENDPOINT,
|
||||
DEFAULT_ANTHROPIC_ENDPOINT,
|
||||
DEFAULT_DEEPSEEK_ENDPOINT,
|
||||
DEFAULT_GEMINI_ENDPOINT,
|
||||
DEFAULT_REQUEST_INTERVAL,
|
||||
DEFAULT_CONTEXT_MESSAGES,
|
||||
API_TIMEOUT,
|
||||
@@ -236,10 +239,17 @@ async def async_setup(hass: HomeAssistant, config: ConfigType) -> bool:
|
||||
async def async_check_api(session, endpoint: str, headers: dict, provider: str) -> bool:
|
||||
"""Check API availability for different providers."""
|
||||
try:
|
||||
if provider == API_PROVIDER_ANTHROPIC:
|
||||
if provider == API_PROVIDER_GEMINI:
|
||||
# Gemini API does not support GET /models for validation, just check key presence
|
||||
if headers.get("Authorization", "").replace("Bearer ", ""):
|
||||
return True
|
||||
else:
|
||||
_LOGGER.error("Gemini API key is missing or empty")
|
||||
return False
|
||||
elif provider == API_PROVIDER_ANTHROPIC:
|
||||
check_url = f"{endpoint}/v1/models"
|
||||
elif provider == API_PROVIDER_DEEPSEEK:
|
||||
check_url = f"{endpoint}/models" # DeepSeek
|
||||
check_url = f"{endpoint}/models"
|
||||
else: # OpenAI
|
||||
check_url = f"{endpoint}/models"
|
||||
|
||||
@@ -248,7 +258,8 @@ async def async_check_api(session, endpoint: str, headers: dict, provider: str)
|
||||
if response.status in [200, 404]:
|
||||
return True
|
||||
elif response.status == 401:
|
||||
raise ConfigEntryNotReady("Invalid API key")
|
||||
_LOGGER.error("Invalid API key")
|
||||
return False
|
||||
elif response.status == 429:
|
||||
_LOGGER.warning("Rate limit exceeded during API check")
|
||||
return False
|
||||
|
||||
@@ -11,6 +11,7 @@ import asyncio
|
||||
from typing import Any, Dict, List, Optional
|
||||
from aiohttp import ClientSession, ClientTimeout
|
||||
from async_timeout import timeout
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from homeassistant.core import HomeAssistant
|
||||
from homeassistant.exceptions import HomeAssistantError
|
||||
@@ -20,6 +21,7 @@ from .const import (
|
||||
API_PROVIDER_ANTHROPIC,
|
||||
API_PROVIDER_DEEPSEEK,
|
||||
API_PROVIDER_OPENAI,
|
||||
API_PROVIDER_GEMINI,
|
||||
MIN_TEMPERATURE,
|
||||
MAX_TEMPERATURE,
|
||||
MIN_MAX_TOKENS,
|
||||
@@ -115,6 +117,10 @@ class APIClient:
|
||||
return await self._create_deepseek_completion(
|
||||
model, messages, temperature, max_tokens
|
||||
)
|
||||
elif self.api_provider == API_PROVIDER_GEMINI:
|
||||
return await self._create_gemini_completion(
|
||||
model, messages, temperature, max_tokens
|
||||
)
|
||||
else:
|
||||
return await self._create_openai_completion(
|
||||
model, messages, temperature, max_tokens
|
||||
@@ -238,6 +244,143 @@ class APIClient:
|
||||
_LOGGER.error(f"Connection check failed: {str(e)}")
|
||||
return False
|
||||
|
||||
async def _create_gemini_completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[Dict[str, str]],
|
||||
temperature: float,
|
||||
max_tokens: int,
|
||||
) -> Dict[str, Any]:
|
||||
"""Create completion using Gemini API with google-genai library.
|
||||
|
||||
Args:
|
||||
model: The model name to use
|
||||
messages: List of message dictionaries with role and content
|
||||
temperature: Sampling temperature between 0.0 and 2.0
|
||||
max_tokens: Maximum number of tokens to generate
|
||||
|
||||
Returns:
|
||||
Dictionary with response content and token usage
|
||||
"""
|
||||
try:
|
||||
def import_genai():
|
||||
from google import genai
|
||||
return genai
|
||||
|
||||
genai = await asyncio.to_thread(import_genai)
|
||||
|
||||
# Extract API key from headers (Bearer token)
|
||||
api_key = self.headers.get("Authorization", "").replace("Bearer ", "")
|
||||
|
||||
def create_client():
|
||||
if self.endpoint and self.endpoint != "https://generativelanguage.googleapis.com/v1beta":
|
||||
return genai.Client(api_key=api_key, transport="rest",
|
||||
client_options={"api_endpoint": self.endpoint})
|
||||
else:
|
||||
return genai.Client(api_key=api_key)
|
||||
|
||||
client = await asyncio.to_thread(create_client)
|
||||
|
||||
# Process messages to extract system instruction and chat history
|
||||
system_instruction = ""
|
||||
contents = []
|
||||
|
||||
for msg in messages:
|
||||
if msg['role'] == 'system':
|
||||
system_instruction += msg['content'] + "\n"
|
||||
else:
|
||||
# For chat history, we need to convert to the format Gemini expects
|
||||
role = "user" if msg['role'] == 'user' else "model"
|
||||
contents.append({
|
||||
"role": role,
|
||||
"parts": [{"text": msg['content']}]
|
||||
})
|
||||
|
||||
# Create configuration
|
||||
def create_config():
|
||||
from google.genai import types
|
||||
config = types.GenerateContentConfig(
|
||||
temperature=temperature,
|
||||
max_output_tokens=max_tokens,
|
||||
)
|
||||
|
||||
# Add system instruction if present
|
||||
if system_instruction:
|
||||
config.system_instruction = system_instruction.strip()
|
||||
|
||||
return config
|
||||
|
||||
config = await asyncio.to_thread(create_config)
|
||||
|
||||
def generate_content():
|
||||
# For single message without history, use generate_content
|
||||
if len(contents) <= 1:
|
||||
# If we have no content yet, create a simple prompt
|
||||
if not contents:
|
||||
prompt = "I need your assistance."
|
||||
else:
|
||||
prompt = contents[0]["parts"][0]["text"]
|
||||
|
||||
return client.models.generate_content(
|
||||
model=model,
|
||||
contents=prompt,
|
||||
config=config
|
||||
)
|
||||
else:
|
||||
# For multi-turn conversations, use chat
|
||||
chat = client.chats.create(model=model, config=config)
|
||||
|
||||
# Send all messages in sequence
|
||||
for content in contents:
|
||||
if content["role"] == "user":
|
||||
response = chat.send_message(content["parts"][0]["text"])
|
||||
# We don't send assistant messages as they're already part of the history
|
||||
|
||||
return response
|
||||
|
||||
response = await asyncio.to_thread(generate_content)
|
||||
|
||||
# Extract response text
|
||||
def extract_response():
|
||||
response_text = response.text if hasattr(response, 'text') else ""
|
||||
|
||||
# Try to get token usage if available
|
||||
usage = {}
|
||||
if hasattr(response, 'usage_metadata'):
|
||||
usage = {
|
||||
"prompt_tokens": getattr(response.usage_metadata, 'prompt_token_count', 0),
|
||||
"completion_tokens": getattr(response.usage_metadata, 'candidates_token_count', 0),
|
||||
"total_tokens": getattr(response.usage_metadata, 'total_token_count', 0)
|
||||
}
|
||||
else:
|
||||
# Estimate token count as fallback
|
||||
usage = {
|
||||
"prompt_tokens": len(" ".join([m["content"] for m in messages]).split()) // 3,
|
||||
"completion_tokens": len(response_text.split()) // 3,
|
||||
"total_tokens": 0 # Will be calculated below
|
||||
}
|
||||
usage["total_tokens"] = usage["prompt_tokens"] + usage["completion_tokens"]
|
||||
|
||||
return response_text, usage
|
||||
|
||||
response_text, usage = await asyncio.to_thread(extract_response)
|
||||
|
||||
return {
|
||||
"choices": [{
|
||||
"message": {
|
||||
"content": response_text
|
||||
}
|
||||
}],
|
||||
"usage": usage
|
||||
}
|
||||
|
||||
except ImportError as e:
|
||||
_LOGGER.error(f"Google Gemini library not installed: {str(e)}")
|
||||
raise HomeAssistantError(f"Missing dependency: {str(e)}. Please install google-genai.")
|
||||
except Exception as e:
|
||||
_LOGGER.error(f"Gemini API error: {str(e)}")
|
||||
raise HomeAssistantError(f"Gemini API error: {str(e)}")
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
"""Shutdown API client."""
|
||||
_LOGGER.debug("Shutting down API client")
|
||||
|
||||
@@ -8,6 +8,7 @@ Config flow for HA text AI integration.
|
||||
"""
|
||||
import logging
|
||||
from typing import Any, Dict, Optional
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import voluptuous as vol
|
||||
from homeassistant import config_entries
|
||||
@@ -29,15 +30,18 @@ from .const import (
|
||||
API_PROVIDER_OPENAI,
|
||||
API_PROVIDER_ANTHROPIC,
|
||||
API_PROVIDER_DEEPSEEK,
|
||||
API_PROVIDER_GEMINI,
|
||||
API_PROVIDERS,
|
||||
DEFAULT_MODEL,
|
||||
DEFAULT_DEEPSEEK_MODEL,
|
||||
DEFAULT_GEMINI_MODEL,
|
||||
DEFAULT_TEMPERATURE,
|
||||
DEFAULT_MAX_TOKENS,
|
||||
DEFAULT_REQUEST_INTERVAL,
|
||||
DEFAULT_OPENAI_ENDPOINT,
|
||||
DEFAULT_ANTHROPIC_ENDPOINT,
|
||||
DEFAULT_DEEPSEEK_ENDPOINT,
|
||||
DEFAULT_GEMINI_ENDPOINT,
|
||||
DEFAULT_CONTEXT_MESSAGES,
|
||||
MIN_TEMPERATURE,
|
||||
MAX_TEMPERATURE,
|
||||
@@ -93,15 +97,20 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
||||
self._errors = {}
|
||||
|
||||
if user_input is None:
|
||||
# Выбор endpoint по провайдеру
|
||||
# Selecting an endpoint by provider
|
||||
default_endpoint = {
|
||||
API_PROVIDER_OPENAI: DEFAULT_OPENAI_ENDPOINT,
|
||||
API_PROVIDER_ANTHROPIC: DEFAULT_ANTHROPIC_ENDPOINT,
|
||||
API_PROVIDER_DEEPSEEK: DEFAULT_DEEPSEEK_ENDPOINT,
|
||||
API_PROVIDER_GEMINI: DEFAULT_GEMINI_ENDPOINT,
|
||||
}.get(self._provider, DEFAULT_OPENAI_ENDPOINT)
|
||||
|
||||
# Выбор модели по умолчанию по провайдеру
|
||||
default_model = DEFAULT_DEEPSEEK_MODEL if self._provider == API_PROVIDER_DEEPSEEK else DEFAULT_MODEL
|
||||
# Selecting the default model by provider
|
||||
default_model = (
|
||||
DEFAULT_DEEPSEEK_MODEL if self._provider == API_PROVIDER_DEEPSEEK else
|
||||
DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else
|
||||
DEFAULT_MODEL
|
||||
)
|
||||
|
||||
return self.async_show_form(
|
||||
step_id="provider",
|
||||
@@ -139,41 +148,173 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
||||
})
|
||||
)
|
||||
|
||||
# Debug log to identify what's in the input
|
||||
_LOGGER.debug(f"Provider step input data: {user_input}")
|
||||
|
||||
input_copy = user_input.copy()
|
||||
|
||||
# Check if CONF_NAME exists in input_copy and ensure it's not empty
|
||||
if CONF_NAME not in input_copy or not input_copy[CONF_NAME]:
|
||||
_LOGGER.warning(f"Missing name in configuration input: {input_copy}")
|
||||
input_copy[CONF_NAME] = f"gemini_assistant_{datetime.now().strftime('%Y%m%d_%H%M%S')}"
|
||||
_LOGGER.info(f"Auto-generated name: {input_copy[CONF_NAME]}")
|
||||
|
||||
# Ensure API key is present
|
||||
if CONF_API_KEY not in input_copy or not input_copy[CONF_API_KEY]:
|
||||
self._errors["base"] = "invalid_auth"
|
||||
_LOGGER.error("API validation error: 'api_key'")
|
||||
return self.async_show_form(
|
||||
step_id="provider",
|
||||
data_schema=vol.Schema({
|
||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
||||
vol.Required(CONF_API_KEY): str,
|
||||
vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else DEFAULT_MODEL)): str,
|
||||
vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): str,
|
||||
vol.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All(
|
||||
vol.Coerce(float),
|
||||
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
||||
),
|
||||
vol.Optional(CONF_MAX_TOKENS, default=input_copy.get(CONF_MAX_TOKENS, DEFAULT_MAX_TOKENS)): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=MIN_MAX_TOKENS, max=MAX_MAX_TOKENS)
|
||||
),
|
||||
vol.Optional(CONF_REQUEST_INTERVAL, default=input_copy.get(CONF_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL)): vol.All(
|
||||
vol.Coerce(float),
|
||||
vol.Range(min=MIN_REQUEST_INTERVAL)
|
||||
),
|
||||
vol.Optional(
|
||||
CONF_CONTEXT_MESSAGES,
|
||||
default=input_copy.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=20)
|
||||
),
|
||||
vol.Optional(
|
||||
CONF_MAX_HISTORY_SIZE,
|
||||
default=input_copy.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=100)
|
||||
),
|
||||
}),
|
||||
errors=self._errors
|
||||
)
|
||||
|
||||
try:
|
||||
# Validate and normalize the name
|
||||
normalized_name = self._validate_and_normalize_name(input_copy[CONF_NAME])
|
||||
input_copy[CONF_NAME] = normalized_name
|
||||
except ValueError as e:
|
||||
return self.async_show_form(
|
||||
step_id="provider",
|
||||
data_schema=vol.Schema({
|
||||
vol.Required(CONF_NAME, default=input_copy[CONF_NAME]): str,
|
||||
vol.Required(CONF_API_KEY, default=input_copy[CONF_API_KEY]): str,
|
||||
vol.Required(CONF_MODEL, default=input_copy[CONF_MODEL]): str,
|
||||
vol.Required(CONF_API_ENDPOINT, default=input_copy[CONF_API_ENDPOINT]): str,
|
||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
||||
vol.Required(CONF_API_KEY, default=input_copy.get(CONF_API_KEY, "")): str,
|
||||
vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else DEFAULT_MODEL)): str,
|
||||
vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): str,
|
||||
vol.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All(
|
||||
vol.Coerce(float),
|
||||
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
||||
),
|
||||
vol.Optional(CONF_MAX_TOKENS, default=input_copy.get(CONF_MAX_TOKENS, DEFAULT_MAX_TOKENS)): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=MIN_MAX_TOKENS, max=MAX_MAX_TOKENS)
|
||||
),
|
||||
vol.Optional(CONF_REQUEST_INTERVAL, default=input_copy.get(CONF_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL)): vol.All(
|
||||
vol.Coerce(float),
|
||||
vol.Range(min=MIN_REQUEST_INTERVAL)
|
||||
),
|
||||
vol.Optional(
|
||||
CONF_CONTEXT_MESSAGES,
|
||||
default=input_copy.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=20)
|
||||
),
|
||||
vol.Optional(
|
||||
CONF_MAX_HISTORY_SIZE,
|
||||
default=input_copy.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=100)
|
||||
),
|
||||
}),
|
||||
errors={"name": str(e)}
|
||||
)
|
||||
|
||||
try:
|
||||
if not await self._async_validate_api(input_copy):
|
||||
return self.async_show_form(
|
||||
step_id="provider",
|
||||
data_schema=vol.Schema({}),
|
||||
errors=self._errors
|
||||
)
|
||||
# Special handling for Gemini API validation
|
||||
if self._provider == API_PROVIDER_GEMINI:
|
||||
# For Gemini, we just check if API key is present as there's no simple endpoint to validate
|
||||
if not input_copy.get(CONF_API_KEY):
|
||||
self._errors["base"] = "invalid_auth"
|
||||
_LOGGER.error("API validation error: 'api_key'")
|
||||
return self.async_show_form(
|
||||
step_id="provider",
|
||||
data_schema=vol.Schema({
|
||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
||||
vol.Required(CONF_API_KEY): str,
|
||||
vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL)): str,
|
||||
vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_GEMINI_ENDPOINT)): str,
|
||||
# Other fields remain the same
|
||||
}),
|
||||
errors=self._errors
|
||||
)
|
||||
else:
|
||||
# For other providers, validate API connection
|
||||
if not await self._async_validate_api(input_copy):
|
||||
return self.async_show_form(
|
||||
step_id="provider",
|
||||
data_schema=vol.Schema({
|
||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
||||
vol.Required(CONF_API_KEY, default=input_copy.get(CONF_API_KEY, "")): str,
|
||||
vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_MODEL)): str,
|
||||
vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_OPENAI_ENDPOINT)): str,
|
||||
vol.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All(
|
||||
vol.Coerce(float),
|
||||
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE)
|
||||
),
|
||||
vol.Optional(CONF_MAX_TOKENS, default=input_copy.get(CONF_MAX_TOKENS, DEFAULT_MAX_TOKENS)): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=MIN_MAX_TOKENS, max=MAX_MAX_TOKENS)
|
||||
),
|
||||
vol.Optional(CONF_REQUEST_INTERVAL, default=input_copy.get(CONF_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL)): vol.All(
|
||||
vol.Coerce(float),
|
||||
vol.Range(min=MIN_REQUEST_INTERVAL)
|
||||
),
|
||||
vol.Optional(
|
||||
CONF_CONTEXT_MESSAGES,
|
||||
default=input_copy.get(CONF_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=20)
|
||||
),
|
||||
vol.Optional(
|
||||
CONF_MAX_HISTORY_SIZE,
|
||||
default=input_copy.get(CONF_MAX_HISTORY_SIZE, DEFAULT_MAX_HISTORY)
|
||||
): vol.All(
|
||||
vol.Coerce(int),
|
||||
vol.Range(min=1, max=100)
|
||||
),
|
||||
}),
|
||||
errors=self._errors
|
||||
)
|
||||
except Exception as e:
|
||||
# Handle any unexpected exceptions during validation
|
||||
_LOGGER.exception("Unexpected error during API validation")
|
||||
return self.async_show_form(
|
||||
step_id="provider",
|
||||
data_schema=vol.Schema({}),
|
||||
data_schema=vol.Schema({
|
||||
vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
|
||||
vol.Required(CONF_API_KEY, default=input_copy.get(CONF_API_KEY, "")): str,
|
||||
vol.Required(CONF_MODEL, default=input_copy.get(CONF_MODEL, DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else DEFAULT_MODEL)): str,
|
||||
vol.Required(CONF_API_ENDPOINT, default=input_copy.get(CONF_API_ENDPOINT, DEFAULT_GEMINI_ENDPOINT if self._provider == API_PROVIDER_GEMINI else DEFAULT_OPENAI_ENDPOINT)): str,
|
||||
# Other fields remain the same
|
||||
}),
|
||||
errors={"base": str(e)}
|
||||
)
|
||||
|
||||
# All validation passed, create the entry
|
||||
return await self._create_entry(input_copy)
|
||||
|
||||
def _validate_and_normalize_name(self, name: str) -> str:
|
||||
@@ -211,23 +352,34 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
||||
async def _async_validate_api(self, user_input: Dict[str, Any]) -> bool:
|
||||
"""Validate API connection."""
|
||||
try:
|
||||
if CONF_API_KEY not in user_input:
|
||||
_LOGGER.error("API validation error: 'api_key'")
|
||||
self._errors["base"] = "invalid_auth"
|
||||
return False
|
||||
|
||||
session = async_get_clientsession(self.hass)
|
||||
headers = self._get_api_headers(user_input)
|
||||
endpoint = user_input[CONF_API_ENDPOINT].rstrip('/')
|
||||
|
||||
check_url = (
|
||||
f"{endpoint}/v1/models" if self._provider == API_PROVIDER_ANTHROPIC
|
||||
else f"{endpoint}/models"
|
||||
)
|
||||
|
||||
async with session.get(check_url, headers=headers) as response:
|
||||
if response.status == 401:
|
||||
if self._provider == API_PROVIDER_GEMINI:
|
||||
if not user_input[CONF_API_KEY]:
|
||||
self._errors["base"] = "invalid_auth"
|
||||
return False
|
||||
elif response.status not in [200, 404]:
|
||||
self._errors["base"] = "cannot_connect"
|
||||
return False
|
||||
return True
|
||||
else:
|
||||
check_url = (
|
||||
f"{endpoint}/v1/models" if self._provider == API_PROVIDER_ANTHROPIC
|
||||
else f"{endpoint}/models"
|
||||
)
|
||||
|
||||
async with session.get(check_url, headers=headers) as response:
|
||||
if response.status == 401:
|
||||
self._errors["base"] = "invalid_auth"
|
||||
return False
|
||||
elif response.status not in [200, 404]:
|
||||
self._errors["base"] = "cannot_connect"
|
||||
return False
|
||||
return True
|
||||
|
||||
except Exception as err:
|
||||
_LOGGER.error("API validation error: %s", str(err))
|
||||
@@ -236,6 +388,9 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
||||
|
||||
def _get_api_headers(self, user_input: Dict[str, Any]) -> Dict[str, str]:
|
||||
"""Get API headers based on provider."""
|
||||
if CONF_API_KEY not in user_input:
|
||||
return {"Content-Type": "application/json"}
|
||||
|
||||
api_key = user_input[CONF_API_KEY]
|
||||
|
||||
if self._provider == API_PROVIDER_ANTHROPIC:
|
||||
@@ -244,6 +399,11 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
||||
"anthropic-version": "2023-06-01",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
elif self._provider == API_PROVIDER_GEMINI:
|
||||
return {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
return {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json"
|
||||
@@ -256,7 +416,11 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
||||
|
||||
unique_id = f"{DOMAIN}_{normalized_name}_{self._provider}".lower()
|
||||
|
||||
default_model = DEFAULT_DEEPSEEK_MODEL if self._provider == API_PROVIDER_DEEPSEEK else DEFAULT_MODEL
|
||||
default_model = (
|
||||
DEFAULT_DEEPSEEK_MODEL if self._provider == API_PROVIDER_DEEPSEEK else
|
||||
DEFAULT_GEMINI_MODEL if self._provider == API_PROVIDER_GEMINI else
|
||||
DEFAULT_MODEL
|
||||
)
|
||||
|
||||
entry_data = {
|
||||
CONF_API_PROVIDER: self._provider,
|
||||
@@ -304,6 +468,13 @@ class OptionsFlowHandler(config_entries.OptionsFlow):
|
||||
return self.async_create_entry(title="", data=user_input)
|
||||
|
||||
current_data = {**self.config_entry.data, **self.config_entry.options}
|
||||
provider = current_data.get(CONF_API_PROVIDER)
|
||||
|
||||
default_model = (
|
||||
DEFAULT_DEEPSEEK_MODEL if provider == API_PROVIDER_DEEPSEEK else
|
||||
DEFAULT_GEMINI_MODEL if provider == API_PROVIDER_GEMINI else
|
||||
DEFAULT_MODEL
|
||||
)
|
||||
|
||||
return self.async_show_form(
|
||||
step_id="init",
|
||||
|
||||
@@ -12,6 +12,8 @@ from typing import Final
|
||||
import voluptuous as vol
|
||||
from homeassistant.const import Platform, CONF_API_KEY, CONF_NAME
|
||||
from homeassistant.helpers import config_validation as cv
|
||||
import logging
|
||||
_LOGGER = logging.getLogger(__name__)
|
||||
|
||||
# Domain and platforms
|
||||
DOMAIN: Final = "ha_text_ai"
|
||||
@@ -22,11 +24,13 @@ CONF_API_PROVIDER: Final = "api_provider"
|
||||
API_PROVIDER_OPENAI: Final = "openai"
|
||||
API_PROVIDER_ANTHROPIC: Final = "anthropic"
|
||||
API_PROVIDER_DEEPSEEK: Final = "deepseek"
|
||||
API_PROVIDER_GEMINI: Final = "gemini"
|
||||
|
||||
API_PROVIDERS: Final = [
|
||||
API_PROVIDER_OPENAI,
|
||||
API_PROVIDER_ANTHROPIC,
|
||||
API_PROVIDER_DEEPSEEK
|
||||
API_PROVIDER_DEEPSEEK,
|
||||
API_PROVIDER_GEMINI
|
||||
]
|
||||
|
||||
# Read version from manifest.json
|
||||
@@ -49,6 +53,7 @@ except Exception as err:
|
||||
DEFAULT_OPENAI_ENDPOINT: Final = "https://api.openai.com/v1"
|
||||
DEFAULT_ANTHROPIC_ENDPOINT: Final = "https://api.anthropic.com"
|
||||
DEFAULT_DEEPSEEK_ENDPOINT: Final = "https://api.deepseek.com"
|
||||
DEFAULT_GEMINI_ENDPOINT: Final = "https://generativelanguage.googleapis.com/v1beta"
|
||||
|
||||
# Configuration constants
|
||||
CONF_MODEL: Final = "model"
|
||||
@@ -69,6 +74,7 @@ ICONS_SUBDOMAIN = "icons"
|
||||
# Default values
|
||||
DEFAULT_MODEL: Final = "gpt-4o-mini"
|
||||
DEFAULT_DEEPSEEK_MODEL: Final = "deepseek-chat"
|
||||
DEFAULT_GEMINI_MODEL: Final = "gemini-2.0-flash"
|
||||
DEFAULT_TEMPERATURE: Final = 0.1
|
||||
DEFAULT_MAX_TOKENS: Final = 1000
|
||||
DEFAULT_REQUEST_INTERVAL: Final = 1.0
|
||||
|
||||
@@ -13,16 +13,17 @@
|
||||
"loggers": ["custom_components.ha_text_ai"],
|
||||
"mqtt": [],
|
||||
"quality_scale": "silver",
|
||||
"requirements": [
|
||||
"openai>=1.12.0",
|
||||
"anthropic>=0.8.0",
|
||||
"aiohttp>=3.8.0",
|
||||
"async-timeout>=4.0.0",
|
||||
"certifi>=2024.2.2"
|
||||
],
|
||||
"requirements": [
|
||||
"openai>=1.12.0",
|
||||
"anthropic>=0.8.0",
|
||||
"google-genai>=1.16.0",
|
||||
"aiohttp>=3.8.0",
|
||||
"async-timeout>=4.0.0",
|
||||
"certifi>=2024.2.2"
|
||||
],
|
||||
"single_config_entry": false,
|
||||
"ssdp": [],
|
||||
"usb": [],
|
||||
"version": "2.1.1",
|
||||
"version": "2.1.7",
|
||||
"zeroconf": []
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ Sensor platform for HA Text AI.
|
||||
import logging
|
||||
import math
|
||||
from typing import Any, Dict
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from homeassistant.components.sensor import (
|
||||
SensorEntity,
|
||||
|
||||
@@ -63,13 +63,13 @@ ask_question:
|
||||
|
||||
max_tokens:
|
||||
name: Max Tokens
|
||||
description: Maximum length of the response (1-4096 tokens)
|
||||
description: Maximum length of the response (tokens)
|
||||
required: false
|
||||
default: 1000
|
||||
selector:
|
||||
number:
|
||||
min: 1
|
||||
max: 4096
|
||||
max: 100000
|
||||
step: 1
|
||||
mode: box
|
||||
|
||||
|
||||
@@ -88,12 +88,13 @@
|
||||
"selector": {
|
||||
"api_provider": {
|
||||
"options": {
|
||||
"openai": "OpenAI (kompatibel)",
|
||||
"anthropic": "Anthropic (kompatibel)",
|
||||
"deepseek": "DeepSeek"
|
||||
}
|
||||
"openai": "OpenAI (compatible)",
|
||||
"anthropic": "Anthropic (compatible)",
|
||||
"deepseek": "DeepSeek",
|
||||
"gemini": "Google Gemini"
|
||||
}
|
||||
},
|
||||
}
|
||||
},
|
||||
"services": {
|
||||
"ask_question": {
|
||||
"name": "Frage stellen (HA Text AI)",
|
||||
|
||||
@@ -86,11 +86,12 @@
|
||||
}
|
||||
},
|
||||
"selector": {
|
||||
"api_provider": {
|
||||
"options": {
|
||||
"openai": "OpenAI (compatible)",
|
||||
"anthropic": "Anthropic (compatible)",
|
||||
"deepseek": "DeepSeek"
|
||||
"api_provider": {
|
||||
"options": {
|
||||
"openai": "OpenAI (compatible)",
|
||||
"anthropic": "Anthropic (compatible)",
|
||||
"deepseek": "DeepSeek",
|
||||
"gemini": "Google Gemini"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -90,10 +90,11 @@
|
||||
"options": {
|
||||
"openai": "OpenAI (compatible)",
|
||||
"anthropic": "Anthropic (compatible)",
|
||||
"deepseek": "DeepSeek"
|
||||
}
|
||||
"deepseek": "DeepSeek",
|
||||
"gemini": "Google Gemini"
|
||||
}
|
||||
},
|
||||
}
|
||||
},
|
||||
"services": {
|
||||
"ask_question": {
|
||||
"name": "Hacer Pregunta (HA Text AI)",
|
||||
|
||||
@@ -90,7 +90,8 @@
|
||||
"options": {
|
||||
"openai": "OpenAI (अनुकूलित)",
|
||||
"anthropic": "Anthropic (अनुकूलित)",
|
||||
"deepseek": "DeepSeek"
|
||||
"deepseek": "DeepSeek",
|
||||
"gemini": "Google Gemini"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -90,7 +90,8 @@
|
||||
"options": {
|
||||
"openai": "OpenAI (compatibile)",
|
||||
"anthropic": "Anthropic (compatibile)",
|
||||
"deepseek": "DeepSeek"
|
||||
"deepseek": "DeepSeek",
|
||||
"gemini": "Google Gemini"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -90,7 +90,8 @@
|
||||
"options": {
|
||||
"openai": "OpenAI (совместимый)",
|
||||
"anthropic": "Anthropic (совместимый)",
|
||||
"deepseek": "DeepSeek"
|
||||
"deepseek": "DeepSeek",
|
||||
"gemini": "Google Gemini"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -90,7 +90,8 @@
|
||||
"options": {
|
||||
"openai": "OpenAI (компатибилан)",
|
||||
"anthropic": "Anthropic (компатибилан)",
|
||||
"deepseek": "DeepSeek"
|
||||
"deepseek": "DeepSeek",
|
||||
"gemini": "Google Gemini"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -90,7 +90,8 @@
|
||||
"options": {
|
||||
"openai": "OpenAI(兼容)",
|
||||
"anthropic": "Anthropic(兼容)",
|
||||
"deepseek": "DeepSeek "
|
||||
"deepseek": "DeepSeek",
|
||||
"gemini": "Google Gemini"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
Reference in New Issue
Block a user