diff --git a/custom_components/ha_text_ai/api_client.py b/custom_components/ha_text_ai/api_client.py index adeee68..92fa12d 100644 --- a/custom_components/ha_text_ai/api_client.py +++ b/custom_components/ha_text_ai/api_client.py @@ -20,6 +20,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 +116,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 +243,57 @@ 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.""" + # 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: """Shutdown API client.""" _LOGGER.debug("Shutting down API client") diff --git a/custom_components/ha_text_ai/config_flow.py b/custom_components/ha_text_ai/config_flow.py index a263520..945a0e7 100644 --- a/custom_components/ha_text_ai/config_flow.py +++ b/custom_components/ha_text_ai/config_flow.py @@ -29,6 +29,7 @@ from .const import ( API_PROVIDER_OPENAI, API_PROVIDER_ANTHROPIC, API_PROVIDER_DEEPSEEK, + API_PROVIDER_GEMINI, API_PROVIDERS, DEFAULT_MODEL, DEFAULT_DEEPSEEK_MODEL, @@ -98,10 +99,15 @@ class HATextAIConfigFlow(config_entries.ConfigFlow, domain=DOMAIN): 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 + 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", diff --git a/custom_components/ha_text_ai/const.py b/custom_components/ha_text_ai/const.py index 640670e..2f90eeb 100644 --- a/custom_components/ha_text_ai/const.py +++ b/custom_components/ha_text_ai/const.py @@ -22,11 +22,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 +51,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 +72,7 @@ ICONS_SUBDOMAIN = "icons" # Default values DEFAULT_MODEL: Final = "gpt-4o-mini" DEFAULT_DEEPSEEK_MODEL: Final = "deepseek-chat" +DEFAULT_GEMINI_MODEL: Final = "gemini-pro" DEFAULT_TEMPERATURE: Final = 0.1 DEFAULT_MAX_TOKENS: Final = 1000 DEFAULT_REQUEST_INTERVAL: Final = 1.0 diff --git a/custom_components/ha_text_ai/manifest.json b/custom_components/ha_text_ai/manifest.json index 4f08669..1c2c884 100644 --- a/custom_components/ha_text_ai/manifest.json +++ b/custom_components/ha_text_ai/manifest.json @@ -16,6 +16,7 @@ "requirements": [ "openai>=1.12.0", "anthropic>=0.8.0", + "google-generativeai>=0.3.0", "aiohttp>=3.8.0", "async-timeout>=4.0.0", "certifi>=2024.2.2" diff --git a/custom_components/ha_text_ai/translations/en.json b/custom_components/ha_text_ai/translations/en.json index 6f4dfc9..4ec4230 100644 --- a/custom_components/ha_text_ai/translations/en.json +++ b/custom_components/ha_text_ai/translations/en.json @@ -90,7 +90,8 @@ "options": { "openai": "OpenAI (compatible)", "anthropic": "Anthropic (compatible)", - "deepseek": "DeepSeek" + "deepseek": "DeepSeek", + "gemini": "Google Gemini" } } },