from __future__ import annotations import asyncio import time from dataclasses import dataclass, field from enum import StrEnum from yuxi.channels.models import ChannelIdentity, ChannelResponse, ChannelType from yuxi.utils.logging_config import logger class StreamMode(StrEnum): TEXT = "text" LANE = "lane" REASONING = "reasoning" @dataclass class StreamManager: chat_id: str send_fn: callable metadata: dict | None = None chunk_size: int = 2000 throttle_ms: int = 300 enabled: bool = True _buffer: str = field(default="", init=False) _last_flush: float = field(default=0, init=False) _sent_count: int = field(default=0, init=False) _finalized: bool = field(default=False, init=False) _lock: asyncio.Lock = field(default_factory=asyncio.Lock, init=False) def _make_response(self, content: str) -> ChannelResponse: return ChannelResponse( identity=ChannelIdentity( channel_id="yuanbao", channel_type=ChannelType.YUANBAO, channel_user_id="", channel_chat_id=self.chat_id, ), content=content, metadata=self.metadata or {}, ) async def append(self, text: str) -> None: if not self.enabled or self._finalized: return async with self._lock: self._buffer += text now = time.monotonic() elapsed_ms = (now - self._last_flush) * 1000 if len(self._buffer) >= self.chunk_size or elapsed_ms >= self.throttle_ms: await self._flush() async def finalize(self) -> int: async with self._lock: self._finalized = True if self._buffer.strip(): await self._flush() logger.debug(f"[Yuanbao] StreamManager finalized, {self._sent_count} blocks sent to {self.chat_id}") return self._sent_count async def cancel(self) -> None: async with self._lock: self._finalized = True self._buffer = "" async def _flush(self) -> None: if not self._buffer.strip(): return content = self._buffer.strip() self._buffer = "" self._last_flush = time.monotonic() response = self._make_response(content) await self.send_fn(response) self._sent_count += 1 @dataclass class LaneStreamManager: chat_id: str send_fn: callable metadata: dict | None = None chunk_size: int = 2000 throttle_ms: int = 300 _lanes: dict[str, StreamManager] = field(default_factory=dict, init=False) _lock: asyncio.Lock = field(default_factory=asyncio.Lock, init=False) async def get_lane(self, lane_id: str) -> StreamManager: async with self._lock: if lane_id not in self._lanes: self._lanes[lane_id] = StreamManager( chat_id=self.chat_id, send_fn=self.send_fn, metadata={**(self.metadata or {}), "lane_id": lane_id}, chunk_size=self.chunk_size, throttle_ms=self.throttle_ms, ) return self._lanes[lane_id] async def append_to_lane(self, lane_id: str, text: str) -> None: lane = await self.get_lane(lane_id) await lane.append(text) async def finalize_lane(self, lane_id: str) -> int: lane = await self.get_lane(lane_id) return await lane.finalize() async def finalize_all(self) -> dict[str, int]: results = {} async with self._lock: for lane_id, lane in self._lanes.items(): if not lane._finalized: results[lane_id] = await lane.finalize() return results @dataclass class ReasoningStreamManager: chat_id: str send_fn: callable metadata: dict | None = None reasoning_prefix: str = "💭 " answer_prefix: str = "" _reasoning: StreamManager | None = field(default=None, init=False) _answer: StreamManager | None = field(default=None, init=False) _lock: asyncio.Lock = field(default_factory=asyncio.Lock, init=False) async def _get_reasoning(self) -> StreamManager: if self._reasoning is None: self._reasoning = StreamManager( chat_id=self.chat_id, send_fn=self.send_fn, metadata={**(self.metadata or {}), "stream_mode": "reasoning", "role": "reasoning"}, chunk_size=1000, throttle_ms=500, ) return self._reasoning async def _get_answer(self) -> StreamManager: if self._answer is None: self._answer = StreamManager( chat_id=self.chat_id, send_fn=self.send_fn, metadata={**(self.metadata or {}), "stream_mode": "reasoning", "role": "answer"}, chunk_size=2000, throttle_ms=300, ) return self._answer async def append_reasoning(self, text: str) -> None: rm = await self._get_reasoning() await rm.append(text) async def append_answer(self, text: str) -> None: am = await self._get_answer() await am.append(text) async def finalize_reasoning(self) -> int: if self._reasoning: return await self._reasoning.finalize() return 0 async def finalize_answer(self) -> int: if self._answer: return await self._answer.finalize() return 0 async def finalize_all(self) -> dict[str, int]: result = {} async with self._lock: if self._reasoning and not self._reasoning._finalized: result["reasoning"] = await self._reasoning.finalize() if self._answer and not self._answer._finalized: result["answer"] = await self._answer.finalize() return result def create_stream_manager( chat_id: str, send_fn: callable, mode: StreamMode = StreamMode.TEXT, metadata: dict | None = None, chunk_size: int = 2000, throttle_ms: int = 300, **kwargs, ) -> StreamManager | LaneStreamManager | ReasoningStreamManager: if mode == StreamMode.LANE: return LaneStreamManager( chat_id=chat_id, send_fn=send_fn, metadata=metadata, chunk_size=chunk_size, throttle_ms=throttle_ms, ) elif mode == StreamMode.REASONING: return ReasoningStreamManager( chat_id=chat_id, send_fn=send_fn, metadata=metadata, reasoning_prefix=kwargs.get("reasoning_prefix", "💭 "), answer_prefix=kwargs.get("answer_prefix", ""), ) else: return StreamManager( chat_id=chat_id, send_fn=send_fn, metadata=metadata, chunk_size=chunk_size, throttle_ms=throttle_ms, ) async def send_blocks_stream( chat_id: str, text: str, send_fn, metadata: dict | None = None, chunk_size: int = 2000, ) -> int: paragraphs = text.split("\n\n") accumulated = "" sent_count = 0 for para in paragraphs: accumulated += para + "\n\n" if len(accumulated) >= chunk_size: response = ChannelResponse( identity=ChannelIdentity( channel_id="yuanbao", channel_type=ChannelType.YUANBAO, channel_user_id="", channel_chat_id=chat_id, ), content=accumulated.strip(), metadata=metadata or {}, ) await send_fn(response) accumulated = "" sent_count += 1 if accumulated.strip(): response = ChannelResponse( identity=ChannelIdentity( channel_id="yuanbao", channel_type=ChannelType.YUANBAO, channel_user_id="", channel_chat_id=chat_id, ), content=accumulated.strip(), metadata=metadata or {}, ) await send_fn(response) sent_count += 1 logger.debug(f"[Yuanbao] Stream sent {sent_count} blocks to {chat_id}") return sent_count