feat: voice message from bot
This commit is contained in:
120
bot.py
120
bot.py
@@ -2,8 +2,10 @@ import asyncio
|
||||
import io
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
import re
|
||||
import sqlite3
|
||||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
@@ -13,7 +15,7 @@ from aiogram.client.session.aiohttp import AiohttpSession
|
||||
from aiogram.enums import ChatAction, ParseMode
|
||||
from aiogram.exceptions import TelegramBadRequest
|
||||
from aiogram.filters import Command, CommandStart
|
||||
from aiogram.types import Message
|
||||
from aiogram.types import BufferedInputFile, Message
|
||||
from dotenv import load_dotenv
|
||||
from httpx import AsyncClient as HttpxAsyncClient
|
||||
from openai import APIConnectionError, APIError, APITimeoutError, AsyncOpenAI, NotFoundError, RateLimitError
|
||||
@@ -33,6 +35,8 @@ class Settings:
|
||||
api_token: str
|
||||
model: str
|
||||
transcription_model: str
|
||||
tts_model: str | None
|
||||
tts_voice: str
|
||||
telegram_proxy: str | None
|
||||
whitelist_user_ids: frozenset[int]
|
||||
system_prompt: str
|
||||
@@ -43,6 +47,33 @@ class Settings:
|
||||
database_path: Path
|
||||
max_context_messages: int
|
||||
message_batch_delay_seconds: float
|
||||
voice_reply_enabled: bool
|
||||
voice_reply_chance: float
|
||||
voice_reply_min_chars: int
|
||||
|
||||
|
||||
def _bool_env(name: str, default: bool) -> bool:
|
||||
value = os.getenv(name, "").strip().lower()
|
||||
if not value:
|
||||
return default
|
||||
if value in {"1", "true", "yes", "on"}:
|
||||
return True
|
||||
if value in {"0", "false", "no", "off"}:
|
||||
return False
|
||||
raise RuntimeError(f"{name} must be a boolean")
|
||||
|
||||
|
||||
def _probability_env(name: str, default: float) -> float:
|
||||
value = os.getenv(name, "").strip()
|
||||
if not value:
|
||||
return default
|
||||
try:
|
||||
parsed = float(value)
|
||||
except ValueError as exc:
|
||||
raise RuntimeError(f"{name} must be a number between 0 and 1") from exc
|
||||
if parsed < 0 or parsed > 1:
|
||||
raise RuntimeError(f"{name} must be between 0 and 1")
|
||||
return parsed
|
||||
|
||||
|
||||
def _required_env(name: str) -> str:
|
||||
@@ -126,6 +157,8 @@ def load_settings() -> Settings:
|
||||
api_token=_required_env("API_TOKEN"),
|
||||
model=_required_env("MODEL"),
|
||||
transcription_model=_required_env("TRANSCRIPTION_MODEL"),
|
||||
tts_model=os.getenv("TTS_MODEL", "").strip() or None,
|
||||
tts_voice=os.getenv("TTS_VOICE", "Kore").strip() or "Kore",
|
||||
telegram_proxy=os.getenv("TELEGRAM_PROXY", "").strip() or None,
|
||||
whitelist_user_ids=_whitelist_env(),
|
||||
system_prompt=_load_system_prompt(),
|
||||
@@ -136,6 +169,9 @@ def load_settings() -> Settings:
|
||||
database_path=Path(os.getenv("DATABASE_PATH", "data/bot.sqlite3")).expanduser(),
|
||||
max_context_messages=_int_env("MAX_CONTEXT_MESSAGES", 20),
|
||||
message_batch_delay_seconds=_float_env("MESSAGE_BATCH_DELAY_SECONDS", 3),
|
||||
voice_reply_enabled=_bool_env("VOICE_REPLY_ENABLED", False),
|
||||
voice_reply_chance=_probability_env("VOICE_REPLY_CHANCE", 0.1),
|
||||
voice_reply_min_chars=_int_env("VOICE_REPLY_MIN_CHARS", 250),
|
||||
)
|
||||
|
||||
|
||||
@@ -245,6 +281,11 @@ def voice_transcript_text(transcript: str) -> str:
|
||||
return f"[Voice message transcript]\n{transcript.strip()}"
|
||||
|
||||
|
||||
def voice_reply_text(reply: str) -> str:
|
||||
text = "\n\n".join(text for text, _ in split_reply(reply) if text).strip()
|
||||
return text or "The model returned an empty response."
|
||||
|
||||
|
||||
async def keep_typing(bot: Bot, chat_id: int, stop: asyncio.Event) -> None:
|
||||
while not stop.is_set():
|
||||
try:
|
||||
@@ -319,6 +360,66 @@ async def safe_answer(message: Message, text: str) -> None:
|
||||
await message.answer(text[: settings.max_response_chars])
|
||||
|
||||
|
||||
def should_send_voice_reply(reply: str) -> bool:
|
||||
if not settings.voice_reply_enabled or not settings.tts_model:
|
||||
return False
|
||||
plain_text = voice_reply_text(reply)
|
||||
return len(plain_text) >= settings.voice_reply_min_chars and random.random() < settings.voice_reply_chance
|
||||
|
||||
|
||||
async def synthesize_voice_reply(text: str) -> bytes:
|
||||
response = await client.audio.speech.create(
|
||||
model=settings.tts_model,
|
||||
voice=settings.tts_voice,
|
||||
input=text,
|
||||
response_format="pcm",
|
||||
)
|
||||
audio_bytes = response.read()
|
||||
if not audio_bytes:
|
||||
raise RuntimeError("TTS model returned empty audio")
|
||||
return audio_bytes
|
||||
|
||||
|
||||
async def convert_pcm_to_ogg_opus(audio_bytes: bytes) -> bytes:
|
||||
def _convert() -> bytes:
|
||||
process = subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-f",
|
||||
"s16le",
|
||||
"-ar",
|
||||
"24000",
|
||||
"-ac",
|
||||
"1",
|
||||
"-i",
|
||||
"pipe:0",
|
||||
"-c:a",
|
||||
"libopus",
|
||||
"-b:a",
|
||||
"48k",
|
||||
"-f",
|
||||
"ogg",
|
||||
"pipe:1",
|
||||
],
|
||||
input=audio_bytes,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
check=True,
|
||||
)
|
||||
if not process.stdout:
|
||||
raise RuntimeError("ffmpeg returned empty output")
|
||||
return process.stdout
|
||||
|
||||
return await asyncio.to_thread(_convert)
|
||||
|
||||
|
||||
async def send_voice_reply(message: Message, reply: str) -> None:
|
||||
voice_text = voice_reply_text(reply)
|
||||
pcm_audio = await synthesize_voice_reply(voice_text)
|
||||
ogg_audio = await convert_pcm_to_ogg_opus(pcm_audio)
|
||||
await message.answer_voice(BufferedInputFile(ogg_audio, filename="reply.ogg"))
|
||||
|
||||
|
||||
async def wait_between_reply_parts(bot: Bot, chat_id: int, timeout: int) -> None:
|
||||
if timeout <= 0:
|
||||
await bot.send_chat_action(chat_id=chat_id, action=ChatAction.TYPING)
|
||||
@@ -398,6 +499,23 @@ async def process_pending_messages(message: Message, bot: Bot, user_id: int) ->
|
||||
if should_store_reply:
|
||||
await history.append_turn(user_id, user_text, reply)
|
||||
|
||||
if should_send_voice_reply(reply):
|
||||
try:
|
||||
await send_voice_reply(message, reply)
|
||||
return
|
||||
except RateLimitError:
|
||||
logger.warning("TTS API rate limit exceeded", exc_info=True)
|
||||
except (APITimeoutError, APIConnectionError):
|
||||
logger.warning("TTS API connection problem", exc_info=True)
|
||||
except NotFoundError:
|
||||
logger.exception("TTS endpoint or model was not found")
|
||||
except APIError:
|
||||
logger.exception("TTS API returned an error")
|
||||
except subprocess.CalledProcessError:
|
||||
logger.exception("ffmpeg failed to convert synthesized audio")
|
||||
except Exception:
|
||||
logger.exception("Unexpected voice reply error")
|
||||
|
||||
await send_reply(message, bot, reply)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
|
||||
Reference in New Issue
Block a user