fix(avatar): prevent takeover replies blocking across chats

This commit is contained in:
stefanfeng
2026-08-21 16:21:00 +08:00
parent e720baa21e
commit 672019830d
4 changed files with 98 additions and 27 deletions
+16 -2
View File
@@ -143,14 +143,28 @@ def on_startup():
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()
takeover_scheduler.add_job( takeover_scheduler.add_job(
takeover_service.poll_and_process_messages, takeover_service.poll_messages,
trigger=IntervalTrigger(seconds=poll_interval), trigger=IntervalTrigger(seconds=poll_interval),
id="takeover_message_poll", id="takeover_message_poll",
max_instances=1, max_instances=1,
coalesce=True, coalesce=True,
) )
process_interval = max(
0.25, float(os.getenv("TAKEOVER_PROCESS_INTERVAL_SECONDS", "0.5"))
)
takeover_scheduler.add_job(
takeover_service.process_reply_tasks,
trigger=IntervalTrigger(seconds=process_interval),
id="takeover_reply_process",
max_instances=1,
coalesce=True,
)
takeover_scheduler.start() takeover_scheduler.start()
logger.info("BOXIM takeover scheduler started (interval=%ss)", poll_interval) logger.info(
"BOXIM takeover scheduler started (poll=%ss, process=%ss)",
poll_interval,
process_interval,
)
except Exception as e: except Exception as e:
stop_takeover_scheduler() stop_takeover_scheduler()
logger.warning(f"Failed to initialize takeover scheduler, app will continue without it: {e}") logger.warning(f"Failed to initialize takeover scheduler, app will continue without it: {e}")
@@ -86,24 +86,34 @@ class TakeoverService:
self.reply_delay_seconds = reply_delay_seconds self.reply_delay_seconds = reply_delay_seconds
self.now = now self.now = now
self._sessions: dict[str, dict] = {} self._sessions: dict[str, dict] = {}
self._run_lock = asyncio.Lock() self._poll_lock = asyncio.Lock()
self._process_lock = asyncio.Lock()
async def poll_and_process_messages(self): async def poll_and_process_messages(self):
"""Run one complete cycle; polling always happens before reply dispatch.""" """Run one complete cycle for callers that do not use the split scheduler."""
if self._run_lock.locked(): await self.poll_messages()
await self.process_reply_tasks()
async def poll_messages(self):
"""Fetch BOXIM events without blocking reply generation and dispatch."""
if self._poll_lock.locked():
return return
async with self._run_lock: async with self._poll_lock:
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: for avatar_id in avatar_ids:
await self._sync_avatar(avatar_id) await self._sync_avatar(avatar_id)
generated = await self._prepare_replies() async def process_reply_tasks(self):
if generated: """Generate and send replies independently from BOXIM's long poll."""
# Catch a human reply sent while the model was preparing its answer. if self._process_lock.locked():
for avatar_id in avatar_ids: return
await self._sync_avatar(avatar_id) async with self._process_lock:
self._recover_stuck_tasks()
avatar_ids = set(self._enabled_avatar_ids())
self._cancel_disabled_tasks(avatar_ids)
await self._prepare_replies()
await self._dispatch_ready_replies() await self._dispatch_ready_replies()
def _enabled_avatar_ids(self) -> list[str]: def _enabled_avatar_ids(self) -> list[str]:
@@ -454,11 +464,19 @@ class TakeoverService:
finally: finally:
db.close() db.close()
generated = 0 if not task_ids:
for task_id in task_ids: return 0
if await asyncio.to_thread(self._generate_reply, task_id):
generated += 1 # Each conversation owns its task, so unrelated contacts can generate in
return generated # parallel instead of one slow model response delaying every other peer.
semaphore = asyncio.Semaphore(4)
async def generate(task_id: str) -> bool:
async with semaphore:
return await asyncio.to_thread(self._generate_reply, task_id)
results = await asyncio.gather(*(generate(task_id) for task_id in task_ids))
return sum(bool(result) for result in results)
def _generate_reply(self, task_id: str) -> bool: def _generate_reply(self, task_id: str) -> bool:
db = self.session_factory() db = self.session_factory()
@@ -548,8 +566,8 @@ class TakeoverService:
finally: finally:
db.close() db.close()
for task_id in task_ids: if task_ids:
await self._send_task(task_id) await asyncio.gather(*(self._send_task(task_id) for task_id in task_ids))
async def _send_task(self, task_id: str) -> bool: async def _send_task(self, task_id: str) -> bool:
db = self.session_factory() db = self.session_factory()
@@ -25,7 +25,8 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
boxim = MagicMock() boxim = MagicMock()
mock_boxim_class.return_value = boxim mock_boxim_class.return_value = boxim
takeover = MagicMock() takeover = MagicMock()
takeover.poll_and_process_messages = AsyncMock() takeover.poll_messages = AsyncMock()
takeover.process_reply_tasks = AsyncMock()
mock_takeover_class.return_value = takeover mock_takeover_class.return_value = takeover
environment = { environment = {
@@ -46,14 +47,18 @@ def test_scheduler_uses_boxim_and_restart_safe_service(
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)
scheduler.add_job.assert_called_once() assert scheduler.add_job.call_count == 2
scheduled_callable = scheduler.add_job.call_args.args[0] poll_call, process_call = scheduler.add_job.call_args_list
job_options = scheduler.add_job.call_args.kwargs assert poll_call.args[0] is takeover.poll_messages
assert scheduled_callable is takeover.poll_and_process_messages assert poll_call.kwargs["id"] == "takeover_message_poll"
assert job_options["id"] == "takeover_message_poll" assert poll_call.kwargs["trigger"].interval.total_seconds() == 1
assert job_options["trigger"].interval.total_seconds() == 1 assert poll_call.kwargs["max_instances"] == 1
assert job_options["max_instances"] == 1 assert poll_call.kwargs["coalesce"] is True
assert job_options["coalesce"] is True assert process_call.args[0] is takeover.process_reply_tasks
assert process_call.kwargs["id"] == "takeover_reply_process"
assert process_call.kwargs["trigger"].interval.total_seconds() == 0.5
assert process_call.kwargs["max_instances"] == 1
assert process_call.kwargs["coalesce"] is True
scheduler.start.assert_called_once_with() scheduler.start.assert_called_once_with()
main.takeover_scheduler = None main.takeover_scheduler = None
@@ -1,6 +1,7 @@
"""End-to-end service tests for BOXIM takeover timing and human priority.""" """End-to-end service tests for BOXIM takeover timing and human priority."""
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from threading import Barrier
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, patch
import pytest import pytest
@@ -143,6 +144,39 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c
db.close() db.close()
@pytest.mark.asyncio
async def test_different_contacts_generate_without_blocking_each_other(service_context):
session_factory, service, boxim, clock = service_context
await service.poll_and_process_messages()
boxim.messages.extend(
[
{"id": 13, "localId": 31, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "联系人甲"},
{"id": 14, "localId": 32, "sendId": 300, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "联系人乙"},
]
)
both_generating = Barrier(2, timeout=2)
def resolve(_db, _avatar, prompt, _history):
both_generating.wait()
return {"answer": f"回复{prompt[-1]}"}
with patch("routers.chat._resolve_reply", side_effect=resolve):
await service.poll_and_process_messages()
clock.advance(3)
await service.process_reply_tasks()
assert {(item["peerId"], item["content"]) for item in boxim.sent} == {
("200", "回复甲"),
("300", "回复乙"),
}
db = session_factory()
try:
assert {task.status for task in db.query(TakeoverReplyTask).all()} == {"sent"}
finally:
db.close()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_read_receipt_failure_does_not_advance_cursor(service_context): async def test_read_receipt_failure_does_not_advance_cursor(service_context):
session_factory, service, boxim, clock = service_context session_factory, service, boxim, clock = service_context