fix(avatar): recover BOXIM image replies

This commit is contained in:
stefanfeng
2026-09-02 14:12:04 +08:00
parent 03c32309a8
commit 7cac96356d
3 changed files with 216 additions and 25 deletions
+20 -2
View File
@@ -1,16 +1,30 @@
import os import os
from sqlalchemy import create_engine from sqlalchemy import create_engine, event
from sqlalchemy.orm import sessionmaker, declarative_base, Session from sqlalchemy.orm import sessionmaker, declarative_base, Session
BASE_DIR = os.path.dirname(os.path.abspath(__file__)) BASE_DIR = os.path.dirname(os.path.abspath(__file__))
DB_FILE = os.path.join(BASE_DIR, "avatar.db") DB_FILE = os.path.join(BASE_DIR, "avatar.db")
DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}") DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}")
IS_SQLITE = DATABASE_URL.startswith("sqlite:")
engine = create_engine( engine = create_engine(
DATABASE_URL, 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) SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
Base = declarative_base() Base = declarative_base()
@@ -26,6 +40,10 @@ def get_db():
def init_db(): def init_db():
import models 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) Base.metadata.create_all(bind=engine)
# 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试) # 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试)
@@ -49,6 +49,8 @@ BOXIM_TEXT_MESSAGE_TYPE = 0
BOXIM_IMAGE_MESSAGE_TYPE = 1 BOXIM_IMAGE_MESSAGE_TYPE = 1
BOXIM_IMAGE_PROMPT = "请看看这张图片。" BOXIM_IMAGE_PROMPT = "请看看这张图片。"
BOXIM_IMAGE_UNAVAILABLE_REPLY = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。" BOXIM_IMAGE_UNAVAILABLE_REPLY = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。"
IMAGE_CONTEXT_LOOKBACK_SECONDS = 1800
MAX_RECENT_IMAGE_CONTEXTS = 3
def _utcnow() -> datetime: def _utcnow() -> datetime:
@@ -149,6 +151,7 @@ class TakeoverService:
self._sessions: dict[str, dict] = {} self._sessions: dict[str, dict] = {}
self._poll_lock = asyncio.Lock() self._poll_lock = asyncio.Lock()
self._process_lock = asyncio.Lock() self._process_lock = asyncio.Lock()
self._persist_lock = asyncio.Lock()
async def poll_and_process_messages(self): async def poll_and_process_messages(self):
"""Run one complete cycle for callers that do not use the split scheduler.""" """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) max_message_id = _numeric_id(cursor.last_message_id)
read_receipts: dict[str, int] = {} read_receipts: dict[str, int] = {}
for message in messages: 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")) message_id = _numeric_id(message.get("id"))
max_message_id = max(max_message_id, message_id) max_message_id = max(max_message_id, message_id)
send_id = str(message.get("sendId") or "") send_id = str(message.get("sendId") or "")
@@ -430,6 +426,17 @@ class TakeoverService:
session["access_token"], peer_id, message_id session["access_token"], peer_id, message_id
) )
# 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.last_message_id = str(max_message_id)
cursor.initialized = True cursor.initialized = True
cursor.last_polled_at = self.now() cursor.last_polled_at = self.now()
@@ -645,6 +652,15 @@ class TakeoverService:
task.status = "cancelled" task.status = "cancelled"
task.cancel_reason = "newer_incoming_message" task.cancel_reason = "newer_incoming_message"
task.locked_at = None 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)) prompt_parts.append(_event_prompt(event))
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:]
@@ -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: async def _prepare_replies(self) -> int:
db = self.session_factory() db = self.session_factory()
try: try:
@@ -766,6 +830,21 @@ class TakeoverService:
db.commit() db.commit()
excluded_ids = set(task.source_message_ids or []) 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 = ( events = (
db.query(TakeoverMessage) db.query(TakeoverMessage)
.filter( .filter(
@@ -777,11 +856,6 @@ class TakeoverService:
.limit(30) .limit(30)
.all() .all()
) )
source_events = {
event.boxim_message_id: event
for event in events
if event.boxim_message_id in excluded_ids
}
image_attachments = [] image_attachments = []
image_failed = False image_failed = False
for message_id in (task.source_message_ids or [])[-3:]: for message_id in (task.source_message_ids or [])[-3:]:
@@ -821,11 +895,7 @@ class TakeoverService:
from routers.chat import _attachment_contexts, _resolve_reply from routers.chat import _attachment_contexts, _resolve_reply
image_contexts = _attachment_contexts(image_attachments) image_contexts = _attachment_contexts(image_attachments)
has_source_text = any( if image_failed and not image_contexts:
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:
answer = BOXIM_IMAGE_UNAVAILABLE_REPLY answer = BOXIM_IMAGE_UNAVAILABLE_REPLY
else: else:
result = _resolve_reply( result = _resolve_reply(
@@ -85,6 +85,29 @@ class ConcurrentPollingBoxIM(FakeBoxIM):
return [] 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 @pytest.fixture
def service_context(tmp_path): def service_context(tmp_path):
engine = create_engine( engine = create_engine(
@@ -202,6 +225,7 @@ async def test_incoming_image_is_analyzed_and_used_in_takeover_reply(service_con
assert scheduled.prompt == "请看看这张图片。" assert scheduled.prompt == "请看看这张图片。"
finally: finally:
db.close() db.close()
clock.advance(3) clock.advance(3)
def analyze(db, avatar, content, **kwargs): 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() 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 @pytest.mark.asyncio
async def test_invalid_image_message_is_recorded_but_not_scheduled(service_context): async def test_invalid_image_message_is_recorded_but_not_scheduled(service_context):
session_factory, service, boxim, clock = 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: finally:
db.close() db.close()
boxim = ConcurrentPollingBoxIM() boxim = ConcurrentMessagePollingBoxIM()
service = TakeoverService( service = TakeoverService(
session_factory, session_factory,
boxim, boxim,
@@ -361,6 +462,8 @@ async def test_multiple_avatar_owners_are_polled_concurrently(service_context):
db = session_factory() db = session_factory()
try: try:
assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2 assert db.query(TakeoverCursor).filter(TakeoverCursor.initialized.is_(True)).count() == 2
assert db.query(TakeoverMessage).count() == 2
assert len(boxim.read_receipts) == 2
finally: finally:
db.close() db.close()