From 7cac96356dfc4bc222eef75646396aa242c15ae6 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Wed, 2 Sep 2026 14:12:04 +0800 Subject: [PATCH] fix(avatar): recover BOXIM image replies --- digital-avatar-app/backend/database.py | 22 +++- .../backend/services/takeover_service.py | 114 ++++++++++++++---- .../backend/tests/test_takeover_service.py | 105 +++++++++++++++- 3 files changed, 216 insertions(+), 25 deletions(-) diff --git a/digital-avatar-app/backend/database.py b/digital-avatar-app/backend/database.py index 9884cf1..9937efa 100644 --- a/digital-avatar-app/backend/database.py +++ b/digital-avatar-app/backend/database.py @@ -1,16 +1,30 @@ import os -from sqlalchemy import create_engine +from sqlalchemy import create_engine, event from sqlalchemy.orm import sessionmaker, declarative_base, Session BASE_DIR = os.path.dirname(os.path.abspath(__file__)) DB_FILE = os.path.join(BASE_DIR, "avatar.db") DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}") +IS_SQLITE = DATABASE_URL.startswith("sqlite:") engine = create_engine( DATABASE_URL, - connect_args={"check_same_thread": False} if DATABASE_URL.startswith("sqlite:") else {}, + connect_args={"check_same_thread": False, "timeout": 30} if IS_SQLITE else {}, ) + + +if IS_SQLITE: + @event.listens_for(engine, "connect") + def _configure_sqlite_connection(dbapi_connection, _connection_record): + cursor = dbapi_connection.cursor() + try: + cursor.execute("PRAGMA synchronous=NORMAL") + cursor.execute("PRAGMA busy_timeout=30000") + finally: + cursor.close() + + SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False) Base = declarative_base() @@ -26,6 +40,10 @@ def get_db(): def init_db(): import models + if IS_SQLITE: + with engine.connect() as conn: + conn.exec_driver_sql("PRAGMA journal_mode=WAL") + conn.commit() Base.metadata.create_all(bind=engine) # 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试) diff --git a/digital-avatar-app/backend/services/takeover_service.py b/digital-avatar-app/backend/services/takeover_service.py index d15da8e..47c1600 100644 --- a/digital-avatar-app/backend/services/takeover_service.py +++ b/digital-avatar-app/backend/services/takeover_service.py @@ -49,6 +49,8 @@ BOXIM_TEXT_MESSAGE_TYPE = 0 BOXIM_IMAGE_MESSAGE_TYPE = 1 BOXIM_IMAGE_PROMPT = "请看看这张图片。" BOXIM_IMAGE_UNAVAILABLE_REPLY = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。" +IMAGE_CONTEXT_LOOKBACK_SECONDS = 1800 +MAX_RECENT_IMAGE_CONTEXTS = 3 def _utcnow() -> datetime: @@ -149,6 +151,7 @@ class TakeoverService: self._sessions: dict[str, dict] = {} self._poll_lock = asyncio.Lock() self._process_lock = asyncio.Lock() + self._persist_lock = asyncio.Lock() async def poll_and_process_messages(self): """Run one complete cycle for callers that do not use the split scheduler.""" @@ -409,13 +412,6 @@ class TakeoverService: max_message_id = _numeric_id(cursor.last_message_id) read_receipts: dict[str, int] = {} for message in messages: - self._record_message( - db, - avatar, - cursor.boxim_owner_id, - message, - schedule_reply=not priming, - ) message_id = _numeric_id(message.get("id")) max_message_id = max(max_message_id, message_id) send_id = str(message.get("sendId") or "") @@ -430,11 +426,22 @@ class TakeoverService: session["access_token"], peer_id, message_id ) - cursor.last_message_id = str(max_message_id) - cursor.initialized = True - cursor.last_polled_at = self.now() - cursor.last_error = "" - db.commit() + # Keep SQLite write transactions short. The read-receipt request above + # can block on the network and must not hold the database write lock. + async with self._persist_lock: + for message in messages: + self._record_message( + db, + avatar, + cursor.boxim_owner_id, + message, + schedule_reply=not priming, + ) + cursor.last_message_id = str(max_message_id) + cursor.initialized = True + cursor.last_polled_at = self.now() + cursor.last_error = "" + db.commit() return True except Exception: db.rollback() @@ -645,6 +652,15 @@ class TakeoverService: task.status = "cancelled" task.cancel_reason = "newer_incoming_message" task.locked_at = None + if event.message_type == BOXIM_TEXT_MESSAGE_TYPE: + for image_event in self._recent_unhandled_images( + db, + avatar, + event, + source_ids, + ): + prompt_parts.append(_event_prompt(image_event)) + source_ids.append(image_event.boxim_message_id) prompt_parts.append(_event_prompt(event)) source_ids.append(event.boxim_message_id) prompt = "\n".join(part for part in prompt_parts if part).strip()[-MAX_PROMPT_LENGTH:] @@ -670,6 +686,54 @@ class TakeoverService: ) ) + @staticmethod + def _recent_unhandled_images( + db: Session, + avatar: Avatar, + event: TakeoverMessage, + current_source_ids: list[str], + ) -> list[TakeoverMessage]: + """Recover a recent image that an older deployment recorded without a task.""" + threshold = event.send_time - timedelta(seconds=IMAGE_CONTEXT_LOOKBACK_SECONDS) + candidates = ( + db.query(TakeoverMessage) + .filter( + TakeoverMessage.avatar_id == avatar.id, + TakeoverMessage.owner_id == avatar.owner_id, + TakeoverMessage.peer_id == event.peer_id, + TakeoverMessage.direction == "incoming", + TakeoverMessage.message_type == BOXIM_IMAGE_MESSAGE_TYPE, + TakeoverMessage.is_avatar.is_(False), + TakeoverMessage.send_time >= threshold, + TakeoverMessage.send_time <= event.send_time, + ) + .order_by(TakeoverMessage.send_time.desc()) + .limit(MAX_RECENT_IMAGE_CONTEXTS) + .all() + ) + if not candidates: + return [] + + handled_ids = set(current_source_ids) + task_sources = ( + db.query(TakeoverReplyTask.source_message_ids) + .filter( + TakeoverReplyTask.avatar_id == avatar.id, + TakeoverReplyTask.owner_id == avatar.owner_id, + TakeoverReplyTask.peer_id == event.peer_id, + TakeoverReplyTask.created_at >= threshold, + ) + .all() + ) + for (source_message_ids,) in task_sources: + handled_ids.update(source_message_ids or []) + + return [ + image + for image in reversed(candidates) + if image.boxim_message_id not in handled_ids + ] + async def _prepare_replies(self) -> int: db = self.session_factory() try: @@ -766,6 +830,21 @@ class TakeoverService: db.commit() excluded_ids = set(task.source_message_ids or []) + source_events = { + event.boxim_message_id: event + for event in ( + db.query(TakeoverMessage) + .filter( + TakeoverMessage.owner_id == task.owner_id, + TakeoverMessage.peer_id == task.peer_id, + TakeoverMessage.avatar_id == task.avatar_id, + TakeoverMessage.boxim_message_id.in_(excluded_ids), + ) + .all() + if excluded_ids + else [] + ) + } events = ( db.query(TakeoverMessage) .filter( @@ -777,11 +856,6 @@ class TakeoverService: .limit(30) .all() ) - source_events = { - event.boxim_message_id: event - for event in events - if event.boxim_message_id in excluded_ids - } image_attachments = [] image_failed = False for message_id in (task.source_message_ids or [])[-3:]: @@ -821,11 +895,7 @@ class TakeoverService: from routers.chat import _attachment_contexts, _resolve_reply image_contexts = _attachment_contexts(image_attachments) - has_source_text = any( - event.message_type == BOXIM_TEXT_MESSAGE_TYPE and event.content.strip() - for event in source_events.values() - ) - if image_failed and not image_contexts and not has_source_text: + if image_failed and not image_contexts: answer = BOXIM_IMAGE_UNAVAILABLE_REPLY else: result = _resolve_reply( diff --git a/digital-avatar-app/backend/tests/test_takeover_service.py b/digital-avatar-app/backend/tests/test_takeover_service.py index f9d3de4..ed17dcf 100644 --- a/digital-avatar-app/backend/tests/test_takeover_service.py +++ b/digital-avatar-app/backend/tests/test_takeover_service.py @@ -85,6 +85,29 @@ class ConcurrentPollingBoxIM(FakeBoxIM): return [] +class ConcurrentMessagePollingBoxIM(ConcurrentPollingBoxIM): + async def fetch_private_messages(self, access_token, min_id="0"): + await super().fetch_private_messages(access_token, min_id) + owner_id = 100 if access_token == "prod-huihui-token" else 101 + return [ + { + "id": owner_id, + "localId": owner_id, + "sendId": owner_id + 100, + "recvId": owner_id, + "sendTime": 1_700_000_000_000, + "type": 0, + "content": "并发写入测试", + } + ] + + async def mark_private_messages_read(self, access_token, friend_id, message_id): + await asyncio.sleep(0.05) + self.read_receipts.append( + {"friendId": str(friend_id), "messageId": str(message_id)} + ) + + @pytest.fixture def service_context(tmp_path): engine = create_engine( @@ -202,6 +225,7 @@ async def test_incoming_image_is_analyzed_and_used_in_takeover_reply(service_con assert scheduled.prompt == "请看看这张图片。" finally: db.close() + clock.advance(3) def analyze(db, avatar, content, **kwargs): @@ -262,6 +286,83 @@ async def test_incoming_image_is_analyzed_and_used_in_takeover_reply(service_con db.close() +@pytest.mark.asyncio +async def test_followup_text_recovers_recent_image_recorded_without_task(service_context): + session_factory, service, boxim, clock = service_context + await service.poll_and_process_messages() + image_message = { + "id": 113, + "localId": 113, + "sendId": 200, + "recvId": 100, + "sendTime": clock.millis(), + "type": 1, + "content": json.dumps( + { + "originUrl": "https://cdn.example/case.png", + "thumbUrl": "https://cdn.example/case-thumb.png", + } + ), + } + + db = session_factory() + try: + avatar = db.get(Avatar, "avatar-1") + service._record_message(db, avatar, "100", image_message, schedule_reply=False) + cursor = db.query(TakeoverCursor).one() + cursor.last_message_id = "113" + db.commit() + finally: + db.close() + + clock.advance(60) + boxim.messages.extend( + [ + image_message, + { + "id": 114, + "localId": 114, + "sendId": 200, + "recvId": 100, + "sendTime": clock.millis(), + "type": 0, + "content": "请帮我看看这张图", + }, + ] + ) + await service.poll_messages() + + db = session_factory() + try: + task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="114").one() + assert task.source_message_ids == ["113", "114"] + assert task.prompt == "请看看这张图片。\n请帮我看看这张图" + finally: + db.close() + + clock.advance(1) + boxim.messages.append( + { + "id": 115, + "localId": 115, + "sendId": 200, + "recvId": 100, + "sendTime": clock.millis(), + "type": 0, + "content": "图里写了什么", + } + ) + await service.poll_messages() + + db = session_factory() + try: + latest = db.query(TakeoverReplyTask).filter_by(trigger_message_id="115").one() + assert latest.source_message_ids == ["113", "114", "115"] + assert latest.source_message_ids.count("113") == 1 + finally: + db.close() + + @pytest.mark.asyncio async def test_invalid_image_message_is_recorded_but_not_scheduled(service_context): session_factory, service, boxim, clock = service_context @@ -347,7 +448,7 @@ async def test_multiple_avatar_owners_are_polled_concurrently(service_context): finally: db.close() - boxim = ConcurrentPollingBoxIM() + boxim = ConcurrentMessagePollingBoxIM() service = TakeoverService( session_factory, boxim, @@ -361,6 +462,8 @@ async def test_multiple_avatar_owners_are_polled_concurrently(service_context): db = session_factory() try: assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2 + assert db.query(TakeoverMessage).count() == 2 + assert len(boxim.read_receipts) == 2 finally: db.close()