Compare commits

...
14 Commits
Author SHA1 Message Date
SMKRV 76c5629fa0 refactor(google-gemini): rewrite integration using google-genai 1.16.0
Completely rewrote the Google Gemini integration logic based on google-genai 1.16.0 to fix issue #6.
Key changes:
- Updated to the latest google-genai library
- Made API endpoint abstract while retaining option for custom endpoint configuration
- Refactored logic and classes exclusively within Google Gemini implementation
- All changes are limited to Google Gemini integration refactoring with no impact on other functionality.
2025-05-21 01:27:47 +03:00
SMKRV 7958bd010b refactor(google-gemini): rewrite integration using google-genai 1.16.0
Completely rewrote the Google Gemini integration logic based on google-genai 1.16.0 to fix issue #6.
Key changes:
- Updated to the latest google-genai library
- Made API endpoint abstract while retaining option for custom endpoint configuration
- Refactored logic and classes exclusively within Google Gemini implementation
- All changes are limited to Google Gemini integration refactoring with no impact on other functionality.
2025-05-21 01:26:42 +03:00
SMKRV 8cd876195a Bump to version 2.1.6 2025-05-20 01:50:06 +03:00
SMKRV 376753e001 fix: correct field naming in Gemini API requests from camelCase to snake_case and improve message handling 2025-05-20 01:42:38 +03:00
SMKRV b6e73e847d fix(api_client): correct Google Gemini API integration
- Change JSON field names from camelCase to snake_case as required by Gemini API
  (generation_config, max_output_tokens, system_instruction)
- Improve message handling to ensure proper role alternation (user/model)
- Add safety checks for empty contents and ensure first message is always from user
- Implement robust error handling and response parsing
- Handle edge cases where candidatesTokenCount might be returned as a list

