fix(qqofficial): prevent leading char loss in streaming buffer (#9444)

* fix(qqofficial): prevent leading char loss in streaming buffer

Copy stream deltas into an owned buffer instead of holding references
to yielded MessageChain/Plain objects. Upstream reuse/mutation dropped
the first character(s) on group (and C2C) streaming accumulation.

Add regression tests covering reference-mutation, independent deltas,
and C2C path.

* fix(qqofficial): harden stream delta copy for review feedback

- deepcopy non-Plain components to avoid shared-reference mutation
- preserve Plain.text as-is instead of coercing falsy values with or "
This commit is contained in:
王纯纯
2026-07-31 12:40:07 +08:00
committed by GitHub
parent 894ea532b0
commit 8162d84376
2 changed files with 324 additions and 10 deletions
@@ -1,5 +1,6 @@
import asyncio
import base64
import copy
import logging
import os
import random
@@ -123,11 +124,8 @@ class QQOfficialMessageEvent(AstrMessageEvent):
source = self.message_obj.raw_message
if not isinstance(source, botpy.message.C2CMessage):
# 非 C2C 场景:直接累积,最后统一发
if not self.send_buffer:
self.send_buffer = chain
else:
self.send_buffer.chain.extend(chain.chain)
# 非 C2C 场景:直接累积,最后统一发(拷贝 delta,避免引用丢首字)
self._append_stream_delta(chain)
continue
# ---- C2C 流式场景 ----
@@ -150,11 +148,8 @@ class QQOfficialMessageEvent(AstrMessageEvent):
last_edit_time = 0
continue
# 累积内容
if not self.send_buffer:
self.send_buffer = chain
else:
self.send_buffer.chain.extend(chain.chain)
# 累积内容(拷贝,避免上游复用 MessageChain 改写 buffer
self._append_stream_delta(chain)
# 节流:按时间间隔发送中间分片
current_time = asyncio.get_running_loop().time()
@@ -185,6 +180,26 @@ class QQOfficialMessageEvent(AstrMessageEvent):
return None
def _append_stream_delta(self, chain: MessageChain) -> None:
"""Append stream delta into an owned buffer (copy components).
Holding the yielded MessageChain by reference drops leading characters
when upstream reuses/mutates the same chain between yields. Non-Plain
components are deep-copied for the same reason.
"""
if not self.send_buffer:
self.send_buffer = MessageChain(
use_t2i_=chain.use_t2i_,
use_markdown_=chain.use_markdown_,
type=chain.type,
)
for comp in chain.chain:
if isinstance(comp, Plain):
# Preserve original text value (do not coerce falsy with `or ""`).
self.send_buffer.chain.append(Plain(text=comp.text))
else:
self.send_buffer.chain.append(copy.deepcopy(comp))
@staticmethod
def _extract_response_message_id(ret) -> str | None:
"""兼容 qq-botpy 返回 Message 对象或 dict 两种形态。"""
+299
View File
@@ -0,0 +1,299 @@
"""Regression tests for QQ Official streaming buffer leading-character loss.
Production logs showed group streaming dropping the first delta:
delta#1 head='' buf=''
delta#2 head='' buf='' # wrong, expected '不稀'
Root cause: send_buffer held a reference to the yielded MessageChain; upstream
reused/mutated that object. Fix: _append_stream_delta copies Plain text.
"""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock
import botpy.message
import pytest
from astrbot.api.event import MessageChain
from astrbot.api.message_components import Plain
from astrbot.api.platform import (
AstrBotMessage,
MessageMember,
MessageType,
PlatformMetadata,
)
from astrbot.core.platform.sources.qqofficial.qqofficial_message_event import (
QQOfficialMessageEvent,
)
def _extract_send_text(kwargs: dict) -> str:
text = kwargs.get("content")
if text:
return str(text)
md = kwargs.get("markdown")
if isinstance(md, dict):
return str(md.get("content") or "")
if md is not None:
return str(getattr(md, "content", None) or "")
return ""
def _make_group_event() -> QQOfficialMessageEvent:
raw = botpy.message.GroupMessage(
api=None,
event_id="event-1",
data={
"id": "msg-1",
"author": {"member_openid": "member-1"},
"group_openid": "group-1",
"content": "ping",
"timestamp": "0",
},
)
abm = AstrBotMessage()
abm.message_id = "msg-1"
abm.session_id = "group-1"
abm.group_id = "group-1"
abm.self_id = "bot-1"
abm.sender = MessageMember(user_id="member-1", nickname="u")
abm.type = MessageType.GROUP_MESSAGE
abm.message_str = "ping"
abm.message = []
abm.raw_message = raw
meta = PlatformMetadata(name="qq_official", description="t", id="qq_official")
bot = SimpleNamespace(api=SimpleNamespace(post_group_message=AsyncMock()))
return QQOfficialMessageEvent(
message_str="ping",
message_obj=abm,
platform_meta=meta,
session_id="group-1",
bot=bot, # type: ignore[arg-type]
)
def _make_c2c_event() -> QQOfficialMessageEvent:
raw = botpy.message.C2CMessage(
api=None,
event_id="event-1",
data={
"id": "msg-1",
"author": {"user_openid": "user-1"},
"content": "ping",
"timestamp": "0",
},
)
abm = AstrBotMessage()
abm.message_id = "msg-1"
abm.session_id = "user-1"
abm.self_id = "bot-1"
abm.sender = MessageMember(user_id="user-1", nickname="u")
abm.type = MessageType.FRIEND_MESSAGE
abm.message_str = "ping"
abm.message = []
abm.raw_message = raw
meta = PlatformMetadata(name="qq_official", description="t", id="qq_official")
bot = SimpleNamespace(api=SimpleNamespace())
return QQOfficialMessageEvent(
message_str="ping",
message_obj=abm,
platform_meta=meta,
session_id="user-1",
bot=bot, # type: ignore[arg-type]
)
def test_append_stream_delta_copies_plain_and_survives_source_mutation() -> None:
"""Unit-level: owned buffer must not track later mutations of the delta."""
event = _make_group_event()
shared = MessageChain(chain=[Plain("")])
event._append_stream_delta(shared)
shared.chain[0].text = "" # mutate after append
event._append_stream_delta(shared)
shared.chain[0].text = ""
event._append_stream_delta(shared)
texts = [c.text for c in event.send_buffer.chain if isinstance(c, Plain)]
assert texts == ["", "", ""]
assert "".join(texts) == "不稀罕"
def test_append_stream_delta_old_reference_style_loses_first_char() -> None:
"""Document the broken pre-fix behavior (reference assign + extend)."""
event = _make_group_event()
shared = MessageChain(chain=[Plain("")])
# Pre-fix group path:
# if not send_buffer: send_buffer = chain
# else: send_buffer.chain.extend(chain.chain)
event.send_buffer = shared
shared.chain[0].text = ""
event.send_buffer.chain.extend(shared.chain)
# After mutation + extend-on-self, leading "不" is gone.
joined = "".join(c.text for c in event.send_buffer.chain if isinstance(c, Plain))
assert "" not in joined
assert joined.startswith("")
@pytest.mark.asyncio
async def test_group_stream_keeps_first_character_when_delta_reused() -> None:
"""End-to-end group send_streaming with reused/mutated MessageChain."""
event = _make_group_event()
captured: list[str] = []
async def capture(**kwargs):
captured.append(_extract_send_text(kwargs))
return {"id": "out-1"}
event.bot.api.post_group_message = AsyncMock(side_effect=capture)
shared = MessageChain(chain=[Plain("")])
async def gen():
shared.chain[0].text = ""
yield shared
shared.chain[0].text = ""
yield shared
shared.chain[0].text = "罕?"
yield shared
await event.send_streaming(gen())
assert len(captured) == 1
assert captured[0].startswith("不稀罕?")
assert "" in captured[0]
@pytest.mark.asyncio
async def test_group_stream_accumulates_independent_delta_chains() -> None:
"""Normal path: each yield is a fresh MessageChain (openai-style deltas)."""
event = _make_group_event()
captured: list[str] = []
async def capture(**kwargs):
captured.append(_extract_send_text(kwargs))
return {"id": "out-1"}
event.bot.api.post_group_message = AsyncMock(side_effect=capture)
async def gen():
yield MessageChain().message("")
yield MessageChain().message("")
yield MessageChain().message("")
yield MessageChain().message("?认识。")
await event.send_streaming(gen())
assert len(captured) == 1
assert captured[0].startswith("不稀罕?认识。")
@pytest.mark.asyncio
async def test_group_stream_preserves_empty_and_multi_char_deltas() -> None:
event = _make_group_event()
captured: list[str] = []
async def capture(**kwargs):
captured.append(_extract_send_text(kwargs))
return {"id": "out-1"}
event.bot.api.post_group_message = AsyncMock(side_effect=capture)
async def gen():
yield MessageChain().message("你好")
yield MessageChain().message("\n\n")
yield MessageChain().message("又来了?")
await event.send_streaming(gen())
assert len(captured) == 1
assert captured[0] == "你好\n\n又来了?"
@pytest.mark.asyncio
async def test_group_stream_keeps_non_plain_components() -> None:
event = _make_group_event()
captured_kwargs: list[dict] = []
async def capture(**kwargs):
captured_kwargs.append(kwargs)
return {"id": "out-1"}
event.bot.api.post_group_message = AsyncMock(side_effect=capture)
async def gen():
yield MessageChain().message("")
# Image may force media path; still ensure text buffer kept "前缀"
yield MessageChain(chain=[Plain("")])
await event.send_streaming(gen())
assert captured_kwargs
text = _extract_send_text(captured_kwargs[0])
assert text.startswith("前缀")
@pytest.mark.asyncio
async def test_c2c_stream_append_keeps_first_char_before_throttle_flush() -> None:
"""C2C also uses _append_stream_delta; keep time <1s so only final state=10 sends."""
event = _make_c2c_event()
sent_texts: list[str] = []
async def fake_post_send(stream=None):
# Capture buffer text at send time (before _post_send clears it).
parts = []
if event.send_buffer:
for c in event.send_buffer.chain:
if isinstance(c, Plain) and c.text:
parts.append(c.text)
sent_texts.append("".join(parts))
event.send_buffer = None
return {"id": f"stream-{len(sent_texts)}"}
shared = MessageChain(chain=[Plain("")])
async def gen():
shared.chain[0].text = ""
yield shared
shared.chain[0].text = ""
yield shared
shared.chain[0].text = ""
yield shared
from unittest.mock import patch
with (
patch.object(event, "_post_send", side_effect=fake_post_send),
patch("asyncio.get_running_loop") as mock_loop,
):
# last_edit_time starts at 0; keep now < 1 so intermediate throttle never fires.
mock_loop.return_value.time.return_value = 0.5
await event.send_streaming(gen())
# Only final state=10 flush with full accumulated text.
assert len(sent_texts) == 1
assert sent_texts[0] == "不稀罕"
@pytest.mark.asyncio
async def test_group_stream_sends_once_after_all_deltas() -> None:
event = _make_group_event()
calls = 0
async def capture(**kwargs):
nonlocal calls
calls += 1
return {"id": f"out-{calls}"}
event.bot.api.post_group_message = AsyncMock(side_effect=capture)
async def gen():
for ch in "不稀罕":
yield MessageChain().message(ch)
await event.send_streaming(gen())
assert calls == 1