fix(avatar): prevent BOXIM polling starvation #11
@@ -163,7 +163,14 @@ def on_startup():
|
||||
boxim_client = BoxIMClient(boxim_config)
|
||||
|
||||
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")))
|
||||
takeover_scheduler = AsyncIOScheduler()
|
||||
|
||||
@@ -126,7 +126,7 @@ class TakeoverMessage(Base):
|
||||
|
||||
|
||||
class TakeoverReplyTask(Base):
|
||||
"""Restart-safe three-second BOXIM reply task."""
|
||||
"""Restart-safe delayed BOXIM reply task."""
|
||||
|
||||
__tablename__ = "takeover_reply_tasks"
|
||||
__table_args__ = (
|
||||
|
||||
@@ -25,7 +25,8 @@ logger = logging.getLogger(__name__)
|
||||
ACTIVE_TASK_STATUSES = ("pending", "generating", "ready", "sending")
|
||||
GENERATABLE_TASK_STATUSES = ("pending",)
|
||||
MAX_PROMPT_LENGTH = 4000
|
||||
MAX_STALE_SECONDS = 120
|
||||
DEFAULT_MAX_MESSAGE_AGE_SECONDS = 600
|
||||
MAX_SEND_OVERDUE_SECONDS = 120
|
||||
STUCK_LOCK_SECONDS = 90
|
||||
TAKEOVER_PERMISSION = "takeover"
|
||||
TAKEOVER_DELAY_KEY = "takeoverReplyDelaySeconds"
|
||||
@@ -115,11 +116,15 @@ class TakeoverService:
|
||||
boxim_client: BoxIMClient,
|
||||
*,
|
||||
reply_delay_seconds: int | None = None,
|
||||
poll_concurrency: int = 8,
|
||||
max_message_age_seconds: int = DEFAULT_MAX_MESSAGE_AGE_SECONDS,
|
||||
now: Callable[[], datetime] = _utcnow,
|
||||
):
|
||||
self.session_factory = session_factory
|
||||
self.boxim = boxim_client
|
||||
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._sessions: dict[str, dict] = {}
|
||||
self._poll_lock = asyncio.Lock()
|
||||
@@ -138,8 +143,48 @@ class TakeoverService:
|
||||
self._recover_stuck_tasks()
|
||||
avatar_ids = self._enabled_avatar_ids()
|
||||
self._cancel_disabled_tasks(set(avatar_ids))
|
||||
for avatar_id in avatar_ids:
|
||||
await self._sync_avatar(avatar_id)
|
||||
self._ensure_takeover_cursors(avatar_ids)
|
||||
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):
|
||||
"""Generate and send replies independently from BOXIM's long poll."""
|
||||
@@ -290,7 +335,10 @@ class TakeoverService:
|
||||
if not cursor:
|
||||
cursor = TakeoverCursor(avatar_id=avatar.id, owner_id=avatar.owner_id)
|
||||
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:
|
||||
self._record_connection_failure(
|
||||
db,
|
||||
@@ -452,16 +500,37 @@ class TakeoverService:
|
||||
return
|
||||
if not schedule_reply or event.message_type != 0 or not event.content.strip():
|
||||
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
|
||||
if is_avatar:
|
||||
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
|
||||
if self._human_pause_active(db, avatar.owner_id, peer_id, now):
|
||||
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
|
||||
if self._conversation_rate_limited(db, avatar.owner_id, peer_id, now):
|
||||
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
|
||||
self._schedule_reply(db, avatar, event)
|
||||
|
||||
@@ -540,8 +609,10 @@ class TakeoverService:
|
||||
prompt_parts.append(event.content.strip())
|
||||
source_ids.append(event.boxim_message_id)
|
||||
prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:]
|
||||
due_at = event.send_time + timedelta(
|
||||
seconds=_configured_reply_delay(avatar, self.reply_delay_seconds)
|
||||
due_at = max(
|
||||
event.send_time
|
||||
+ timedelta(seconds=_configured_reply_delay(avatar, self.reply_delay_seconds)),
|
||||
self.now(),
|
||||
)
|
||||
task_id = secrets.token_hex(16)
|
||||
local_id = _avatar_local_id(avatar.owner_id, event.boxim_message_id)
|
||||
@@ -701,7 +772,7 @@ class TakeoverService:
|
||||
task.cancel_reason = "takeover_disabled"
|
||||
db.commit()
|
||||
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.cancel_reason = "stale_reply"
|
||||
db.commit()
|
||||
|
||||
@@ -46,7 +46,12 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
|
||||
config = mock_boxim_class.call_args.args[0]
|
||||
assert config["HUIHUI_PLATFORM_BASE_URL"] == "https://open.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()
|
||||
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."""
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from threading import Barrier
|
||||
from unittest.mock import AsyncMock, patch
|
||||
@@ -62,6 +63,26 @@ class FakeBoxIM:
|
||||
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
|
||||
def service_context(tmp_path):
|
||||
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] == ["好的"]
|
||||
|
||||
|
||||
@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
|
||||
async def test_avatar_origin_message_never_schedules_a_reply(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_CLIENT_CODE=<production-client-code>
|
||||
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_CALLBACK_BASE_URL=https://digital.99hui.com
|
||||
HUIHUI_PAYMENT_CALLBACK_SECRET=<至少32位随机密钥>
|
||||
|
||||
Reference in New Issue
Block a user