import asyncio import logging from collections.abc import Callable, Coroutine from typing import Any from yuxi.channel.message.models import ChunkType, ReplyPayload, ReplyStage, StreamingChunk, UnifiedMessage from yuxi.channel.message.reply_dispatcher import ReplyDispatcher from yuxi.channel.streaming.block_chunker import BlockReplyCoalescer from yuxi.channel.streaming.models import BlockReplyCoalescing logger = logging.getLogger(__name__) _DEFAULT_COALESCE_MIN_CHARS = 80 _DEFAULT_COALESCE_MAX_CHARS = 2000 _DEFAULT_COALESCE_IDLE_MS = 150.0 _DEFAULT_TIMEOUT_MS = 10_000.0 def _handle_pipeline_task_exception(task: asyncio.Task) -> None: try: task.result() except Exception: logger.exception("Unhandled exception in pipeline background task") class BlockReplyPipeline: def __init__( self, dispatcher: ReplyDispatcher, *, min_chars: int = _DEFAULT_COALESCE_MIN_CHARS, max_chars: int = _DEFAULT_COALESCE_MAX_CHARS, debounce_ms: float = _DEFAULT_COALESCE_IDLE_MS, timeout_ms: float = _DEFAULT_TIMEOUT_MS, ): self._dispatcher = dispatcher self._timeout_ms = timeout_ms self._sent_payload_keys: set[str] = set() self._sent_content_keys: set[str] = set() self._sent_media_urls: set[str] = set() self._pending_keys: set[str] = set() self._payload_seq: int = 0 self._streamed_text_fragments: list[str] = [] self._aborted = False self._did_stream = False self._did_log_timeout = False self._send_chain = asyncio.Lock() coalesce_config = BlockReplyCoalescing( min_chars=max(1, min_chars), max_chars=max(max(1, min_chars), max_chars), idle_ms=int(debounce_ms), joiner="", ) self._coalescer = BlockReplyCoalescer( config=coalesce_config, on_flush=self._on_coalesced_flush, is_stopped=lambda: self._aborted, ) async def enqueue(self, chunk: StreamingChunk) -> None: if self._aborted: return if chunk.stage == ReplyStage.TOOL: self._coalescer.flush() await self._enqueue_tool(chunk) return if chunk.stage == ReplyStage.BLOCK and chunk.chunk_type == ChunkType.TEXT_DELTA: if chunk.content: self._coalescer.enqueue(chunk.content) return if chunk.stage == ReplyStage.FINAL: self._coalescer.flush() await self._dispatcher.enqueue(chunk) return self._coalescer.flush() await self._dispatcher.enqueue(chunk) def _on_coalesced_flush(self, text: str) -> None: payload = ReplyPayload(target_id="", content=text) task = asyncio.create_task(self._send_payload(payload)) task.add_done_callback(_handle_pipeline_task_exception) async def _enqueue_tool(self, chunk: StreamingChunk) -> None: payload = ReplyPayload( target_id="", content=chunk.content or "", ) await self._send_payload(payload) async def _send_payload(self, payload: ReplyPayload) -> None: if self._aborted: return self._payload_seq += 1 dedup_key = f"{self._payload_seq}|{payload.payload_key}" content_key = payload.content_key if dedup_key in self._sent_payload_keys or dedup_key in self._pending_keys: return self._pending_keys.add(dedup_key) try: async with self._send_chain: if self._aborted: return chunk = StreamingChunk( stage=ReplyStage.BLOCK, chunk_type=ChunkType.TEXT_DELTA, content=payload.content, ) try: await asyncio.wait_for( self._dispatcher.enqueue(chunk), timeout=self._timeout_ms / 1000, ) except TimeoutError: self._aborted = True if not self._did_log_timeout: self._did_log_timeout = True logger.warning( "Block reply delivery timed out after %.0fms; " "aborting remaining replies to preserve ordering", self._timeout_ms, ) return self._sent_payload_keys.add(dedup_key) self._sent_content_keys.add(content_key) for url in payload.media_urls: self._sent_media_urls.add(url) if not payload.media_urls and payload.content.strip(): self._streamed_text_fragments.append(payload.content.strip()) self._did_stream = True finally: self._pending_keys.discard(dedup_key) async def flush(self, *, force: bool = False) -> None: if force or self._coalescer.has_buffered(): self._coalescer.flush() async with self._send_chain: pass def stop(self) -> None: self._coalescer.stop() def has_buffered(self) -> bool: return self._coalescer.has_buffered() def did_stream(self) -> bool: return self._did_stream def is_aborted(self) -> bool: return self._aborted def get_sent_media_urls(self) -> list[str]: return list(self._sent_media_urls) def has_sent_payload(self, content: str) -> bool: stripped = content.strip() if stripped in self._sent_content_keys: return True if not self._did_stream or not self._streamed_text_fragments: return False def _normalize(s: str) -> str: return "".join(s.split()) return _normalize("".join(self._streamed_text_fragments)) == _normalize(stripped) def track_media_sent(self, media_url: str) -> None: self._sent_media_urls.add(media_url) def has_sent_media(self, media_url: str) -> bool: return media_url in self._sent_media_urls @property def streamed_text(self) -> str: return "".join(self._streamed_text_fragments) @property def sent_payload_count(self) -> int: return len(self._sent_payload_keys) @property def sent_content_count(self) -> int: return len(self._sent_content_keys) @property def sent_media_count(self) -> int: return len(self._sent_media_urls) async def abort(self) -> None: self._aborted = True await self._dispatcher.abort() async def wait_idle(self, timeout_ms: float | None = None) -> None: await self._dispatcher.wait_idle(timeout_ms) def create_block_reply_pipeline( send_fn: Callable[..., Coroutine[Any, Any, str | None]], msg: UnifiedMessage, *, response_prefix: str = "", human_delay: tuple[float, float] | None = None, min_chars: int = _DEFAULT_COALESCE_MIN_CHARS, max_chars: int = _DEFAULT_COALESCE_MAX_CHARS, debounce_ms: float = _DEFAULT_COALESCE_IDLE_MS, timeout_ms: float = _DEFAULT_TIMEOUT_MS, on_idle: Callable[[], None] | None = None, on_error: Callable[[Exception], None] | None = None, before_deliver: Callable[[str], Coroutine[Any, Any, str | None]] | None = None, ) -> BlockReplyPipeline: dispatcher = ReplyDispatcher( send_fn=send_fn, msg=msg, response_prefix=response_prefix, human_delay=human_delay, timeout_ms=timeout_ms, on_idle=on_idle, on_error=on_error, before_deliver=before_deliver, ) return BlockReplyPipeline( dispatcher, min_chars=min_chars, max_chars=max_chars, debounce_ms=debounce_ms, timeout_ms=timeout_ms, )