275 lines
9.3 KiB
Python
275 lines
9.3 KiB
Python
"""
|
|
私聊图片询价:识图当口播、新开一票、复用文字流程。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
import sys
|
|
import unittest
|
|
|
|
ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))
|
|
if ROOT not in sys.path:
|
|
sys.path.insert(0, ROOT)
|
|
|
|
os.environ.setdefault("YTD_ENV", "test")
|
|
os.environ.setdefault("REDIS_BACKEND", "memory")
|
|
os.environ.setdefault("CHECKPOINT_BACKEND", "memory")
|
|
os.environ.setdefault("MESSAGE_STORE_BACKEND", "memory")
|
|
os.environ.setdefault("LLM_DATA_USAGE_CONFIRMED", "false")
|
|
os.environ.setdefault("LLM_ALLOW_NETWORK", "false")
|
|
os.environ.setdefault("LEDGER_BACKEND", "memory")
|
|
|
|
from agent.channel.queue import MemoryMessageStore
|
|
from agent.channel.wecom.models import InboundMessage
|
|
from agent.channel.outbox.sender import WeComAppClient
|
|
from agent.handlers.image_inquiry import (
|
|
flush_active_wave,
|
|
handle_image_inquiry,
|
|
)
|
|
from agent.handlers.text_inquiry import handle_text_inquiry
|
|
from agent.ledger.memory_ledger import MemoryLedger
|
|
from agent.policy import inquiry_copy as copy
|
|
from agent.policy.air_text_flow import AirTextInquiryFlow, FlowSession, reset_air_text_flow_for_test
|
|
from agent.policy.image_wave import reset_image_wave_book_for_test
|
|
from agent.policy.land_text_flow import get_land_text_flow, reset_land_text_flow_for_test
|
|
from agent.policy.sea_text_flow import get_sea_text_flow, reset_sea_text_flow_for_test
|
|
from agent.routing.decision import RouteDecision
|
|
from agent.routing.dispatch import dispatch_inbound
|
|
from tests.test_sea_text_inquiry import COMPLETE_SEA
|
|
|
|
SEA_ORAL = """海运
|
|
起运港:上海
|
|
目的港:洛杉矶
|
|
品名:普货
|
|
货量:20吨
|
|
整柜或拼柜:整柜
|
|
箱型箱量:1x40HQ
|
|
贸易条款:FOB
|
|
运输分类:Port to Port
|
|
"""
|
|
|
|
LAND_ORAL = """陆运
|
|
国内运输拼车
|
|
国内长途/零担
|
|
急件
|
|
起运港:广州
|
|
目的港:深圳
|
|
品名:衣服
|
|
体积:100CBM
|
|
毛重:50kg
|
|
"""
|
|
|
|
|
|
def _vision(text: str, *, ok: bool = True, error: str = ""):
|
|
def _fn(**_kwargs):
|
|
body = (text or "").strip()
|
|
return {"ok": ok, "text": body, "empty": not bool(body), "error": error}
|
|
|
|
return _fn
|
|
|
|
|
|
def _img(sender: str, mid: str, msgid: str = "") -> InboundMessage:
|
|
return InboundMessage(
|
|
sender_id=sender,
|
|
message_id=msgid or mid,
|
|
content="",
|
|
msg_type="image",
|
|
media={"media_id": mid, "filename": f"{mid}.jpg"},
|
|
)
|
|
|
|
|
|
class ImageBookmarkTests(unittest.TestCase):
|
|
def test_drop_keeps_ticket_session(self) -> None:
|
|
flow = AirTextInquiryFlow(MemoryLedger())
|
|
sess = FlowSession(
|
|
sender_id="s1", thread_id="s1", work_order_no="WO1", phase="quoted"
|
|
)
|
|
flow._save(sess)
|
|
dropped = flow.drop_current_bookmark("s1")
|
|
self.assertEqual(dropped.work_order_no, "WO1")
|
|
self.assertIsNone(flow.session_of("s1"))
|
|
self.assertIsNotNone(flow.session_by_work_order("WO1"))
|
|
|
|
def test_download_empty_media_id(self) -> None:
|
|
client = WeComAppClient(corp_id="c", secret="s", agent_id=1)
|
|
raw, err = client.download_media("")
|
|
self.assertEqual(raw, b"")
|
|
self.assertEqual(err, "media_empty")
|
|
|
|
|
|
class ImageInquiryFlowTests(unittest.TestCase):
|
|
def setUp(self) -> None:
|
|
reset_sea_text_flow_for_test()
|
|
reset_land_text_flow_for_test()
|
|
reset_air_text_flow_for_test()
|
|
reset_image_wave_book_for_test()
|
|
self.store = MemoryMessageStore()
|
|
|
|
def _bodies(self) -> list[str]:
|
|
return [i.content for i in self.store._outbox.values()]
|
|
|
|
def test_only_image_asks_mode(self) -> None:
|
|
phase = handle_image_inquiry(
|
|
_img("s1", "mid-1"),
|
|
store=self.store,
|
|
vision_fn=_vision(""),
|
|
download_fn=lambda _mid: (b"x", ""),
|
|
)
|
|
bodies = self._bodies()
|
|
self.assertEqual(bodies[0], copy.IMAGE_RECOGNIZING)
|
|
self.assertIn("没看清", bodies[-1])
|
|
self.assertIn("海运", bodies[-1])
|
|
self.assertIn("空运", bodies[-1])
|
|
self.assertEqual(phase, "need_mode")
|
|
|
|
def test_sea_visible_creates_ticket(self) -> None:
|
|
handle_image_inquiry(
|
|
_img("s1", "mid-sea"),
|
|
store=self.store,
|
|
vision_fn=_vision(SEA_ORAL),
|
|
download_fn=lambda _mid: (b"x", ""),
|
|
put_fn=lambda **k: {"ok": True, "key": k.get("key") or "k1"},
|
|
)
|
|
sea = get_sea_text_flow()
|
|
sess = sea.session_of("s1")
|
|
self.assertIsNotNone(sess)
|
|
self.assertTrue((sess.work_order_no or "").strip())
|
|
ticket = sea.ledger.get_ticket(work_order_no=sess.work_order_no)
|
|
self.assertTrue(ticket.attachments)
|
|
|
|
def test_burst_two_images_one_ack(self) -> None:
|
|
handle_image_inquiry(
|
|
_img("s1", "mid-a", "m-a"),
|
|
store=self.store,
|
|
vision_fn=_vision(""),
|
|
download_fn=lambda _mid: (b"x", ""),
|
|
process_now=False,
|
|
)
|
|
handle_image_inquiry(
|
|
_img("s1", "mid-b", "m-b"),
|
|
store=self.store,
|
|
vision_fn=_vision(SEA_ORAL),
|
|
download_fn=lambda _mid: (b"x", ""),
|
|
process_now=False,
|
|
)
|
|
flush_active_wave(
|
|
"s1",
|
|
store=self.store,
|
|
vision_fn=_vision(SEA_ORAL),
|
|
download_fn=lambda _mid: (b"x", ""),
|
|
seed_message=_img("s1", "mid-b", "m-b"),
|
|
)
|
|
acks = [c for c in self._bodies() if c == copy.IMAGE_RECOGNIZING]
|
|
self.assertEqual(len(acks), 1)
|
|
sess = get_sea_text_flow().session_of("s1")
|
|
self.assertTrue(sess and sess.work_order_no)
|
|
|
|
def test_precede_text_joins_new_ticket(self) -> None:
|
|
handle_text_inquiry(
|
|
InboundMessage(
|
|
sender_id="s1",
|
|
message_id="txt1",
|
|
content="空运 上海到纽约 3件 180KGS 2CBM 托盘",
|
|
),
|
|
store=self.store,
|
|
)
|
|
handle_image_inquiry(
|
|
_img("s1", "mid-air"),
|
|
store=self.store,
|
|
vision_fn=_vision("包装方式:托盘"),
|
|
download_fn=lambda _mid: (b"x", ""),
|
|
)
|
|
from agent.policy.air_text_flow import get_air_text_flow
|
|
|
|
air = get_air_text_flow().session_of("s1")
|
|
self.assertIsNotNone(air)
|
|
self.assertEqual((air.business_line or "").upper(), "AIR")
|
|
|
|
def test_land_image_confirm_card_has_no_wo(self) -> None:
|
|
handle_image_inquiry(
|
|
_img("s1", "mid-land"),
|
|
store=self.store,
|
|
vision_fn=_vision(LAND_ORAL),
|
|
download_fn=lambda _mid: (b"x", ""),
|
|
)
|
|
last = self._bodies()[-1]
|
|
self.assertIn("请回复“确定”", last)
|
|
self.assertNotIn("WO", last)
|
|
sess = get_land_text_flow().session_of("s1")
|
|
self.assertEqual(sess.phase, "wait_confirm")
|
|
self.assertFalse((sess.work_order_no or "").strip())
|
|
|
|
def test_quoted_then_new_image_new_ticket(self) -> None:
|
|
handle_text_inquiry(
|
|
InboundMessage(sender_id="s1", message_id="old", content="海运出号"),
|
|
store=self.store,
|
|
injected_facts=COMPLETE_SEA,
|
|
injected_mode="SEA",
|
|
)
|
|
wo_old = get_sea_text_flow().session_of("s1").work_order_no
|
|
self.assertTrue(wo_old)
|
|
handle_image_inquiry(
|
|
_img("s1", "mid-new"),
|
|
store=self.store,
|
|
vision_fn=_vision(SEA_ORAL),
|
|
download_fn=lambda _mid: (b"x", ""),
|
|
)
|
|
wo_new = get_sea_text_flow().session_of("s1").work_order_no
|
|
self.assertTrue(wo_new)
|
|
self.assertNotEqual(wo_old, wo_new)
|
|
old = get_sea_text_flow().ledger.get_ticket(work_order_no=wo_old)
|
|
self.assertIsNotNone(old)
|
|
|
|
def test_group_image_not_this_handler(self) -> None:
|
|
msg = InboundMessage(
|
|
sender_id="s1",
|
|
message_id="g1",
|
|
content="",
|
|
msg_type="image",
|
|
chat_id="wr_group",
|
|
chat_type="group",
|
|
media={"media_id": "mid-g", "filename": "现场.jpg"},
|
|
)
|
|
name = dispatch_inbound(msg, RouteDecision(intent="ordinary_text_other"))
|
|
self.assertNotEqual(name, "image_inquiry")
|
|
|
|
def test_worker_kind_vision_finishes_wave(self) -> None:
|
|
from agent.handlers.image_inquiry import finish_queued_image
|
|
from agent.policy.image_wave import get_image_wave_book
|
|
|
|
book = reset_image_wave_book_for_test()
|
|
book.open_or_join("s1", "mid-1", now=1.0)
|
|
got = finish_queued_image(
|
|
{
|
|
"sender_id": "s1",
|
|
"media_id": "mid-1",
|
|
"message_id": "m1",
|
|
"__vision": {"ok": True, "text": "", "empty": True, "error": ""},
|
|
"__bytes": b"img",
|
|
}
|
|
)
|
|
self.assertTrue(got.get("ok"))
|
|
self.assertTrue(book._waves["s1"].business_done)
|
|
|
|
def test_png_download_sends_png_mime(self) -> None:
|
|
seen: dict = {}
|
|
png = b"\x89PNG\r\n\x1a\n" + b"\x00" * 20
|
|
|
|
def _fn(**kwargs):
|
|
seen.update(kwargs.get("payload") or {})
|
|
return {"ok": True, "text": "起运地:深圳", "empty": False, "error": ""}
|
|
|
|
handle_image_inquiry(
|
|
_img("s1", "mid-png"),
|
|
store=self.store,
|
|
vision_fn=_fn,
|
|
download_fn=lambda _mid: (png, ""),
|
|
)
|
|
self.assertEqual(seen.get("mime"), "image/png")
|
|
self.assertTrue(seen.get("image_b64"))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|