Merge pull request 'fix(avatar): recover BOXIM image replies' (#13) from codex/avatar-boxim-vision-hotfix-20260902 into main
Reviewed-on: #13
This commit was merged in pull request #13.
This commit is contained in:
@@ -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()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user