Fixes #6
2025-05-20 01:16:41 +03:00
SMKRV 440c734214 Bump release version to v2.1.4 2025-05-19 23:20:27 +03:00
SMKRV 73788373cd Release v2.1.3 2025-05-19 23:12:55 +03:00
SMKRV 4bfc96019b fix: DEFAULT_GEMINI_ENDPOINT 2025-05-19 15:53:43 +03:00
SMKRV 2138fc7654 fix: DEFAULT_GEMINI_ENDPOINT 2025-05-19 15:36:58 +03:00
SMKRV 95bd2ebb41 Add support for Google Gemini (thanks to @Azzedde) #5 2025-05-19 15:10:19 +03:00
smkrvandGitHub cad0fd7031 Merge pull request #5 from Azzedde/main
Add Gemini API provider support to HA Text AI integration by @Azzedde
2025-05-19 14:44:06 +03:00
Azzedde c003b258f6 Add Gemini API provider support to HA Text AI integration 2025-05-18 13:23:55 +02:00
SMKRV 65a10c77f4 ~ 2025-01-30 01:15:13 +03:00
SMKRV e1463828c9 ~ 2025-01-30 01:14:24 +03:00
17 changed files with 403 additions and 101 deletions
-41
View File
@@ -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
+6 -4
View File
@@ -106,12 +106,13 @@ Transform your smart home experience with powerful AI assistance powered by mult
## 📋 Prerequisites ## 📋 Prerequisites
- Home Assistant 2024.11 or later - Home Assistant 2024.2.2 or later
- Active API key from: - Active API key from:
- OpenAI ([Get key](https://platform.openai.com/account/api-keys)) - OpenAI ([Get key](https://platform.openai.com/account/api-keys))
- Anthropic ([Get key](https://console.anthropic.com/)) - 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)) - 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 - Any OpenAI-compatible API provider
- Python 3.9 or newer - Python 3.9 or newer
- Stable internet connection - Stable internet connection
@@ -159,6 +160,9 @@ To be compatible, a provider should support:
## ⚡ Installation ## ⚡ Installation
### HACS Installation (Recommended) ### 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> <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 1. Open HACS in Home Assistant
2. Click on "Integrations" 2. Click on "Integrations"
@@ -169,8 +173,6 @@ To be compatible, a provider should support:
7. Click "Download" 7. Click "Download"
8. Restart Home Assistant 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 ### Manual Installation
1. Download the latest release 1. Download the latest release
2. Extract and copy `custom_components/ha_text_ai` to your `custom_components` directory 2. Extract and copy `custom_components/ha_text_ai` to your `custom_components` directory
+14 -3
View File
@@ -40,13 +40,16 @@ from .const import (
API_PROVIDER_OPENAI, API_PROVIDER_OPENAI,
API_PROVIDER_ANTHROPIC, API_PROVIDER_ANTHROPIC,
API_PROVIDER_DEEPSEEK, API_PROVIDER_DEEPSEEK,
API_PROVIDER_GEMINI,
DEFAULT_MODEL, DEFAULT_MODEL,
DEFAULT_DEEPSEEK_MODEL, DEFAULT_DEEPSEEK_MODEL,
DEFAULT_GEMINI_MODEL,
DEFAULT_TEMPERATURE, DEFAULT_TEMPERATURE,
DEFAULT_MAX_TOKENS, DEFAULT_MAX_TOKENS,
DEFAULT_OPENAI_ENDPOINT, DEFAULT_OPENAI_ENDPOINT,
DEFAULT_ANTHROPIC_ENDPOINT, DEFAULT_ANTHROPIC_ENDPOINT,
DEFAULT_DEEPSEEK_ENDPOINT, DEFAULT_DEEPSEEK_ENDPOINT,
DEFAULT_GEMINI_ENDPOINT,
DEFAULT_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL,
DEFAULT_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES,
API_TIMEOUT, 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: async def async_check_api(session, endpoint: str, headers: dict, provider: str) -> bool:
"""Check API availability for different providers.""" """Check API availability for different providers."""
try: 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" check_url = f"{endpoint}/v1/models"
elif provider == API_PROVIDER_DEEPSEEK: elif provider == API_PROVIDER_DEEPSEEK:
check_url = f"{endpoint}/models" # DeepSeek check_url = f"{endpoint}/models"
else: # OpenAI else: # OpenAI
check_url = f"{endpoint}/models" 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]: if response.status in [200, 404]:
return True return True
elif response.status == 401: elif response.status == 401:
raise ConfigEntryNotReady("Invalid API key") _LOGGER.error("Invalid API key")
return False
elif response.status == 429: elif response.status == 429:
_LOGGER.warning("Rate limit exceeded during API check") _LOGGER.warning("Rate limit exceeded during API check")
return False return False
+143
View File
@@ -11,6 +11,7 @@ import asyncio
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
from aiohttp import ClientSession, ClientTimeout from aiohttp import ClientSession, ClientTimeout
from async_timeout import timeout from async_timeout import timeout
from datetime import datetime, timedelta
from homeassistant.core import HomeAssistant from homeassistant.core import HomeAssistant
from homeassistant.exceptions import HomeAssistantError from homeassistant.exceptions import HomeAssistantError
@@ -20,6 +21,7 @@ from .const import (
API_PROVIDER_ANTHROPIC, API_PROVIDER_ANTHROPIC,
API_PROVIDER_DEEPSEEK, API_PROVIDER_DEEPSEEK,
API_PROVIDER_OPENAI, API_PROVIDER_OPENAI,
API_PROVIDER_GEMINI,
MIN_TEMPERATURE, MIN_TEMPERATURE,
MAX_TEMPERATURE, MAX_TEMPERATURE,
MIN_MAX_TOKENS, MIN_MAX_TOKENS,
@@ -115,6 +117,10 @@ class APIClient:
return await self._create_deepseek_completion( return await self._create_deepseek_completion(
model, messages, temperature, max_tokens model, messages, temperature, max_tokens
) )
elif self.api_provider == API_PROVIDER_GEMINI:
return await self._create_gemini_completion(
model, messages, temperature, max_tokens
)
else: else:
return await self._create_openai_completion( return await self._create_openai_completion(
model, messages, temperature, max_tokens model, messages, temperature, max_tokens
@@ -238,6 +244,143 @@ class APIClient:
_LOGGER.error(f"Connection check failed: {str(e)}") _LOGGER.error(f"Connection check failed: {str(e)}")
return False 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: async def shutdown(self) -> None:
"""Shutdown API client.""" """Shutdown API client."""
_LOGGER.debug("Shutting down API client") _LOGGER.debug("Shutting down API client")
+196 -25
View File
@@ -8,6 +8,7 @@ Config flow for HA text AI integration.
""" """
import logging import logging
from typing import Any, Dict, Optional from typing import Any, Dict, Optional
from datetime import datetime, timedelta
import voluptuous as vol import voluptuous as vol
from homeassistant import config_entries from homeassistant import config_entries
@@ -29,15 +30,18 @@ from .const import (
API_PROVIDER_OPENAI, API_PROVIDER_OPENAI,
API_PROVIDER_ANTHROPIC, API_PROVIDER_ANTHROPIC,
API_PROVIDER_DEEPSEEK, API_PROVIDER_DEEPSEEK,
API_PROVIDER_GEMINI,
API_PROVIDERS, API_PROVIDERS,
DEFAULT_MODEL, DEFAULT_MODEL,
DEFAULT_DEEPSEEK_MODEL, DEFAULT_DEEPSEEK_MODEL,
DEFAULT_GEMINI_MODEL,
DEFAULT_TEMPERATURE, DEFAULT_TEMPERATURE,
DEFAULT_MAX_TOKENS, DEFAULT_MAX_TOKENS,
DEFAULT_REQUEST_INTERVAL, DEFAULT_REQUEST_INTERVAL,
DEFAULT_OPENAI_ENDPOINT, DEFAULT_OPENAI_ENDPOINT,
DEFAULT_ANTHROPIC_ENDPOINT, DEFAULT_ANTHROPIC_ENDPOINT,
DEFAULT_DEEPSEEK_ENDPOINT, DEFAULT_DEEPSEEK_ENDPOINT,
DEFAULT_GEMINI_ENDPOINT,
DEFAULT_CONTEXT_MESSAGES, DEFAULT_CONTEXT_MESSAGES,
MIN_TEMPERATURE, MIN_TEMPERATURE,
MAX_TEMPERATURE, MAX_TEMPERATURE,
@@ -93,15 +97,20 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
self._errors = {} self._errors = {}
if user_input is None: if user_input is None:
# Выбор endpoint по провайдеру # Selecting an endpoint by provider
default_endpoint = { default_endpoint = {
API_PROVIDER_OPENAI: DEFAULT_OPENAI_ENDPOINT, API_PROVIDER_OPENAI: DEFAULT_OPENAI_ENDPOINT,
API_PROVIDER_ANTHROPIC: DEFAULT_ANTHROPIC_ENDPOINT, API_PROVIDER_ANTHROPIC: DEFAULT_ANTHROPIC_ENDPOINT,
API_PROVIDER_DEEPSEEK: DEFAULT_DEEPSEEK_ENDPOINT, API_PROVIDER_DEEPSEEK: DEFAULT_DEEPSEEK_ENDPOINT,
API_PROVIDER_GEMINI: DEFAULT_GEMINI_ENDPOINT,
}.get(self._provider, DEFAULT_OPENAI_ENDPOINT) }.get(self._provider, DEFAULT_OPENAI_ENDPOINT)
# Выбор модели по умолчанию по провайдеру # Selecting the default model by provider
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
)
return self.async_show_form( return self.async_show_form(
step_id="provider", 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() 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: try:
# Validate and normalize the name
normalized_name = self._validate_and_normalize_name(input_copy[CONF_NAME]) normalized_name = self._validate_and_normalize_name(input_copy[CONF_NAME])
input_copy[CONF_NAME] = normalized_name input_copy[CONF_NAME] = normalized_name
except ValueError as e: except ValueError as e:
return self.async_show_form( return self.async_show_form(
step_id="provider", step_id="provider",
data_schema=vol.Schema({ data_schema=vol.Schema({
vol.Required(CONF_NAME, default=input_copy[CONF_NAME]): str, vol.Required(CONF_NAME, default=input_copy.get(CONF_NAME, "my_assistant")): str,
vol.Required(CONF_API_KEY, default=input_copy[CONF_API_KEY]): str, vol.Required(CONF_API_KEY, default=input_copy.get(CONF_API_KEY, "")): str,
vol.Required(CONF_MODEL, default=input_copy[CONF_MODEL]): 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[CONF_API_ENDPOINT]): 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.Optional(CONF_TEMPERATURE, default=input_copy.get(CONF_TEMPERATURE, DEFAULT_TEMPERATURE)): vol.All(
vol.Coerce(float), vol.Coerce(float),
vol.Range(min=MIN_TEMPERATURE, max=MAX_TEMPERATURE) 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)} errors={"name": str(e)}
) )
try: try:
if not await self._async_validate_api(input_copy): # Special handling for Gemini API validation
return self.async_show_form( if self._provider == API_PROVIDER_GEMINI:
step_id="provider", # For Gemini, we just check if API key is present as there's no simple endpoint to validate
data_schema=vol.Schema({}), if not input_copy.get(CONF_API_KEY):
errors=self._errors 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: except Exception as e:
# Handle any unexpected exceptions during validation
_LOGGER.exception("Unexpected error during API validation")
return self.async_show_form( return self.async_show_form(
step_id="provider", 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)} errors={"base": str(e)}
) )
# All validation passed, create the entry
return await self._create_entry(input_copy) return await self._create_entry(input_copy)
def _validate_and_normalize_name(self, name: str) -> str: 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: async def _async_validate_api(self, user_input: Dict[str, Any]) -> bool:
"""Validate API connection.""" """Validate API connection."""
try: 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) session = async_get_clientsession(self.hass)
headers = self._get_api_headers(user_input) headers = self._get_api_headers(user_input)
endpoint = user_input[CONF_API_ENDPOINT].rstrip('/') endpoint = user_input[CONF_API_ENDPOINT].rstrip('/')
check_url = ( if self._provider == API_PROVIDER_GEMINI:
f"{endpoint}/v1/models" if self._provider == API_PROVIDER_ANTHROPIC if not user_input[CONF_API_KEY]:
else f"{endpoint}/models"
)
async with session.get(check_url, headers=headers) as response:
if response.status == 401:
self._errors["base"] = "invalid_auth" self._errors["base"] = "invalid_auth"
return False return False
elif response.status not in [200, 404]:
self._errors["base"] = "cannot_connect"
return False
return True 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: except Exception as err:
_LOGGER.error("API validation error: %s", str(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]: def _get_api_headers(self, user_input: Dict[str, Any]) -> Dict[str, str]:
"""Get API headers based on provider.""" """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] api_key = user_input[CONF_API_KEY]
if self._provider == API_PROVIDER_ANTHROPIC: if self._provider == API_PROVIDER_ANTHROPIC:
@@ -244,6 +399,11 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
"anthropic-version": "2023-06-01", "anthropic-version": "2023-06-01",
"Content-Type": "application/json" "Content-Type": "application/json"
} }
elif self._provider == API_PROVIDER_GEMINI:
return {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json"
}
return { return {
"Authorization": f"Bearer {api_key}", "Authorization": f"Bearer {api_key}",
"Content-Type": "application/json" "Content-Type": "application/json"
@@ -256,7 +416,11 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
unique_id = f"{DOMAIN}_{normalized_name}_{self._provider}".lower() 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 = { entry_data = {
CONF_API_PROVIDER: self._provider, CONF_API_PROVIDER: self._provider,
@@ -304,6 +468,13 @@ class OptionsFlowHandler(config_entries.OptionsFlow):
return self.async_create_entry(title="", data=user_input) return self.async_create_entry(title="", data=user_input)
current_data = {**self.config_entry.data, **self.config_entry.options} 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( return self.async_show_form(
step_id="init", step_id="init",
+7 -1
View File
@@ -12,6 +12,8 @@ from typing import Final
import voluptuous as vol import voluptuous as vol
from homeassistant.const import Platform, CONF_API_KEY, CONF_NAME from homeassistant.const import Platform, CONF_API_KEY, CONF_NAME
from homeassistant.helpers import config_validation as cv from homeassistant.helpers import config_validation as cv
import logging
_LOGGER = logging.getLogger(__name__)
# Domain and platforms # Domain and platforms
DOMAIN: Final = "ha_text_ai" DOMAIN: Final = "ha_text_ai"
@@ -22,11 +24,13 @@ CONF_API_PROVIDER: Final = "api_provider"
API_PROVIDER_OPENAI: Final = "openai" API_PROVIDER_OPENAI: Final = "openai"
API_PROVIDER_ANTHROPIC: Final = "anthropic" API_PROVIDER_ANTHROPIC: Final = "anthropic"
API_PROVIDER_DEEPSEEK: Final = "deepseek" API_PROVIDER_DEEPSEEK: Final = "deepseek"
API_PROVIDER_GEMINI: Final = "gemini"
API_PROVIDERS: Final = [ API_PROVIDERS: Final = [
API_PROVIDER_OPENAI, API_PROVIDER_OPENAI,
API_PROVIDER_ANTHROPIC, API_PROVIDER_ANTHROPIC,
API_PROVIDER_DEEPSEEK API_PROVIDER_DEEPSEEK,
API_PROVIDER_GEMINI
] ]
# Read version from manifest.json # Read version from manifest.json
@@ -49,6 +53,7 @@ except Exception as err:
DEFAULT_OPENAI_ENDPOINT: Final = "https://api.openai.com/v1" DEFAULT_OPENAI_ENDPOINT: Final = "https://api.openai.com/v1"
DEFAULT_ANTHROPIC_ENDPOINT: Final = "https://api.anthropic.com" DEFAULT_ANTHROPIC_ENDPOINT: Final = "https://api.anthropic.com"
DEFAULT_DEEPSEEK_ENDPOINT: Final = "https://api.deepseek.com" DEFAULT_DEEPSEEK_ENDPOINT: Final = "https://api.deepseek.com"
DEFAULT_GEMINI_ENDPOINT: Final = "https://generativelanguage.googleapis.com/v1beta"
# Configuration constants # Configuration constants
CONF_MODEL: Final = "model" CONF_MODEL: Final = "model"
@@ -69,6 +74,7 @@ ICONS_SUBDOMAIN = "icons"
# Default values # Default values
DEFAULT_MODEL: Final = "gpt-4o-mini" DEFAULT_MODEL: Final = "gpt-4o-mini"
DEFAULT_DEEPSEEK_MODEL: Final = "deepseek-chat" DEFAULT_DEEPSEEK_MODEL: Final = "deepseek-chat"
DEFAULT_GEMINI_MODEL: Final = "gemini-2.0-flash"
DEFAULT_TEMPERATURE: Final = 0.1 DEFAULT_TEMPERATURE: Final = 0.1
DEFAULT_MAX_TOKENS: Final = 1000 DEFAULT_MAX_TOKENS: Final = 1000
DEFAULT_REQUEST_INTERVAL: Final = 1.0 DEFAULT_REQUEST_INTERVAL: Final = 1.0
+9 -8
View File
@@ -13,16 +13,17 @@
"loggers": ["custom_components.ha_text_ai"], "loggers": ["custom_components.ha_text_ai"],
"mqtt": [], "mqtt": [],
"quality_scale": "silver", "quality_scale": "silver",
"requirements": [ "requirements": [
"openai>=1.12.0", "openai>=1.12.0",
"anthropic>=0.8.0", "anthropic>=0.8.0",
"aiohttp>=3.8.0", "google-genai>=1.16.0",
"async-timeout>=4.0.0", "aiohttp>=3.8.0",
"certifi>=2024.2.2" "async-timeout>=4.0.0",
], "certifi>=2024.2.2"
],
"single_config_entry": false, "single_config_entry": false,
"ssdp": [], "ssdp": [],
"usb": [], "usb": [],
"version": "2.1.1", "version": "2.1.7",
"zeroconf": [] "zeroconf": []
} }
+1
View File
@@ -9,6 +9,7 @@ Sensor platform for HA Text AI.
import logging import logging
import math import math
from typing import Any, Dict from typing import Any, Dict
from datetime import datetime, timedelta
from homeassistant.components.sensor import ( from homeassistant.components.sensor import (
SensorEntity, SensorEntity,
@@ -88,12 +88,13 @@
"selector": { "selector": {
"api_provider": { "api_provider": {
"options": { "options": {
"openai": "OpenAI (kompatibel)", "openai": "OpenAI (compatible)",
"anthropic": "Anthropic (kompatibel)", "anthropic": "Anthropic (compatible)",
"deepseek": "DeepSeek" "deepseek": "DeepSeek",
} "gemini": "Google Gemini"
} }
}, }
},
"services": { "services": {
"ask_question": { "ask_question": {
"name": "Frage stellen (HA Text AI)", "name": "Frage stellen (HA Text AI)",
@@ -86,11 +86,12 @@
} }
}, },
"selector": { "selector": {
"api_provider": { "api_provider": {
"options": { "options": {
"openai": "OpenAI (compatible)", "openai": "OpenAI (compatible)",
"anthropic": "Anthropic (compatible)", "anthropic": "Anthropic (compatible)",
"deepseek": "DeepSeek" "deepseek": "DeepSeek",
"gemini": "Google Gemini"
} }
} }
}, },
@@ -90,10 +90,11 @@
"options": { "options": {
"openai": "OpenAI (compatible)", "openai": "OpenAI (compatible)",
"anthropic": "Anthropic (compatible)", "anthropic": "Anthropic (compatible)",
"deepseek": "DeepSeek" "deepseek": "DeepSeek",
} "gemini": "Google Gemini"
} }
}, }
},
"services": { "services": {
"ask_question": { "ask_question": {
"name": "Hacer Pregunta (HA Text AI)", "name": "Hacer Pregunta (HA Text AI)",
@@ -90,7 +90,8 @@
"options": { "options": {
"openai": "OpenAI (अनुकूलित)", "openai": "OpenAI (अनुकूलित)",
"anthropic": "Anthropic (अनुकूलित)", "anthropic": "Anthropic (अनुकूलित)",
"deepseek": "DeepSeek" "deepseek": "DeepSeek",
"gemini": "Google Gemini"
} }
} }
}, },
@@ -90,7 +90,8 @@
"options": { "options": {
"openai": "OpenAI (compatibile)", "openai": "OpenAI (compatibile)",
"anthropic": "Anthropic (compatibile)", "anthropic": "Anthropic (compatibile)",
"deepseek": "DeepSeek" "deepseek": "DeepSeek",
"gemini": "Google Gemini"
} }
} }
}, },
@@ -90,7 +90,8 @@
"options": { "options": {
"openai": "OpenAI (совместимый)", "openai": "OpenAI (совместимый)",
"anthropic": "Anthropic (совместимый)", "anthropic": "Anthropic (совместимый)",
"deepseek": "DeepSeek" "deepseek": "DeepSeek",
"gemini": "Google Gemini"
} }
} }
}, },
@@ -90,7 +90,8 @@
"options": { "options": {
"openai": "OpenAI (компатибилан)", "openai": "OpenAI (компатибилан)",
"anthropic": "Anthropic (компатибилан)", "anthropic": "Anthropic (компатибилан)",
"deepseek": "DeepSeek" "deepseek": "DeepSeek",
"gemini": "Google Gemini"
} }
} }
}, },
@@ -90,7 +90,8 @@
"options": { "options": {
"openai": "OpenAI(兼容)", "openai": "OpenAI(兼容)",
"anthropic": "Anthropic(兼容)", "anthropic": "Anthropic(兼容)",
"deepseek": "DeepSeek " "deepseek": "DeepSeek",
"gemini": "Google Gemini"
} }
} }
}, },
+1 -1
View File
@@ -1,5 +1,5 @@
{ {
"name": "HA text AI", "name": "HA Text AI",
"render_readme": true, "render_readme": true,
"homeassistant": "2024.11.0" "homeassistant": "2024.11.0"
} }