add banbanmini backend
This commit is contained in:
153
talkingq-url/services/text_generator.py
Normal file
153
talkingq-url/services/text_generator.py
Normal file
@@ -0,0 +1,153 @@
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Dict, List
|
||||
from utils.logger import session_logger
|
||||
from services.interrupt_handler import interrupt_handler
|
||||
from utils.text_splitter import (
|
||||
PUNCTUATION_MARKS,
|
||||
SENTENCE_ENDINGS,
|
||||
is_numbered_list_item,
|
||||
split_into_sentences,
|
||||
)
|
||||
from services.tts_config import TTSConfig
|
||||
from services.interruption_helper import InterruptionHelper
|
||||
|
||||
|
||||
class TextGenerator:
|
||||
|
||||
|
||||
def __init__(self, device_id: str, session_id: str):
|
||||
self.device_id = device_id
|
||||
self.session_id = session_id
|
||||
self.session_key = (device_id, session_id)
|
||||
|
||||
def check_interruption(self):
|
||||
|
||||
return interrupt_handler.is_interrupted(self.session_key)
|
||||
|
||||
async def generate_text(
|
||||
self,
|
||||
llm_service,
|
||||
transcript: str,
|
||||
history: List[Dict[str, str]],
|
||||
selected_role: dict,
|
||||
text_queue: asyncio.Queue,
|
||||
) -> str:
|
||||
|
||||
reply = ""
|
||||
text_buffer = ""
|
||||
first_text_token_time = None
|
||||
sent_batches_count = 0
|
||||
first_punct_sent = False
|
||||
initial_batches_count = 0
|
||||
llm_first_token_timeout = TTSConfig.get_llm_first_token_timeout()
|
||||
first_token_timeout = False
|
||||
start_time = time.perf_counter()
|
||||
|
||||
async def check_timeout():
|
||||
nonlocal first_token_timeout
|
||||
while True:
|
||||
await asyncio.sleep(0.5) # 每0.5秒检查一次
|
||||
if first_text_token_time is not None or first_token_timeout:
|
||||
break # 已收到首个令牌或已超时,停止检查
|
||||
if time.perf_counter() - start_time > llm_first_token_timeout:
|
||||
session_logger.error(
|
||||
self.device_id,
|
||||
self.session_id,
|
||||
f"LLM首个token超时({llm_first_token_timeout}秒),主动中断",
|
||||
)
|
||||
await InterruptionHelper.handle_llm_timeout(
|
||||
self.device_id, self.session_id, llm_service
|
||||
)
|
||||
break
|
||||
|
||||
try:
|
||||
token_stream = llm_service.generate_response_stream(
|
||||
transcript, history, selected_role
|
||||
)
|
||||
try:
|
||||
first_token = None
|
||||
|
||||
async def get_first_token():
|
||||
nonlocal first_token
|
||||
async for token in token_stream:
|
||||
first_token = token
|
||||
return token
|
||||
|
||||
await asyncio.wait_for(get_first_token(), llm_first_token_timeout)
|
||||
if first_token:
|
||||
first_text_token_time = time.perf_counter()
|
||||
token_latency = first_text_token_time - start_time
|
||||
session_logger.info(
|
||||
self.device_id,
|
||||
self.session_id,
|
||||
f"收到第一个文本 Token 耗时: {token_latency:.3f}秒"
|
||||
)
|
||||
if first_token is None:
|
||||
raise Exception("LLM返回的首个token为None")
|
||||
reply += first_token
|
||||
text_buffer += first_token
|
||||
async for token in token_stream:
|
||||
if self.check_interruption():
|
||||
session_logger.info(
|
||||
self.device_id, self.session_id, "LLM生成被中断"
|
||||
)
|
||||
break
|
||||
if token is None:
|
||||
continue
|
||||
reply += token
|
||||
text_buffer += token
|
||||
should_submit = False
|
||||
if initial_batches_count < 3 and any(
|
||||
p in token for p in PUNCTUATION_MARKS
|
||||
):
|
||||
if not is_numbered_list_item(text_buffer):
|
||||
should_submit = True
|
||||
elif initial_batches_count >= 3 and any(
|
||||
p in token for p in SENTENCE_ENDINGS
|
||||
):
|
||||
if not is_numbered_list_item(text_buffer):
|
||||
sentences = split_into_sentences(text_buffer)
|
||||
batch_size = 2
|
||||
if len(sentences) >= batch_size:
|
||||
to_synthesize = "".join(sentences[:batch_size])
|
||||
text_buffer = "".join(sentences[batch_size:])
|
||||
await text_queue.put(to_synthesize)
|
||||
initial_batches_count += 1
|
||||
sent_batches_count += 1
|
||||
continue
|
||||
if should_submit:
|
||||
await text_queue.put(text_buffer)
|
||||
text_buffer = ""
|
||||
if not first_punct_sent:
|
||||
first_punct_sent = True
|
||||
initial_batches_count += 1
|
||||
sent_batches_count += 1
|
||||
except asyncio.TimeoutError:
|
||||
first_token_timeout = True
|
||||
session_logger.error(
|
||||
self.device_id,
|
||||
self.session_id,
|
||||
f"LLM首个token超时({llm_first_token_timeout}秒),中断生成",
|
||||
)
|
||||
await InterruptionHelper.handle_llm_timeout(
|
||||
self.device_id, self.session_id, llm_service
|
||||
)
|
||||
await text_queue.put(None)
|
||||
return reply
|
||||
if (
|
||||
not self.check_interruption()
|
||||
and not first_token_timeout
|
||||
and text_buffer.strip()
|
||||
):
|
||||
await text_queue.put(text_buffer.strip())
|
||||
await text_queue.put(None)
|
||||
except Exception as e:
|
||||
session_logger.error(
|
||||
self.device_id, self.session_id, f"文本生成时出错: {str(e)}"
|
||||
)
|
||||
await InterruptionHelper.send_error_prompt_sound(
|
||||
self.device_id, self.session_id
|
||||
)
|
||||
await text_queue.put(None)
|
||||
return reply
|
||||
Reference in New Issue
Block a user