Merge pull request 'fix(avatar): prevent BOXIM polling starvation' (#11) from codex/avatar-takeover-poll-20260901 into main

This commit was merged in pull request #11.
This commit is contained in:
2026-09-01 14:04:18 +08:00
6 changed files with 196 additions and 11 deletions
+8 -1
View File
@@ -163,7 +163,14 @@ def on_startup():
boxim_client = BoxIMClient(boxim_config) boxim_client = BoxIMClient(boxim_config)
from services.takeover_service import TakeoverService from services.takeover_service import TakeoverService
takeover_service = TakeoverService(SessionLocal, boxim_client) takeover_service = TakeoverService(
SessionLocal,
boxim_client,
poll_concurrency=int(os.getenv("BOXIM_POLL_CONCURRENCY", "8")),
max_message_age_seconds=int(
os.getenv("BOXIM_MAX_MESSAGE_AGE_SECONDS", "600")
),
)
poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1"))) poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1")))
takeover_scheduler = AsyncIOScheduler() takeover_scheduler = AsyncIOScheduler()
+1 -1
View File
@@ -126,7 +126,7 @@ class TakeoverMessage(Base):
class TakeoverReplyTask(Base): class TakeoverReplyTask(Base):
"""Restart-safe three-second BOXIM reply task.""" """Restart-safe delayed BOXIM reply task."""
__tablename__ = "takeover_reply_tasks" __tablename__ = "takeover_reply_tasks"
__table_args__ = ( __table_args__ = (
@@ -25,7 +25,8 @@ logger = logging.getLogger(__name__)
ACTIVE_TASK_STATUSES = ("pending", "generating", "ready", "sending") ACTIVE_TASK_STATUSES = ("pending", "generating", "ready", "sending")
GENERATABLE_TASK_STATUSES = ("pending",) GENERATABLE_TASK_STATUSES = ("pending",)
MAX_PROMPT_LENGTH = 4000 MAX_PROMPT_LENGTH = 4000
MAX_STALE_SECONDS = 120 DEFAULT_MAX_MESSAGE_AGE_SECONDS = 600
MAX_SEND_OVERDUE_SECONDS = 120
STUCK_LOCK_SECONDS = 90 STUCK_LOCK_SECONDS = 90
TAKEOVER_PERMISSION = "takeover" TAKEOVER_PERMISSION = "takeover"
TAKEOVER_DELAY_KEY = "takeoverReplyDelaySeconds" TAKEOVER_DELAY_KEY = "takeoverReplyDelaySeconds"
@@ -115,11 +116,15 @@ class TakeoverService:
boxim_client: BoxIMClient, boxim_client: BoxIMClient,
*, *,
reply_delay_seconds: int | None = None, reply_delay_seconds: int | None = None,
poll_concurrency: int = 8,
max_message_age_seconds: int = DEFAULT_MAX_MESSAGE_AGE_SECONDS,
now: Callable[[], datetime] = _utcnow, now: Callable[[], datetime] = _utcnow,
): ):
self.session_factory = session_factory self.session_factory = session_factory
self.boxim = boxim_client self.boxim = boxim_client
self.reply_delay_seconds = reply_delay_seconds self.reply_delay_seconds = reply_delay_seconds
self.poll_concurrency = max(1, min(int(poll_concurrency), 64))
self.max_message_age_seconds = max(60, int(max_message_age_seconds))
self.now = now self.now = now
self._sessions: dict[str, dict] = {} self._sessions: dict[str, dict] = {}
self._poll_lock = asyncio.Lock() self._poll_lock = asyncio.Lock()
@@ -138,8 +143,48 @@ class TakeoverService:
self._recover_stuck_tasks() self._recover_stuck_tasks()
avatar_ids = self._enabled_avatar_ids() avatar_ids = self._enabled_avatar_ids()
self._cancel_disabled_tasks(set(avatar_ids)) self._cancel_disabled_tasks(set(avatar_ids))
for avatar_id in avatar_ids: self._ensure_takeover_cursors(avatar_ids)
await self._sync_avatar(avatar_id) semaphore = asyncio.Semaphore(self.poll_concurrency)
async def sync(avatar_id: str):
async with semaphore:
return await self._sync_avatar(avatar_id)
results = await asyncio.gather(
*(sync(avatar_id) for avatar_id in avatar_ids),
return_exceptions=True,
)
for avatar_id, result in zip(avatar_ids, results):
if isinstance(result, Exception):
logger.warning("BOXIM poll crashed for avatar %s: %s", avatar_id, result)
def _ensure_takeover_cursors(self, avatar_ids: list[str]):
"""Create durable cursors before concurrent network polling starts."""
if not avatar_ids:
return
db = self.session_factory()
try:
existing = {
row[0]
for row in db.query(TakeoverCursor.avatar_id)
.filter(TakeoverCursor.avatar_id.in_(avatar_ids))
.all()
}
avatars = (
db.query(Avatar.id, Avatar.owner_id)
.filter(
Avatar.id.in_(
[avatar_id for avatar_id in avatar_ids if avatar_id not in existing]
)
)
.all()
)
for avatar_id, owner_id in avatars:
db.add(TakeoverCursor(avatar_id=avatar_id, owner_id=owner_id))
if avatars:
db.commit()
finally:
db.close()
async def process_reply_tasks(self): async def process_reply_tasks(self):
"""Generate and send replies independently from BOXIM's long poll.""" """Generate and send replies independently from BOXIM's long poll."""
@@ -290,7 +335,10 @@ class TakeoverService:
if not cursor: if not cursor:
cursor = TakeoverCursor(avatar_id=avatar.id, owner_id=avatar.owner_id) cursor = TakeoverCursor(avatar_id=avatar.id, owner_id=avatar.owner_id)
db.add(cursor) db.add(cursor)
db.flush() db.commit()
else:
# Release SQLite's read transaction before the long network poll.
db.commit()
if not user or not user.huihui_token: if not user or not user.huihui_token:
self._record_connection_failure( self._record_connection_failure(
db, db,
@@ -452,16 +500,37 @@ class TakeoverService:
return return
if not schedule_reply or event.message_type != 0 or not event.content.strip(): if not schedule_reply or event.message_type != 0 or not event.content.strip():
return return
if (now - send_time).total_seconds() > MAX_STALE_SECONDS: if (now - send_time).total_seconds() > self.max_message_age_seconds:
logger.info(
"Ignored stale BOXIM message %s for avatar %s (age=%ss)",
message_id,
avatar.id,
int((now - send_time).total_seconds()),
)
return return
if is_avatar: if is_avatar:
self._cancel_conversation(db, avatar.owner_id, peer_id, "peer_avatar_message") self._cancel_conversation(db, avatar.owner_id, peer_id, "peer_avatar_message")
logger.info(
"Skipped BOXIM reply for avatar %s message %s: peer_avatar_message",
avatar.id,
message_id,
)
return return
if self._human_pause_active(db, avatar.owner_id, peer_id, now): if self._human_pause_active(db, avatar.owner_id, peer_id, now):
self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_active") self._cancel_conversation(db, avatar.owner_id, peer_id, "owner_active")
logger.info(
"Skipped BOXIM reply for avatar %s message %s: owner_active",
avatar.id,
message_id,
)
return return
if self._conversation_rate_limited(db, avatar.owner_id, peer_id, now): if self._conversation_rate_limited(db, avatar.owner_id, peer_id, now):
self._cancel_conversation(db, avatar.owner_id, peer_id, "rate_limited") self._cancel_conversation(db, avatar.owner_id, peer_id, "rate_limited")
logger.info(
"Skipped BOXIM reply for avatar %s message %s: rate_limited",
avatar.id,
message_id,
)
return return
self._schedule_reply(db, avatar, event) self._schedule_reply(db, avatar, event)
@@ -540,8 +609,10 @@ class TakeoverService:
prompt_parts.append(event.content.strip()) prompt_parts.append(event.content.strip())
source_ids.append(event.boxim_message_id) source_ids.append(event.boxim_message_id)
prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:] prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:]
due_at = event.send_time + timedelta( due_at = max(
seconds=_configured_reply_delay(avatar, self.reply_delay_seconds) event.send_time
+ timedelta(seconds=_configured_reply_delay(avatar, self.reply_delay_seconds)),
self.now(),
) )
task_id = secrets.token_hex(16) task_id = secrets.token_hex(16)
local_id = _avatar_local_id(avatar.owner_id, event.boxim_message_id) local_id = _avatar_local_id(avatar.owner_id, event.boxim_message_id)
@@ -701,7 +772,7 @@ class TakeoverService:
task.cancel_reason = "takeover_disabled" task.cancel_reason = "takeover_disabled"
db.commit() db.commit()
return False return False
if (self.now() - task.scheduled_at).total_seconds() > MAX_STALE_SECONDS: if (self.now() - task.scheduled_at).total_seconds() > MAX_SEND_OVERDUE_SECONDS:
task.status = "cancelled" task.status = "cancelled"
task.cancel_reason = "stale_reply" task.cancel_reason = "stale_reply"
db.commit() db.commit()
@@ -46,7 +46,12 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
config = mock_boxim_class.call_args.args[0] config = mock_boxim_class.call_args.args[0]
assert config["HUIHUI_PLATFORM_BASE_URL"] == "https://open.example/api" assert config["HUIHUI_PLATFORM_BASE_URL"] == "https://open.example/api"
assert config["BOXIM_API_BASE_URL"] == "https://im.example/api" assert config["BOXIM_API_BASE_URL"] == "https://im.example/api"
mock_takeover_class.assert_called_once_with(main.SessionLocal, boxim) mock_takeover_class.assert_called_once_with(
main.SessionLocal,
boxim,
poll_concurrency=8,
max_message_age_seconds=600,
)
maintenance_scheduler.add_job.assert_called_once() maintenance_scheduler.add_job.assert_called_once()
assert maintenance_scheduler.add_job.call_args.kwargs["id"] == "chat_attachment_cleanup" assert maintenance_scheduler.add_job.call_args.kwargs["id"] == "chat_attachment_cleanup"
@@ -1,5 +1,6 @@
"""End-to-end service tests for BOXIM takeover timing and human priority.""" """End-to-end service tests for BOXIM takeover timing and human priority."""
import asyncio
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from threading import Barrier from threading import Barrier
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, patch
@@ -62,6 +63,26 @@ class FakeBoxIM:
return {"id": 900 + len(self.sent), "localId": int(local_id)} return {"id": 900 + len(self.sent), "localId": int(local_id)}
class ConcurrentPollingBoxIM(FakeBoxIM):
def __init__(self):
super().__init__()
self.active_polls = 0
self.peak_active_polls = 0
async def exchange_access_token(self, huihui_token):
return {"accessToken": huihui_token, "accessTokenExpiresIn": 3600}
async def get_self(self, access_token):
return {"id": 100 if access_token == "prod-huihui-token" else 101}
async def fetch_private_messages(self, access_token, min_id="0"):
self.active_polls += 1
self.peak_active_polls = max(self.peak_active_polls, self.active_polls)
await asyncio.sleep(0.05)
self.active_polls -= 1
return []
@pytest.fixture @pytest.fixture
def service_context(tmp_path): def service_context(tmp_path):
engine = create_engine( engine = create_engine(
@@ -183,6 +204,85 @@ async def test_default_reply_delay_is_three_minutes(service_context):
assert [item["content"] for item in boxim.sent] == ["好的"] assert [item["content"] for item in boxim.sent] == ["好的"]
@pytest.mark.asyncio
async def test_multiple_avatar_owners_are_polled_concurrently(service_context):
session_factory, _service, _boxim, clock = service_context
db = session_factory()
try:
db.add_all(
[
User(
id="owner-local-2",
huihui_user_id="owner-huihui-2",
huihui_token="prod-huihui-token-2",
app_token="app-token-2",
),
Avatar(
id="avatar-2",
owner_id="owner-huihui-2",
name="分身二",
status="active",
config={"authorizationPermissions": ["chat", "takeover"]},
),
]
)
db.commit()
finally:
db.close()
boxim = ConcurrentPollingBoxIM()
service = TakeoverService(
session_factory,
boxim,
poll_concurrency=2,
now=clock.now,
)
await service.poll_messages()
assert boxim.peak_active_polls == 2
db = session_factory()
try:
assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2
finally:
db.close()
@pytest.mark.asyncio
async def test_delayed_poll_still_schedules_recent_message(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_messages()
delayed_send_time = int(
(clock.value - timedelta(seconds=150)).replace(tzinfo=timezone.utc).timestamp()
* 1000
)
boxim.messages.append(
{
"id": 13,
"localId": 13,
"sendId": 200,
"recvId": 100,
"sendTime": delayed_send_time,
"type": 0,
"content": "排队后仍需回复",
}
)
await service.poll_messages()
db = session_factory()
try:
task = db.query(TakeoverReplyTask).one()
assert task.status == "pending"
assert task.scheduled_at == clock.now()
finally:
db.close()
with patch("routers.chat._resolve_reply", return_value={"answer": "已经收到"}):
await service.process_reply_tasks()
assert [item["content"] for item in boxim.sent] == ["已经收到"]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_avatar_origin_message_never_schedules_a_reply(service_context): async def test_avatar_origin_message_never_schedules_a_reply(service_context):
session_factory, service, boxim, clock = service_context session_factory, service, boxim, clock = service_context
@@ -39,6 +39,8 @@ HUIHUI_ACCESS_ID=<production-access-id>
HUIHUI_ACCESS_SECRET=<production-access-secret> HUIHUI_ACCESS_SECRET=<production-access-secret>
HUIHUI_CLIENT_CODE=<production-client-code> HUIHUI_CLIENT_CODE=<production-client-code>
BOXIM_TIMEOUT_SECONDS=20 BOXIM_TIMEOUT_SECONDS=20
BOXIM_POLL_CONCURRENCY=8
BOXIM_MAX_MESSAGE_AGE_SECONDS=600
HUIHUI_PAYMENT_BASE_URL=https://open.99hui.com/api/payment-v3 HUIHUI_PAYMENT_BASE_URL=https://open.99hui.com/api/payment-v3
HUIHUI_PAYMENT_CALLBACK_BASE_URL=https://digital.99hui.com HUIHUI_PAYMENT_CALLBACK_BASE_URL=https://digital.99hui.com
HUIHUI_PAYMENT_CALLBACK_SECRET=<至少32位随机密钥> HUIHUI_PAYMENT_CALLBACK_SECRET=<至少32位随机密钥>