feat: audio transcribing
This commit is contained in:
100
bot.py
100
bot.py
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import io
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
@@ -31,6 +32,7 @@ class Settings:
|
||||
api_base_url: str
|
||||
api_token: str
|
||||
model: str
|
||||
transcription_model: str
|
||||
telegram_proxy: str | None
|
||||
whitelist_user_ids: frozenset[int]
|
||||
system_prompt: str
|
||||
@@ -123,6 +125,7 @@ def load_settings() -> Settings:
|
||||
api_base_url=_api_base_url_env(),
|
||||
api_token=_required_env("API_TOKEN"),
|
||||
model=_required_env("MODEL"),
|
||||
transcription_model=_required_env("TRANSCRIPTION_MODEL"),
|
||||
telegram_proxy=os.getenv("TELEGRAM_PROXY", "").strip() or None,
|
||||
whitelist_user_ids=_whitelist_env(),
|
||||
system_prompt=_load_system_prompt(),
|
||||
@@ -238,6 +241,10 @@ def combine_user_messages(messages: list[str]) -> str:
|
||||
return "\n".join(f"Message {index}: {text}" for index, text in enumerate(messages, start=1))
|
||||
|
||||
|
||||
def voice_transcript_text(transcript: str) -> str:
|
||||
return f"[Voice message transcript]\n{transcript.strip()}"
|
||||
|
||||
|
||||
async def keep_typing(bot: Bot, chat_id: int, stop: asyncio.Event) -> None:
|
||||
while not stop.is_set():
|
||||
try:
|
||||
@@ -280,6 +287,31 @@ async def stream_model_reply(user_text: str, context: list[dict[str, str]]) -> s
|
||||
return "".join(chunks).strip()
|
||||
|
||||
|
||||
async def download_voice_message(bot: Bot, message: Message) -> tuple[bytes, str]:
|
||||
if not message.voice:
|
||||
raise ValueError("Message does not contain a voice attachment")
|
||||
|
||||
file = await bot.get_file(message.voice.file_id)
|
||||
if not file.file_path:
|
||||
raise RuntimeError("Telegram did not return a file path for the voice message")
|
||||
|
||||
buffer = io.BytesIO()
|
||||
await bot.download_file(file.file_path, destination=buffer)
|
||||
return buffer.getvalue(), Path(file.file_path).name
|
||||
|
||||
|
||||
async def transcribe_voice_message(bot: Bot, message: Message) -> str:
|
||||
voice_bytes, filename = await download_voice_message(bot, message)
|
||||
transcription = await client.audio.transcriptions.create(
|
||||
model=settings.transcription_model,
|
||||
file=(filename, voice_bytes, "audio/ogg"),
|
||||
)
|
||||
transcript = transcription.text.strip()
|
||||
if not transcript:
|
||||
raise RuntimeError("Voice message transcription was empty")
|
||||
return transcript
|
||||
|
||||
|
||||
async def safe_answer(message: Message, text: str) -> None:
|
||||
try:
|
||||
await message.answer(text)
|
||||
@@ -371,12 +403,27 @@ async def process_pending_messages(message: Message, bot: Bot, user_id: int) ->
|
||||
raise
|
||||
|
||||
|
||||
async def queue_user_message(message: Message, bot: Bot, user_text: str) -> None:
|
||||
if not message.from_user:
|
||||
await message.answer("Cannot identify Telegram user.")
|
||||
return
|
||||
|
||||
async with user_lock(message.from_user.id):
|
||||
pending_user_messages.setdefault(message.from_user.id, []).append(user_text)
|
||||
task = pending_tasks.get(message.from_user.id)
|
||||
if task:
|
||||
task.cancel()
|
||||
pending_tasks[message.from_user.id] = asyncio.create_task(
|
||||
process_pending_messages(message, bot, message.from_user.id)
|
||||
)
|
||||
|
||||
|
||||
@router.message(CommandStart())
|
||||
async def start(message: Message) -> None:
|
||||
if not is_allowed(message):
|
||||
await message.answer("Access denied.")
|
||||
return
|
||||
await message.answer("Send a message and I will forward it to the configured model.")
|
||||
await message.answer("Send a text or voice message and I will forward it to the configured model.")
|
||||
|
||||
|
||||
@router.message(Command("new"))
|
||||
@@ -407,14 +454,51 @@ async def handle_text(message: Message, bot: Bot) -> None:
|
||||
await message.answer("Only text messages are supported.")
|
||||
return
|
||||
|
||||
async with user_lock(message.from_user.id):
|
||||
pending_user_messages.setdefault(message.from_user.id, []).append(message.text)
|
||||
task = pending_tasks.get(message.from_user.id)
|
||||
if task:
|
||||
task.cancel()
|
||||
pending_tasks[message.from_user.id] = asyncio.create_task(
|
||||
process_pending_messages(message, bot, message.from_user.id)
|
||||
await queue_user_message(message, bot, message.text)
|
||||
|
||||
|
||||
@router.message(F.voice)
|
||||
async def handle_voice(message: Message, bot: Bot) -> None:
|
||||
if not is_allowed(message):
|
||||
await message.answer("Access denied.")
|
||||
return
|
||||
|
||||
if not message.voice:
|
||||
await message.answer("Only Telegram voice messages are supported.")
|
||||
return
|
||||
|
||||
logger.info("Transcribing voice message for user_id=%s", message.from_user.id)
|
||||
|
||||
try:
|
||||
transcript = await transcribe_voice_message(bot, message)
|
||||
except RateLimitError:
|
||||
logger.warning("Transcription API rate limit exceeded", exc_info=True)
|
||||
await message.answer("The transcription model provider is rate-limiting requests. Please try again later.")
|
||||
return
|
||||
except (APITimeoutError, APIConnectionError):
|
||||
logger.warning("Transcription API connection problem", exc_info=True)
|
||||
await message.answer("The transcription model provider is temporarily unreachable. Please try again later.")
|
||||
return
|
||||
except NotFoundError:
|
||||
logger.exception("Transcription endpoint or model was not found")
|
||||
await message.answer(
|
||||
"The transcription model provider returned 404. Check API_BASE_URL and TRANSCRIPTION_MODEL configuration."
|
||||
)
|
||||
return
|
||||
except APIError:
|
||||
logger.exception("Transcription API returned an error")
|
||||
await message.answer("The transcription model provider returned an error. Please try again later.")
|
||||
return
|
||||
except Exception:
|
||||
logger.exception("Unexpected voice transcription error")
|
||||
await message.answer("Unexpected error while transcribing the voice message. Please try again later.")
|
||||
return
|
||||
|
||||
transcript_message = voice_transcript_text(transcript)
|
||||
logger.info("Voice message transcribed for user_id=%s", message.from_user.id)
|
||||
|
||||
# Telegram bots cannot mark a voice note as listened in the client UI.
|
||||
await queue_user_message(message, bot, transcript_message)
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
|
||||
Reference in New Issue
Block a user