这是一个批量整理提交,包含以下主要改动: 1. 删除多处冗余的空行和未使用的导入 2. 修复文件末尾缺少换行符的问题 3. 调整部分模块的导入顺序与代码排版 4. 修复部分配置默认值与策略逻辑 5. 新增多个功能模块与辅助工具 6. 完善异常处理与日志记录 7. 修复速率限制、消息缓存、权限校验等逻辑bug 8. 废弃部分旧有API与配置项并添加警告提示
265 lines
8.1 KiB
Python
265 lines
8.1 KiB
Python
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
|