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()