122 lines
3.5 KiB
Python
122 lines
3.5 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import asyncio
|
||
|
|
import logging
|
||
|
|
from collections.abc import Callable
|
||
|
|
|
||
|
|
from yuxi.channel.extensions.yuanbao.outbound.chunk import (
|
||
|
|
drain_buffer,
|
||
|
|
merge_block_streaming,
|
||
|
|
)
|
||
|
|
from yuxi.channel.extensions.yuanbao.types import MergeTextSession
|
||
|
|
|
||
|
|
logger = logging.getLogger(__name__)
|
||
|
|
|
||
|
|
|
||
|
|
class MergeTextQueue:
|
||
|
|
def __init__(
|
||
|
|
self,
|
||
|
|
send_fn: Callable[[str, str], None],
|
||
|
|
min_chars: int = 2800,
|
||
|
|
max_chars: int = 3000,
|
||
|
|
idle_ms: int = 5000,
|
||
|
|
):
|
||
|
|
self._send_fn = send_fn
|
||
|
|
|
||
|
|
self._min_chars = min_chars
|
||
|
|
self._max_chars = max_chars
|
||
|
|
self._idle_ms = idle_ms
|
||
|
|
|
||
|
|
self._sessions: dict[str, MergeTextSession] = {}
|
||
|
|
self._drain_tasks: dict[str, asyncio.Task] = {}
|
||
|
|
self._closed = False
|
||
|
|
|
||
|
|
def open_session(self, session_key: str, account_id: str, target_id: str = "") -> MergeTextSession:
|
||
|
|
if session_key in self._sessions:
|
||
|
|
session = self._sessions[session_key]
|
||
|
|
if session.closed:
|
||
|
|
session.closed = False
|
||
|
|
if target_id:
|
||
|
|
session.target_id = target_id
|
||
|
|
return session
|
||
|
|
|
||
|
|
session = MergeTextSession(
|
||
|
|
account_id=account_id,
|
||
|
|
session_key=session_key,
|
||
|
|
target_id=target_id,
|
||
|
|
min_chars=self._min_chars,
|
||
|
|
max_chars=self._max_chars,
|
||
|
|
idle_ms=self._idle_ms,
|
||
|
|
)
|
||
|
|
self._sessions[session_key] = session
|
||
|
|
return session
|
||
|
|
|
||
|
|
def push(self, session_key: str, text: str) -> None:
|
||
|
|
session = self._sessions.get(session_key)
|
||
|
|
if session is None or session.closed:
|
||
|
|
return
|
||
|
|
|
||
|
|
session.buffer = merge_block_streaming(session.buffer, text)
|
||
|
|
|
||
|
|
if len(session.buffer) >= self._min_chars:
|
||
|
|
chunks, remainder = drain_buffer(
|
||
|
|
session.buffer,
|
||
|
|
self._min_chars,
|
||
|
|
self._max_chars,
|
||
|
|
)
|
||
|
|
session.buffer = remainder
|
||
|
|
for chunk in chunks:
|
||
|
|
if chunk.strip():
|
||
|
|
self._send_fn(session.target_id, chunk)
|
||
|
|
|
||
|
|
async def drain_async(self, session_key: str) -> None:
|
||
|
|
session = self._sessions.get(session_key)
|
||
|
|
if session is None:
|
||
|
|
return
|
||
|
|
|
||
|
|
if session.draining:
|
||
|
|
return
|
||
|
|
|
||
|
|
session.draining = True
|
||
|
|
session.closed = True
|
||
|
|
buffer = session.buffer
|
||
|
|
session.buffer = ""
|
||
|
|
|
||
|
|
if buffer.strip():
|
||
|
|
self._send_fn(session.target_id, buffer)
|
||
|
|
|
||
|
|
def drain_now(self, session_key: str) -> None:
|
||
|
|
session = self._sessions.get(session_key)
|
||
|
|
if session is None or session.closed:
|
||
|
|
return
|
||
|
|
|
||
|
|
buffer = session.buffer
|
||
|
|
session.buffer = ""
|
||
|
|
if buffer.strip():
|
||
|
|
self._send_fn(session.target_id, buffer)
|
||
|
|
|
||
|
|
def flush(self, session_key: str) -> None:
|
||
|
|
session = self._sessions.get(session_key)
|
||
|
|
if session is None:
|
||
|
|
return
|
||
|
|
buffer = session.buffer
|
||
|
|
session.closed = True
|
||
|
|
session.buffer = ""
|
||
|
|
if buffer.strip():
|
||
|
|
self._send_fn(session.target_id, buffer)
|
||
|
|
self._sessions.pop(session_key, None)
|
||
|
|
|
||
|
|
def abort(self, session_key: str) -> None:
|
||
|
|
session = self._sessions.get(session_key)
|
||
|
|
if session is None:
|
||
|
|
return
|
||
|
|
session.buffer = ""
|
||
|
|
session.closed = True
|
||
|
|
self._sessions.pop(session_key, None)
|
||
|
|
|
||
|
|
def close(self) -> None:
|
||
|
|
self._closed = True
|
||
|
|
for session in list(self._sessions.values()):
|
||
|
|
session.closed = True
|
||
|
|
self._sessions.clear()
|