Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6a4b35c49a | ||
|
|
207bbd02cf | ||
|
|
7cac96356d | ||
|
|
03c32309a8 |
@@ -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,逐列尝试)
|
||||
|
||||
@@ -47,6 +47,25 @@ QA_SEMANTIC_THRESHOLD = 0.72
|
||||
QA_MATCH_MARGIN = 0.06
|
||||
KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42"))
|
||||
|
||||
_IMAGE_ACCESS_DENIAL_PATTERNS = (
|
||||
re.compile(
|
||||
r"(?:我|目前|暂时|这里|本身|系统)?\s*(?:无法|不能|没法|不支持)\s*"
|
||||
r"(?:直接)?\s*(?:查看|看到|看见|识别|读取|访问|打开|分析|理解)"
|
||||
r"(?:\s*(?:或|、|/)\s*(?:查看|看到|看见|识别|读取|访问|打开|分析|理解))*\s*"
|
||||
r"(?:你(?:发|提供|上传)的|这张|该|当前)?\s*(?:图片|图像|照片|影像|文件)"
|
||||
),
|
||||
re.compile(
|
||||
r"(?:我|这里|目前|暂时)?\s*(?:看不到|看不见|未看到|没有看到|没收到|未收到)\s*"
|
||||
r"(?:你(?:发|提供|上传)的|这张|该|当前)?\s*(?:图片|图像|照片|影像)"
|
||||
),
|
||||
re.compile(
|
||||
r"\b(?:i\s+)?(?:can(?:not|'t)|am\s+unable\s+to)\s+(?:directly\s+)?"
|
||||
r"(?:view|see|access|read|analy[sz]e|recogni[sz]e)\s+"
|
||||
r"(?:the\s+|this\s+|your\s+)?(?:image|photo|picture|scan)\b",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
)
|
||||
|
||||
_WRITING_SYSTEM_PATTERNS = {
|
||||
"han": re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff]"),
|
||||
"latin": re.compile(r"[A-Za-z\u00c0-\u024f]"),
|
||||
@@ -178,6 +197,75 @@ def _image_retrieval_question(question: str, image_contexts: list[dict]) -> str:
|
||||
return "\n".join(part for part in parts if part).strip()
|
||||
|
||||
|
||||
def _answer_denies_available_image(answer: str) -> bool:
|
||||
"""Reject only whole-image access denials, not uncertainty about one field."""
|
||||
value = re.sub(r"\s+", " ", answer or "").strip()
|
||||
return any(pattern.search(value) for pattern in _IMAGE_ACCESS_DENIAL_PATTERNS)
|
||||
|
||||
|
||||
def _compact_context_text(value: Any, limit: int) -> str:
|
||||
lines = [re.sub(r"\s+", " ", line).strip() for line in str(value or "").splitlines()]
|
||||
text = "\n".join(line for line in lines if line).strip()
|
||||
return text[:limit].rstrip()
|
||||
|
||||
|
||||
def _grounded_image_fallback(question: str, image_contexts: list[dict]) -> str:
|
||||
"""Build a safe answer from completed vision data when the chat model contradicts it."""
|
||||
summaries: list[str] = []
|
||||
facts: list[str] = []
|
||||
excerpts: list[str] = []
|
||||
warnings: list[str] = []
|
||||
for context in image_contexts:
|
||||
summary = _compact_context_text(context.get("summary"), 500)
|
||||
if summary:
|
||||
summaries.append(summary)
|
||||
structured = context.get("structuredData") or {}
|
||||
if isinstance(structured, dict):
|
||||
for fact in structured.get("key_facts") or []:
|
||||
value = _compact_context_text(fact, 300)
|
||||
if value:
|
||||
facts.append(value)
|
||||
extracted = _compact_context_text(context.get("extractedText"), 900)
|
||||
if extracted:
|
||||
excerpts.append(extracted)
|
||||
warning = _compact_context_text(context.get("warning"), 300)
|
||||
if warning:
|
||||
warnings.append(warning)
|
||||
|
||||
summaries = list(dict.fromkeys(summaries))
|
||||
facts = list(dict.fromkeys(facts))[:6]
|
||||
excerpts = list(dict.fromkeys(excerpts))
|
||||
warnings = list(dict.fromkeys(warnings))
|
||||
writing_system = _dominant_writing_system(question)
|
||||
|
||||
if writing_system == "latin":
|
||||
parts = []
|
||||
if summaries:
|
||||
parts.append("From the image, I can confirm: " + " ".join(summaries))
|
||||
if facts:
|
||||
parts.append("Key details:\n" + "\n".join(
|
||||
f"{index}. {fact}" for index, fact in enumerate(facts, 1)
|
||||
))
|
||||
elif excerpts:
|
||||
parts.append("Visible text:\n" + excerpts[0])
|
||||
if warnings:
|
||||
parts.append("Please note: " + " ".join(warnings))
|
||||
return "\n".join(parts).strip() or "The image is available, but there is not enough clear detail to confirm more."
|
||||
|
||||
parts = []
|
||||
if summaries:
|
||||
parts.append("从这张图中可以确认:" + ";".join(summaries).rstrip("。;") + "。")
|
||||
if facts:
|
||||
parts.append("其中比较明确的信息有:\n" + "\n".join(
|
||||
f"{index}. {fact}" for index, fact in enumerate(facts, 1)
|
||||
))
|
||||
elif excerpts:
|
||||
parts.append("图中可见的主要文字是:\n" + excerpts[0])
|
||||
if warnings:
|
||||
parts.append("需要注意:" + ";".join(warnings).rstrip("。;") + "。")
|
||||
return "\n".join(parts).strip() or "这张图已经看到了,但目前能确认的清晰信息比较有限。"
|
||||
|
||||
|
||||
def _run_billed_vision_call(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
@@ -553,8 +641,10 @@ def _build_prompt(
|
||||
if image_contexts:
|
||||
image_material = json.dumps(image_contexts, ensure_ascii=False, default=str)
|
||||
system += (
|
||||
"\n以下是当前会话图片经过视觉识别后得到的资料:\n"
|
||||
"\n当前会话图片已经成功读取并完成内容识别,以下资料就是可直接使用的图片内容:\n"
|
||||
f"{image_material}"
|
||||
"\n必须直接依据这些图片内容回答当前问题。禁止声称无法查看、看不到、未收到、无法识别、"
|
||||
"无法读取或不能访问图片,也不要要求对方重新上传;只有资料明确标记读取失败时才可以请对方重发。"
|
||||
"\n图片资料可能包含 OCR 错字、模糊内容或用户尚未确认的信息,只能按可见内容谨慎表达。"
|
||||
"标准答题对中的事实优先级高于图片资料,知识库事实优先级高于模型推测;发生冲突时遵循更高优先级资料,"
|
||||
"并自然提醒对方核对原图。不得声称看到了图片中不存在的内容。"
|
||||
@@ -815,6 +905,14 @@ def _resolve_reply(
|
||||
except Exception as exc:
|
||||
release_reservation(db, reservation, str(exc))
|
||||
raise
|
||||
answer = str(answer or "").strip()
|
||||
if image_contexts and _answer_denies_available_image(answer):
|
||||
logger.warning(
|
||||
"chat model contradicted ready image context avatar=%s source=%s",
|
||||
avatar.id,
|
||||
usage_source,
|
||||
)
|
||||
answer = _grounded_image_fallback(question, image_contexts)
|
||||
result = {
|
||||
"answer": answer,
|
||||
"source": "qa" if matched else (
|
||||
|
||||
@@ -49,6 +49,15 @@ BOXIM_TEXT_MESSAGE_TYPE = 0
|
||||
BOXIM_IMAGE_MESSAGE_TYPE = 1
|
||||
BOXIM_IMAGE_PROMPT = "请看看这张图片。"
|
||||
BOXIM_IMAGE_UNAVAILABLE_REPLY = "这张图片我暂时没看清,麻烦重新发送一张清晰的原图。"
|
||||
IMAGE_CONTEXT_LOOKBACK_SECONDS = 1800
|
||||
IMAGE_REFERENCE_LOOKBACK_SECONDS = 172_800
|
||||
MAX_RECENT_IMAGE_CONTEXTS = 3
|
||||
|
||||
_IMAGE_REFERENCE_PATTERN = re.compile(
|
||||
r"(?:图片|图像|照片|截图|这张图|刚才.{0,8}图|病例|病历|检查单|检验单|化验单|报告|影像|"
|
||||
r"\b(?:image|photo|picture|screenshot|scan|report)\b)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _utcnow() -> datetime:
|
||||
@@ -127,6 +136,10 @@ def _event_prompt(event: TakeoverMessage) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
def _references_recent_image(value: str) -> bool:
|
||||
return bool(_IMAGE_REFERENCE_PATTERN.search(value or ""))
|
||||
|
||||
|
||||
class TakeoverService:
|
||||
"""Poll BOXIM, honor the owner grace period, then generate and send one reply."""
|
||||
|
||||
@@ -149,6 +162,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 +423,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 +437,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 +663,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 +697,68 @@ class TakeoverService:
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _recent_unhandled_images(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
event: TakeoverMessage,
|
||||
current_source_ids: list[str],
|
||||
) -> list[TakeoverMessage]:
|
||||
"""Recover missed images, or reuse a referenced image from the last two days."""
|
||||
references_image = _references_recent_image(event.content)
|
||||
lookback_seconds = (
|
||||
IMAGE_REFERENCE_LOOKBACK_SECONDS
|
||||
if references_image
|
||||
else IMAGE_CONTEXT_LOOKBACK_SECONDS
|
||||
)
|
||||
threshold = event.send_time - timedelta(seconds=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 []
|
||||
|
||||
current_ids = set(current_source_ids)
|
||||
if references_image:
|
||||
return [
|
||||
image
|
||||
for image in reversed(candidates)
|
||||
if image.boxim_message_id not in current_ids
|
||||
]
|
||||
|
||||
handled_ids = set(current_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 +855,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 +881,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 +920,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(
|
||||
|
||||
@@ -12,6 +12,7 @@ from main import app
|
||||
from models import ChatAttachment
|
||||
from routers.chat import (
|
||||
ChatIn,
|
||||
_answer_denies_available_image,
|
||||
_attachment_contexts,
|
||||
_load_chat_attachments,
|
||||
_resolve_reply,
|
||||
@@ -274,6 +275,47 @@ def test_image_context_keeps_standard_answer_authoritative():
|
||||
assert "标准答题对中的事实优先级高于图片资料" in system
|
||||
|
||||
|
||||
def test_ready_image_context_never_returns_whole_image_access_denial():
|
||||
avatar = SimpleNamespace(
|
||||
id="avatar-vision",
|
||||
name="测试分身",
|
||||
description="产品顾问",
|
||||
config={},
|
||||
)
|
||||
model = Mock(return_value="抱歉,我无法查看或识别图片,请重新上传。")
|
||||
result = _resolve_reply(
|
||||
None,
|
||||
avatar,
|
||||
"请看看这张图片",
|
||||
[],
|
||||
qa_pairs=[],
|
||||
search_fn=Mock(return_value=[]),
|
||||
model_client=model,
|
||||
image_contexts=[{
|
||||
"id": "attachment",
|
||||
"filename": "report.jpg",
|
||||
"category": "medical_document",
|
||||
"summary": "一份耳鼻喉科门诊记录",
|
||||
"extractedText": "主诉:咽痛三天",
|
||||
"structuredData": {"key_facts": ["主诉为咽痛三天"]},
|
||||
"warning": "请核对原始资料",
|
||||
}],
|
||||
)
|
||||
|
||||
assert result["source"] == "vision"
|
||||
assert "一份耳鼻喉科门诊记录" in result["answer"]
|
||||
assert "主诉为咽痛三天" in result["answer"]
|
||||
assert "无法查看" not in result["answer"]
|
||||
system = model.call_args.kwargs["messages"][0]["content"]
|
||||
assert "当前会话图片已经成功读取" in system
|
||||
assert "禁止声称无法查看" in system
|
||||
|
||||
|
||||
def test_image_denial_detector_allows_uncertain_field_in_ready_image():
|
||||
assert _answer_denies_available_image("我无法查看这张图片") is True
|
||||
assert _answer_denies_available_image("图片中患者姓名无法辨认,主诉为咽痛三天。") is False
|
||||
|
||||
|
||||
def test_attachment_context_does_not_expose_internal_fields():
|
||||
row = SimpleNamespace(
|
||||
id="attachment",
|
||||
|
||||
@@ -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,131 @@ 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_explicit_followup_reuses_handled_image_within_two_days(service_context):
|
||||
session_factory, service, boxim, clock = service_context
|
||||
await service.poll_and_process_messages()
|
||||
image_message = {
|
||||
"id": 116,
|
||||
"localId": 116,
|
||||
"sendId": 200,
|
||||
"recvId": 100,
|
||||
"sendTime": clock.millis(),
|
||||
"type": 1,
|
||||
"content": json.dumps({"originUrl": "https://cdn.example/handled-case.png"}),
|
||||
}
|
||||
boxim.messages.append(image_message)
|
||||
await service.poll_messages()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
image_task = db.query(TakeoverReplyTask).filter_by(trigger_message_id="116").one()
|
||||
image_task.status = "sent"
|
||||
image_task.sent_at = clock.now()
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
clock.advance(47 * 60 * 60)
|
||||
boxim.messages.append(
|
||||
{
|
||||
"id": 117,
|
||||
"localId": 117,
|
||||
"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="117").one()
|
||||
assert task.source_message_ids == ["116", "117"]
|
||||
assert task.prompt == "请看看这张图片。\n重新看一下刚才那张病例图片"
|
||||
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 +496,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 +510,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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user