Files

165 lines
5.5 KiB
Python
Raw Permalink Normal View History

import asyncio
import logging
import os
from pathlib import Path
from collections.abc import Awaitable, Callable
from typing import Optional
from hass.config_loader import load_config
from hass.errors import HomeAssistantRequestError
from hass.ha_client import HomeAssistantClient
from hass.memory import ConversationMemory
from hass.typing import RecognizedSpeechInput
from hass.xiaoai_bridge import XiaoAiBridge
_LOGGER = logging.getLogger(__name__)
class ConversationSession:
def __init__(
self,
*,
enabled: bool,
idle_timeout_seconds: int,
end_keywords: list[str],
) -> None:
self._enabled = enabled
self._idle_timeout_seconds = max(1, int(idle_timeout_seconds))
self._end_keywords = [keyword.strip() for keyword in end_keywords if keyword and keyword.strip()]
self._idle_task: Optional[asyncio.Task] = None
def match_end_keyword(self, text: str) -> Optional[str]:
for keyword in self._end_keywords:
if keyword in text:
return keyword
return None
@staticmethod
def build_area_prefix(area: str) -> str:
return f"当前用户对话所在区域:{area},如后续对话未明确指定区域,则默认为此区域"
def resolve_area(self, client_ip: Optional[str], mapping: dict[str, str]) -> Optional[str]:
if not client_ip:
return None
return mapping.get(client_ip)
def restart_idle_timer(self, on_timeout: Callable[[], Awaitable[None]]) -> None:
if not self._enabled:
return
self.cancel_idle_timer()
async def _timeout_worker() -> None:
try:
await asyncio.sleep(self._idle_timeout_seconds)
await on_timeout()
except asyncio.CancelledError:
return
self._idle_task = asyncio.create_task(_timeout_worker())
def cancel_idle_timer(self) -> None:
if self._idle_task and not self._idle_task.done():
self._idle_task.cancel()
self._idle_task = None
async def close(self) -> None:
self.cancel_idle_timer()
async def run(config_path: str) -> None:
config = load_config(config_path)
logging.basicConfig(level=getattr(logging, config.runtime.log_level.upper(), logging.INFO))
memory = ConversationMemory(config.runtime.persist_state_path)
await memory.load()
ha = HomeAssistantClient(config.homeassistant, memory, concurrency=config.runtime.concurrency)
await ha.start()
bridge = XiaoAiBridge(config.xiaoai)
await bridge.start()
session = ConversationSession(
enabled=config.xiaoai.continuous_conversation_enabled,
idle_timeout_seconds=config.xiaoai.session_idle_timeout_seconds,
end_keywords=config.xiaoai.session_end_keywords,
)
session_end_lock = asyncio.Lock()
async def end_session(reason: str) -> None:
async with session_end_lock:
await ha.reset_conversation()
_LOGGER.info("session_ended reason=%s", reason)
async def on_idle_timeout() -> None:
await end_session("idle_timeout")
async def on_text(recognized: RecognizedSpeechInput) -> None:
text = recognized.text.strip()
if not text:
return
session.restart_idle_timer(on_idle_timeout)
matched_keyword = session.match_end_keyword(text)
if matched_keyword:
await end_session(f"keyword:{matched_keyword}")
await bridge.speak("好的,当前会话已结束")
return
area = session.resolve_area(recognized.client_ip, config.xiaoai.client_ip_area_mapping)
prefix_prompt = None
if area:
prefix_prompt = session.build_area_prefix(area)
_LOGGER.info(
"area_mapping_hit client_ip=%s area=%s",
recognized.client_ip,
area,
)
else:
_LOGGER.info("area_mapping_miss client_ip=%s", recognized.client_ip or "")
try:
if not config.xiaoai.continuous_conversation_enabled:
await end_session("continuous_disabled_new_round")
reply = await ha.process_text(text, prefix_prompt=prefix_prompt)
await bridge.speak(reply.speech)
except HomeAssistantRequestError as exc:
_LOGGER.warning(
"ha_process_failed text=%s client_ip=%s error=%s",
text,
recognized.client_ip or "",
exc,
)
if prefix_prompt:
try:
_LOGGER.info("ha_process_retry_without_prefix client_ip=%s", recognized.client_ip)
reply = await ha.process_text(text, prefix_prompt=None)
await bridge.speak(reply.speech)
return
except HomeAssistantRequestError as retry_exc:
_LOGGER.warning(
"ha_process_retry_failed client_ip=%s error=%s",
recognized.client_ip or "",
retry_exc,
)
await bridge.speak("Home Assistant 连接异常")
bridge.set_on_text(on_text)
try:
await bridge.wait_forever()
finally:
await session.close()
await ha.close()
def main() -> None:
config_path = os.environ.get("HASS_ASSISTANT_CONFIG", "config.json")
if not Path(config_path).exists():
config_path = "config.example.json"
asyncio.run(run(config_path))
if __name__ == "__main__":
main()