mirror of
https://github.com/smkrv/ha-text-ai.git
synced 2026-07-21 22:54:00 +08:00
Merge pull request #5 from Azzedde/main
Add Gemini API provider support to HA Text AI integration by @Azzedde
This commit is contained in:
@@ -20,6 +20,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 +116,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 +243,57 @@ 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."""
|
||||||
|
# Extract API key from headers (Bearer token)
|
||||||
|
api_key = self.headers.get("Authorization", "").replace("Bearer ", "")
|
||||||
|
url = f"{self.endpoint}/models/{model}:generateContent?key={api_key}"
|
||||||
|
|
||||||
|
# Convert messages to Gemini format
|
||||||
|
contents = []
|
||||||
|
system_instruction = ""
|
||||||
|
|
||||||
|
for msg in messages:
|
||||||
|
if msg['role'] == 'system':
|
||||||
|
system_instruction += msg['content'] + "\n"
|
||||||
|
else:
|
||||||
|
contents.append({
|
||||||
|
"role": "user" if msg['role'] == 'user' else "model",
|
||||||
|
"parts": [{"text": msg['content']}]
|
||||||
|
})
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"contents": contents,
|
||||||
|
"generationConfig": {
|
||||||
|
"temperature": temperature,
|
||||||
|
"maxOutputTokens": max_tokens
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if system_instruction:
|
||||||
|
payload["systemInstruction"] = {
|
||||||
|
"parts": [{"text": system_instruction}]
|
||||||
|
}
|
||||||
|
|
||||||
|
data = await self._make_request(url, payload)
|
||||||
|
return {
|
||||||
|
"choices": [{
|
||||||
|
"message": {
|
||||||
|
"content": data["candidates"][0]["content"]["parts"][0]["text"]
|
||||||
|
}
|
||||||
|
}],
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": data["usageMetadata"]["promptTokenCount"],
|
||||||
|
"completion_tokens": data["usageMetadata"]["candidatesTokenCount"],
|
||||||
|
"total_tokens": data["usageMetadata"]["totalTokenCount"]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
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")
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ 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,
|
||||||
@@ -98,10 +99,15 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN):
|
|||||||
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)
|
||||||
|
|
||||||
# Выбор модели по умолчанию по провайдеру
|
# Выбор модели по умолчанию по провайдеру
|
||||||
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",
|
||||||
|
|||||||
@@ -22,11 +22,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 +51,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 +72,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-pro"
|
||||||
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
|
||||||
|
|||||||
@@ -16,6 +16,7 @@
|
|||||||
"requirements": [
|
"requirements": [
|
||||||
"openai>=1.12.0",
|
"openai>=1.12.0",
|
||||||
"anthropic>=0.8.0",
|
"anthropic>=0.8.0",
|
||||||
|
"google-generativeai>=0.3.0",
|
||||||
"aiohttp>=3.8.0",
|
"aiohttp>=3.8.0",
|
||||||
"async-timeout>=4.0.0",
|
"async-timeout>=4.0.0",
|
||||||
"certifi>=2024.2.2"
|
"certifi>=2024.2.2"
|
||||||
|
|||||||
@@ -90,7 +90,8 @@
|
|||||||
"options": {
|
"options": {
|
||||||
"openai": "OpenAI (compatible)",
|
"openai": "OpenAI (compatible)",
|
||||||
"anthropic": "Anthropic (compatible)",
|
"anthropic": "Anthropic (compatible)",
|
||||||
"deepseek": "DeepSeek"
|
"deepseek": "DeepSeek",
|
||||||
|
"gemini": "Google Gemini"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|||||||
Reference in New Issue
Block a user