diff --git a/inquiry-agent/agent/channel/outbox/sender.py b/inquiry-agent/agent/channel/outbox/sender.py index 612c772..72cb0f1 100644 --- a/inquiry-agent/agent/channel/outbox/sender.py +++ b/inquiry-agent/agent/channel/outbox/sender.py @@ -628,6 +628,42 @@ class WeComAppClient: self._token_expire_at = time.time() + expires_in return token + def download_media(self, media_id: str) -> tuple[bytes, str]: + """ + 下载企微临时素材(cgi-bin/media/get)。 + + 返回 (字节, 错误码)。失败时字节为空。禁止把 token 写入日志。 + 调用:Worker 识别槽,禁止回调线程同步下载。 + """ + mid = (media_id or "").strip() + if not mid: + return b"", "media_empty" + try: + token = self._get_token() + except Exception: + logger.exception("下载素材取 token 失败") + return b"", "token_failed" + url = f"{self._api_base}/cgi-bin/media/get" + try: + with httpx.Client(timeout=self._timeout) as client: + resp = client.get(url, params={"access_token": token, "media_id": mid}) + ctype = (resp.headers.get("content-type") or "").lower() + raw = resp.content or b"" + if "application/json" in ctype or raw[:1] in (b"{", b"["): + try: + data = resp.json() + except Exception: + data = {} + if isinstance(data, dict) and int(data.get("errcode") or 0) != 0: + return b"", str(data.get("errmsg") or "media_get_fail") + return b"", "media_json_body" + if len(raw) < 32: + return b"", "media_too_small" + return raw, "" + except Exception: + logger.exception("下载素材失败") + return b"", "media_download_failed" + # 企微侧明确不可重试的错误(IP 白名单、非法 userid 等) PERMANENT_SEND_ERRORS = frozenset( diff --git a/inquiry-agent/agent/channel/runtime.py b/inquiry-agent/agent/channel/runtime.py index 31cb755..4aa30bf 100644 --- a/inquiry-agent/agent/channel/runtime.py +++ b/inquiry-agent/agent/channel/runtime.py @@ -75,6 +75,11 @@ class ChannelRuntime: from agent.routing.decision import RouteDecision decision = RouteDecision(intent="ordinary_text_other") + elif (message.msg_type or "text").lower() == "image": + from agent.routing.decision import RouteDecision + + # 私聊图片不进 DeepSeek A(正文为空会被判闲聊)。 + decision = RouteDecision(intent="image_inquiry", confidence=1.0) else: decision = self._router.route_inbound( text=message.content or "", diff --git a/inquiry-agent/agent/handlers/__init__.py b/inquiry-agent/agent/handlers/__init__.py index f99cc82..7e5bd3d 100644 --- a/inquiry-agent/agent/handlers/__init__.py +++ b/inquiry-agent/agent/handlers/__init__.py @@ -2,10 +2,16 @@ Handler 包:一类意图 / 一类卡片一个模块。 禁止:单文件处理全部卡片;禁止把 Schema 校验复制粘贴进每个 handler(放 schema/policy)。 -当前:文字询价空运/海运闭环 + 海运应用群协同 + 空运 BOT 群文字/附件 + echo 联调残留。 +当前:文字询价空运/海运/陆运闭环 + 私聊图片询价 + 海运应用群协同 + 空运 BOT 群文字/附件 + echo 联调残留。 """ from agent.handlers.echo_text import handle_echo_text +from agent.handlers.image_inquiry import handle_image_inquiry from agent.handlers.text_inquiry import handle_text_inquiry, handle_text_other -__all__ = ["handle_echo_text", "handle_text_inquiry", "handle_text_other"] +__all__ = [ + "handle_echo_text", + "handle_image_inquiry", + "handle_text_inquiry", + "handle_text_other", +] diff --git a/inquiry-agent/agent/handlers/image_inquiry.py b/inquiry-agent/agent/handlers/image_inquiry.py new file mode 100644 index 0000000..a01933b --- /dev/null +++ b/inquiry-agent/agent/handlers/image_inquiry.py @@ -0,0 +1,409 @@ +""" +意图:私聊图片询价。 + +本文件职责:收图 → 回「正在识别」→ 千问抄可见原文 → 拼口播 → 新开后复用文字询价。 +禁止:千问填最终缺项;另写海运/空运/陆运图片专线;回调线程同步等千问;群图走本 Handler。 +""" + +from __future__ import annotations + +import logging +import time +from typing import Any, Callable, Optional + +from agent.channel.queue import get_message_store +from agent.channel.wecom.models import InboundMessage +from agent.policy import inquiry_copy as copy +from agent.policy.image_wave import get_image_wave_book + +logger = logging.getLogger(__name__) + +VisionFn = Callable[..., dict[str, Any]] +DownloadFn = Callable[[str], tuple[bytes, str]] +PutFn = Callable[..., dict[str, Any]] + + +def drop_all_private_bookmarks(sender_id: str) -> None: + """三线当前书签都丢掉;已出号仍能按工单号找回。""" + from agent.policy.air_text_flow import get_air_text_flow + from agent.policy.land_text_flow import get_land_text_flow + from agent.policy.sea_text_flow import get_sea_text_flow + + get_air_text_flow().drop_current_bookmark(sender_id) + get_sea_text_flow().drop_current_bookmark(sender_id) + get_land_text_flow().drop_current_bookmark(sender_id) + + +def is_active_image_wave(sender_id: str) -> bool: + """业务结果尚未发出的图片波次。""" + return get_image_wave_book().active(sender_id) is not None + + +def append_wave_text(sender_id: str, text: str) -> bool: + """认图完成前的字并进本波,不走文字询价。""" + return get_image_wave_book().append_follow_text(sender_id, text) + + +def note_private_text(sender_id: str, text: str, *, now: Optional[float] = None) -> None: + """文字询价记下最近一句,供 10 秒内发图带走。""" + get_image_wave_book().note_text(sender_id, text, now=now if now is not None else time.time()) + + +def handle_image_inquiry( + message: InboundMessage, + *, + store=None, + vision_fn: Optional[VisionFn] = None, + download_fn: Optional[DownloadFn] = None, + put_fn: Optional[PutFn] = None, + process_now: bool = True, +) -> str: + """ + 私聊一张图。inbox 线程只开波、回识别中;认图在 Worker 或单测 vision_fn。 + + 须已有 sender_id。群消息不要调用本函数。 + """ + if not message.has_sender(): + logger.warning("image_inquiry 拒绝:无 sender_id") + return "reject_no_sender" + if (message.chat_type or "single") == "group": + logger.info("image_inquiry 忽略群图 chat=%s", message.chat_id) + return "ignore_group_image" + + book = get_image_wave_book() + msg_store = store or get_message_store() + media_id = str((message.media or {}).get("media_id") or "") + filename = str((message.media or {}).get("filename") or "image.jpg") + wave, how = book.open_or_join(message.sender_id, media_id, now=time.time()) + if how in {"opened", "new_after_done"}: + drop_all_private_bookmarks(message.sender_id) + msg_store.enqueue_outbound( + touser=message.sender_id, + content=copy.IMAGE_RECOGNIZING, + dedupe_key=f"img-ack:{wave.wave_id}", + ) + book.mark_ack(message.sender_id) + logger.info("image_inquiry 开波 how=%s wave=%s", how, wave.wave_id) + else: + logger.info( + "image_inquiry 并入 how=%s wave=%s pending=%s", + how, + wave.wave_id, + wave.pending, + ) + + if process_now and vision_fn is not None: + _recognize_one( + sender_id=message.sender_id, + media_id=media_id, + filename=filename, + vision_fn=vision_fn, + download_fn=download_fn, + put_fn=put_fn, + ) + wave = book.active(message.sender_id) or wave + if wave.pending <= 0: + return finish_wave(message.sender_id, store=msg_store, seed_message=message) + return "recognizing" + + if vision_fn is None: + try: + from agent.jobs.recognition import enqueue_recognition + + enqueue_recognition( + kind="vision", + payload={ + "sender_id": message.sender_id, + "media_id": media_id, + "filename": filename, + "message_id": message.message_id, + "wave_id": wave.wave_id, + }, + idempotency_key=message.message_id or media_id, + ) + except Exception: + logger.exception("image_inquiry 入识别队失败 sender=%s", message.sender_id) + _recognize_one( + sender_id=message.sender_id, + media_id=media_id, + filename=filename, + vision_fn=None, + download_fn=download_fn, + put_fn=put_fn, + ) + return finish_wave(message.sender_id, store=msg_store, seed_message=message) + return "recognizing" + + return "recognizing" + + +def flush_active_wave( + sender_id: str, + *, + store=None, + vision_fn: Optional[VisionFn] = None, + download_fn: Optional[DownloadFn] = None, + put_fn: Optional[PutFn] = None, + seed_message: Optional[InboundMessage] = None, +) -> str: + """单测:两张图都 join 后再认,只出一次业务结果。""" + book = get_image_wave_book() + wave = book.active(sender_id) + if wave is None: + return "no_wave" + pending_ids = list(wave.media_ids[-wave.pending :]) if wave.pending else [] + for mid in pending_ids: + _recognize_one( + sender_id=sender_id, + media_id=mid, + filename="", + vision_fn=vision_fn, + download_fn=download_fn, + put_fn=put_fn, + ) + return finish_wave(sender_id, store=store, seed_message=seed_message) + + +def finish_queued_image(payload: dict[str, Any]) -> dict[str, Any]: + """ + Worker 识别槽:认完一张,pending 为 0 则出业务结果。 + + payload.__vision / __bytes 仅单测注入,生产不得依赖。 + """ + sender_id = str(payload.get("sender_id") or "") + media_id = str(payload.get("media_id") or "") + filename = str(payload.get("filename") or "image.jpg") + vision_hook = payload.get("__vision") + raw = payload.get("__bytes") + + def _vision(**_kwargs: Any) -> dict[str, Any]: + if isinstance(vision_hook, dict): + return dict(vision_hook) + from agent.config import get_settings + from agent.llm.mode_vision import invoke_vision + + settings = get_settings() + return invoke_vision( + payload=_kwargs.get("payload") or {}, + allow_network=bool(settings.llm_allow_network and settings.llm_data_usage_confirmed), + ) + + def _download(mid: str) -> tuple[bytes, str]: + if isinstance(raw, (bytes, bytearray)): + return bytes(raw), "" + from agent.channel.outbox.sender import WeComAppClient + from agent.config import get_settings + + settings = get_settings() + client = WeComAppClient( + corp_id=settings.wecom_corp_id, + secret=settings.wecom_secret, + agent_id=int(settings.wecom_agent_id or "0") or 0, + ) + return client.download_media(mid) + + _recognize_one( + sender_id=sender_id, + media_id=media_id, + filename=filename, + vision_fn=_vision, + download_fn=_download, + put_fn=None, + ) + book = get_image_wave_book() + wave = book.active(sender_id) + if wave is None: + logger.warning( + "image_inquiry Worker 认完但无活跃波次 sender=%s media=%s(多半 HTTP/Worker 账不一致)", + sender_id, + media_id, + ) + return {"ok": False, "phase": "no_active_wave", "pending": 0} + if wave.pending > 0: + return {"ok": True, "phase": "recognizing", "pending": wave.pending} + phase = finish_wave(sender_id, seed_message=InboundMessage( + sender_id=sender_id, + message_id=str(payload.get("message_id") or media_id), + content="", + msg_type="image", + media={"media_id": media_id, "filename": filename}, + )) + return {"ok": True, "phase": phase} + + +def _recognize_one( + *, + sender_id: str, + media_id: str, + filename: str, + vision_fn: Optional[VisionFn], + download_fn: Optional[DownloadFn], + put_fn: Optional[PutFn], +) -> None: + """下载、可选落对象存储、读图,写入波次。""" + import base64 + + from agent.llm.mode_vision import sniff_image_mime + + book = get_image_wave_book() + err = "" + raw = b"" + if download_fn is None: + raw, err = b"", "no_download" + else: + raw, err = download_fn(media_id) + if not raw: + logger.info( + "image_inquiry 无字节 sender=%s err=%s", + sender_id, + err or "media_empty", + ) + book.add_vision( + sender_id, + "", + object_key="", + filename=filename, + error=err or "media_empty", + ) + return + mime = sniff_image_mime(raw) or "image/jpeg" + object_key = "" + if put_fn is not None: + key = f"inquiry-image/{sender_id}/{media_id or filename}" + try: + stored = put_fn(key=key, data=raw, content_type=mime) + if stored.get("ok"): + object_key = str(stored.get("key") or key) + except Exception: + logger.exception("原图落存储失败 sender=%s", sender_id) + else: + object_key = f"inquiry-image/{sender_id}/{media_id or filename}" + try: + from agent.config import get_settings + from agent.storage.s3_client import S3ObjectStorage + + S3ObjectStorage.from_settings(get_settings(), allow_network=False).put_bytes( + key=object_key, data=raw, content_type=mime + ) + except Exception: + logger.info("原图 dry_run 存储跳过 sender=%s", sender_id) + + vision: dict[str, Any] = {"ok": False, "text": "", "empty": True, "error": err or "no_vision"} + if vision_fn is not None: + b64 = base64.b64encode(raw).decode("ascii") + vision = vision_fn(payload={"image_b64": b64, "mime": mime, "media_id": media_id}) or vision + else: + from agent.config import get_settings + from agent.llm.mode_vision import invoke_vision + + settings = get_settings() + vision = invoke_vision( + payload={"image_b64": base64.b64encode(raw).decode("ascii"), "mime": mime}, + allow_network=bool(settings.llm_allow_network and settings.llm_data_usage_confirmed), + ) + text = str(vision.get("text") or "") + if not vision.get("ok"): + err = str(vision.get("error") or err or "vision_failed") + logger.info( + "image_inquiry 千问 bytes=%s mime=%s ok=%s empty=%s text_len=%s err=%s", + len(raw), + mime, + bool(vision.get("ok")), + bool(vision.get("empty")), + len(text.strip()), + err or "-", + ) + book.add_vision( + sender_id, + text, + object_key=object_key, + filename=filename, + error=err, + ) + + +def finish_wave( + sender_id: str, + *, + store=None, + seed_message: Optional[InboundMessage] = None, +) -> str: + """ + 本波图都认完:拼口播、当文字询价、必要时把说明和补问合成一条。 + """ + from agent.handlers.text_inquiry import handle_text_inquiry + from agent.policy.air_text_flow import get_air_text_flow + from agent.policy.land_text_flow import get_land_text_flow + from agent.policy.sea_text_flow import get_sea_text_flow + + book = get_image_wave_book() + wave = book.active(sender_id) + oral = book.compose_oral(sender_id) + reason = "" + if not oral.strip(): + reason = copy.IMAGE_TECH_FAIL if (wave and wave.errors) else copy.IMAGE_UNREADABLE + + msg_store = store or get_message_store() + prefix_used = {"v": False} + + class _StoreProxy: + def enqueue_outbound(self, **kwargs): + body = str(kwargs.get("content") or "") + extra = dict(kwargs.get("payload") or {}) + if reason and not prefix_used["v"] and body: + body = copy.image_fail_then_clarify(reason, body) + kwargs["content"] = body + prefix_used["v"] = True + return msg_store.enqueue_outbound( + touser=kwargs.get("touser") or sender_id, + content=body, + dedupe_key=str(kwargs.get("dedupe_key") or ""), + payload=extra, + ) + + inbound = InboundMessage( + sender_id=sender_id, + message_id=(seed_message.message_id if seed_message else f"img:{sender_id}:{int(time.time())}"), + content=oral, + msg_type="text", + media=dict(seed_message.media) if seed_message else {}, + raw=dict(seed_message.raw) if seed_message else {}, + ) + phase = handle_text_inquiry(inbound, store=_StoreProxy()) + if reason and not prefix_used["v"]: + msg_store.enqueue_outbound( + touser=sender_id, + content=copy.image_fail_then_clarify(reason, copy.ASK_TRANSPORT), + dedupe_key=f"img-empty:{inbound.message_id}", + ) + phase = "need_mode" + book.mark_business_done(sender_id) + book.forget_text(sender_id) + sess = ( + get_sea_text_flow().session_of(sender_id) + or get_air_text_flow().session_of(sender_id) + or get_land_text_flow().session_of(sender_id) + ) + wo = (sess.work_order_no if sess else "") or "" + if wo and wave: + ledger = None + for engine in (get_sea_text_flow(), get_air_text_flow(), get_land_text_flow()): + hit = engine.session_of(sender_id) + if hit and hit.work_order_no == wo: + ledger = engine.ledger + break + if ledger is not None: + for i, key in enumerate(wave.object_keys): + name = wave.filenames[i] if i < len(wave.filenames) else "image.jpg" + try: + ledger.attach_file( + work_order_no=wo, + file_name=name or "image.jpg", + object_key=key, + content_type="image/jpeg", + sender_id=sender_id, + ) + except Exception: + logger.exception("挂图片附件失败 wo=%s", wo) + logger.info("image_inquiry 波次完成 phase=%s oral_len=%s", phase, len(oral)) + return phase diff --git a/inquiry-agent/agent/handlers/text_inquiry.py b/inquiry-agent/agent/handlers/text_inquiry.py index 3ded7ff..b966dec 100644 --- a/inquiry-agent/agent/handlers/text_inquiry.py +++ b/inquiry-agent/agent/handlers/text_inquiry.py @@ -21,11 +21,23 @@ from agent.policy.sea_text_flow import get_sea_text_flow logger = logging.getLogger(__name__) +def _front_bookmark(sess) -> bool: + """补问/核对手稿优先于已出票旧书签。空的 need_mode 不算(那是丢会话后的误伤)。""" + if sess is None: + return False + phase = (sess.phase or "").strip() + line = (sess.business_line or "").upper() + if phase == "need_mode" and not line: + return False + return phase in {"clarify", "need_mode", "need_land_options", "wait_confirm"} + + def _pick_text_flow(*, sender_id: str, text: str, injected_mode: str = ""): """ 按当前书签或本句运输方式选空运/海运/陆运流程。 陆运核对卡上的确定优先于旧海运书签,避免「确定」被海运/空运吃掉。 + 图片识图在 Worker 写下的补问书签,HTTP 必须从 Redis 读到,否则「托盘」会变成再问运输方式。 """ from agent.schema.land_options import looks_like_confirm @@ -38,10 +50,15 @@ def _pick_text_flow(*, sender_id: str, text: str, injected_mode: str = ""): if retry: return land sea = get_sea_text_flow() + air = get_air_text_flow() sea_sess = sea.session_of(sender_id) + land_sess = land.session_of(sender_id) + air_sess = air.session_of(sender_id) + for sess, flow in ((land_sess, land), (air_sess, air), (sea_sess, sea)): + if _front_bookmark(sess): + return flow if sea_sess and (sea_sess.business_line or "").upper() == "SEA": return sea - land_sess = land.session_of(sender_id) if land_sess and (land_sess.business_line or "").upper() == "LAND": return land mode = (injected_mode or detect_transport_mode(text) or "").upper() @@ -49,7 +66,7 @@ def _pick_text_flow(*, sender_id: str, text: str, injected_mode: str = ""): return sea if mode == "LAND": return land - return get_air_text_flow() + return air def is_pending_land_confirm(message: InboundMessage) -> bool: @@ -126,6 +143,9 @@ def handle_text_inquiry( if not message.has_sender(): logger.warning("text_inquiry 拒绝:无 sender_id") return "reject_no_sender" + from agent.handlers.image_inquiry import note_private_text + + note_private_text(message.sender_id, message.content or "") settings = get_settings() msg_store = store or get_message_store() engine = flow or _pick_text_flow( diff --git a/inquiry-agent/agent/http_app.py b/inquiry-agent/agent/http_app.py index 6cc8e69..5d88152 100644 --- a/inquiry-agent/agent/http_app.py +++ b/inquiry-agent/agent/http_app.py @@ -82,6 +82,7 @@ def main() -> None: level=logging.INFO, format="%(asctime)s %(levelname)s [%(name)s] %(message)s", ) + logging.getLogger("httpx").setLevel(logging.WARNING) settings = get_settings() # test 环境默认允许本机注入,便于并存期联调 if settings.ytd_env == "test" and not settings.dev_inject_enabled: diff --git a/inquiry-agent/agent/jobs/__init__.py b/inquiry-agent/agent/jobs/__init__.py index 8b9b0b2..c4df2ea 100644 --- a/inquiry-agent/agent/jobs/__init__.py +++ b/inquiry-agent/agent/jobs/__init__.py @@ -19,7 +19,7 @@ from agent.jobs.archive import archive_stream, process_archive_stub from agent.jobs.libreoffice import LibreOfficeConverter from agent.jobs.pdf_job import pdf_stream, process_pdf_job from agent.jobs.quote_fill import fill_quote_workbook -from agent.jobs.recognition import process_recognition_stub +from agent.jobs.recognition import process_recognition from agent.jobs.wake_consumer import claim_and_apply_wake, wake_stream from agent.llm import LlmGateway, LlmMode from agent.redis_coord.runtime import RedisRuntime @@ -211,7 +211,7 @@ class WorkerRuntime: return entry_id, data = claimed[0] try: - process_recognition_stub(data) + process_recognition(data) finally: stream.ack(entry_id) diff --git a/inquiry-agent/agent/jobs/recognition.py b/inquiry-agent/agent/jobs/recognition.py index dd3178d..2047daa 100644 --- a/inquiry-agent/agent/jobs/recognition.py +++ b/inquiry-agent/agent/jobs/recognition.py @@ -1,7 +1,7 @@ """ 识别任务壳:入队 Redis Stream,Worker HTTP 槽认领后跑占位处理。 -本文件职责:enqueue / process_stub;正式识别(附件/OCR/字段)后续按 payload.kind 分支。 +本文件职责:识别入队;Worker 槽消费。kind=vision 跑私聊图片波次,其它 kind 仍占位。 禁止:回调线程同步等识别完成;禁止本模块写主账六态。 """ @@ -62,6 +62,21 @@ def enqueue_recognition( return RecognitionEnqueueResult(accepted=True, entry_id=entry_id) +def process_recognition(data: dict[str, Any]) -> dict[str, Any]: + """ + Worker 槽:vision 跑图片波次,其它 kind 仍走占位壳。 + + 禁止在本函数写六态。 + """ + kind = str(data.get("kind") or "") + payload = data.get("payload") if isinstance(data.get("payload"), dict) else {} + if kind == "vision": + from agent.handlers.image_inquiry import finish_queued_image + + return finish_queued_image(payload) + return process_recognition_stub(data) + + def process_recognition_stub(data: dict[str, Any]) -> dict[str, Any]: """ 识别占位:回显 kind,不读附件、不调模型。 diff --git a/inquiry-agent/agent/llm/extract_text.py b/inquiry-agent/agent/llm/extract_text.py index 2937cf9..177ebd9 100644 --- a/inquiry-agent/agent/llm/extract_text.py +++ b/inquiry-agent/agent/llm/extract_text.py @@ -24,13 +24,18 @@ _LABEL_LINE = re.compile( r"^(起运地|起运港|目的地|目的港|品名|件数|毛重|体积|重量|重量\(KG\)|体积\(CBM\)|" r"包装方式|包装类型|货量|货物数量|整柜或拼柜|箱型箱量|报价日期|车型/数量|车型数量|" r"贸易条款|货好时间|商品海关编码|海关编码|HS编码|HS|是否含油|是否含电|是否含磁|" - r"客户名称|货值|是否为危险品)\s*[::]\s*(.+)$" + r"客户名称|货值|是否为危险品|通关口岸(非必填)|通关口岸)\s*[::]\s*(.+)$" ) _LABEL_INLINE = re.compile( r"(起运地|起运港|目的地|目的港|品名|件数|毛重|体积|重量|重量\(KG\)|体积\(CBM\)|" r"包装方式|包装类型|货量|货物数量|整柜或拼柜|箱型箱量|报价日期|车型/数量|车型数量|" r"贸易条款|货好时间|商品海关编码|海关编码|HS编码|HS|是否含油|是否含电|是否含磁|" - r"客户名称|货值|是否为危险品)\s*[::]\s*(\S+)" + r"客户名称|货值|是否为危险品|通关口岸(非必填)|通关口岸)\s*[::]\s*(\S+)" +) +# 陆运核对卡非必填:销售可能写「通关口岸皇岗」或照抄「通关口岸(非必填):皇岗」。长词优先。 +_OPTIONAL_INQUIRY_LABELS = ( + "通关口岸(非必填)", + "通关口岸", ) # 群里协同标签,冒号可省略:「客户名称张三」「货值 10万」。长词优先。 _COLLAB_LABELS = ( @@ -122,6 +127,7 @@ def extract_inquiry_snapshot( mode = (injected_mode or detect_transport_mode(text) or "").upper() facts = harvest_oral_measures(text, normalize_facts(dict(injected_facts))) facts = _merge_spoken_collab(text, facts) + facts = _merge_optional_inquiry_labels(text, facts) subtype = str(facts.get("运输类型") or "").strip() return { "business_line": mode, @@ -147,6 +153,7 @@ def extract_inquiry_snapshot( source = "deepseek_b" facts = harvest_oral_measures(text, normalize_facts(facts)) facts = _merge_spoken_collab(text, facts) + facts = _merge_optional_inquiry_labels(text, facts) land_subtype = str((b_payload or {}).get("land_subtype") or facts.get("运输类型") or "").strip() if land_subtype and not str(facts.get("运输类型") or "").strip(): facts["运输类型"] = land_subtype @@ -168,9 +175,36 @@ def harvest_collab_labels(text: str) -> dict[str, str]: 「客户名称张三 包装方式编织袋 货值 10万」切成三项。 只说「纸箱」没有标签则不收。值取到下一个标签之前。 """ + return _harvest_labels(text, _COLLAB_LABELS) + + +def harvest_inquiry_optional_labels(text: str) -> dict[str, str]: + """ + 陆运询价非必填标签:通关口岸。 + + 核对卡写「通关口岸(非必填):」,销售可能照抄或只写「通关口岸皇岗」。 + 有无冒号都收。不猜口岸名。 + """ + return _harvest_labels(text, _OPTIONAL_INQUIRY_LABELS) + + +def _merge_optional_inquiry_labels(text: str, facts: dict[str, str]) -> dict[str, str]: + """把原话里的通关口岸写进询价事实。有值才覆盖。""" + from agent.schema.field_validate import normalize_facts + + merged = dict(facts or {}) + extra = normalize_facts(harvest_inquiry_optional_labels(text)) + for key, val in extra.items(): + if val: + merged[key] = val + return merged + + +def _harvest_labels(text: str, labels: tuple[str, ...]) -> dict[str, str]: + """按标签切原话。同一起点只留最长标签。""" raw = text or "" hits: list[tuple[int, str]] = [] - for label in _COLLAB_LABELS: + for label in labels: start = 0 while True: pos = raw.find(label, start) @@ -181,7 +215,7 @@ def harvest_collab_labels(text: str) -> dict[str, str]: if not hits: return {} hits.sort() - # 同一起点只留最长标签,避免 HS 抢 HS编码。 + # 同一起点只留最长标签,避免 HS 抢 HS编码、通关口岸抢「通关口岸(非必填)」。 kept: list[tuple[int, str]] = [] for pos, label in hits: if kept and kept[-1][0] == pos and len(label) <= len(kept[-1][1]): diff --git a/inquiry-agent/agent/llm/http_client.py b/inquiry-agent/agent/llm/http_client.py index 29d265e..a3f5e37 100644 --- a/inquiry-agent/agent/llm/http_client.py +++ b/inquiry-agent/agent/llm/http_client.py @@ -19,6 +19,42 @@ from agent.llm.modes import LlmMode logger = logging.getLogger(__name__) +def message_plain_text(msg: dict[str, Any]) -> str: + """ + 把 assistant message 收成纯文本。 + + 千问视觉的 content 经常是 [{text: ...}] 数组,不是字符串; + 思考模式打开时正文可能只在 reasoning_content。两者都漏会让识图变「没看清」。 + """ + if not isinstance(msg, dict): + return "" + text = _content_to_text(msg.get("content")) + if text: + return text + reason = msg.get("reasoning_content") + if isinstance(reason, str) and reason.strip(): + return reason.strip() + return "" + + +def _content_to_text(content: Any) -> str: + """字符串或千问多模态数组 → 纯文本。""" + if isinstance(content, str): + return content.strip() + if isinstance(content, list): + parts: list[str] = [] + for item in content: + if isinstance(item, str) and item.strip(): + parts.append(item.strip()) + elif isinstance(item, dict): + piece = item.get("text") if "text" in item else item.get("content") + got = _content_to_text(piece) + if got: + parts.append(got) + return "\n".join(parts).strip() + return "" + + @dataclass class LlmHttpResult: ok: bool @@ -72,7 +108,7 @@ class LlmHttpClient: mode_val = mode.value if isinstance(mode, LlmMode) else str(mode) if provider == "qwen": if mode_val == LlmMode.VISION.value: - return self.settings.qwen_vl_model or "qwen-vl-plus" + return self.settings.qwen_vl_model or "qwen3-vl-plus" return self.settings.qwen_text_model or "qwen-plus" return self.settings.deepseek_model or "deepseek-chat" @@ -115,6 +151,9 @@ class LlmHttpClient: "messages": messages, "temperature": temperature, } + if provider == "qwen": + # 识图只要可见原文。思考开着时 content 常为空,销售侧就会「认不到」。 + body["enable_thinking"] = False if tools: body["tools"] = tools if tool_choice is not None: @@ -147,7 +186,7 @@ class LlmHttpClient: tool_calls: list[dict[str, Any]] = [] if choices: msg = (choices[0].get("message") or {}) - content = str(msg.get("content") or "") + content = message_plain_text(msg if isinstance(msg, dict) else {}) raw_calls = msg.get("tool_calls") or [] if isinstance(raw_calls, list): tool_calls = [c for c in raw_calls if isinstance(c, dict)] diff --git a/inquiry-agent/agent/llm/mode_extract_fields.py b/inquiry-agent/agent/llm/mode_extract_fields.py index 93466c0..aba5fa5 100644 --- a/inquiry-agent/agent/llm/mode_extract_fields.py +++ b/inquiry-agent/agent/llm/mode_extract_fields.py @@ -33,6 +33,7 @@ _FIELD_KEYS = ( "箱型箱量", "车型/数量", "车型数量", + "通关口岸", ) @@ -127,7 +128,10 @@ def build_extract_messages(text: str) -> list[dict[str, str]]: "10件/3pcs→空运件数;海运同样数字+件也要写入货量;100KGS/180kg/1.5KG/100公斤→毛重;2CBM/4.2cbm/2.5立方/2方→体积;" "托盘/散货/卡板/纸箱/木箱→包装方式(销售说什么记什么,散货也要抽)。" "9月15日/9月15号/2026-09-15→报价日期,规范成 yyyy-MM-dd,未写年用今年。" + "陆运「通关口岸:皇岗」「通关口岸皇岗」「过皇岗」有说就写入通关口岸,没说留空,禁止编造口岸。" "「从南京飞吉隆坡」起运港=南京、目的港=吉隆坡;「北京飞吉隆坡」起运港=北京。" + "原话是机场三字码或航线就原样记,禁止翻译成城市名:" + "SZX-CGK / SZX/CGK / SZX到CGK → 起运港=SZX、目的港=CGK,不要写成深圳、雅加达。" "换城市必须覆盖旧值,不要沿用上一句的港口。" "运输方式仅在原话能判断时填 AIR/SEA/LAND,否则 UNKNOWN。" ), diff --git a/inquiry-agent/agent/llm/mode_vision.py b/inquiry-agent/agent/llm/mode_vision.py index 3e8958e..9e59027 100644 --- a/inquiry-agent/agent/llm/mode_vision.py +++ b/inquiry-agent/agent/llm/mode_vision.py @@ -1,30 +1,166 @@ """ -千问视觉:读图 mode 壳。 +千问视觉:只读图上可见原文。 + +本文件职责:把图片交给 qwen3-vl-plus,抄出看得见的文字。 +属于 llm 包。依赖:LlmHttpClient(生产);单测注入 chat_fn。 +禁止:填询价字段 JSON、算最终缺项、发明港口/重量/日期;禁止在回调线程长等(由 Worker 槽调用)。 """ from __future__ import annotations -from typing import Any +import logging +from typing import Any, Optional from agent.llm.modes import LlmMode +logger = logging.getLogger(__name__) + MODE = LlmMode.VISION +VISION_SYSTEM = ( + "只读出图片上能看见的文字和数字,原样抄写,一行不漏。" + "英文、机场三字码、航线(如 SZX-CGK、SZX/CGK)必须保持原样,禁止翻译成中文城市名。" + "看不清的字不要猜。不要补询价缺项,不要发明港口、重量或日期。" + "只有整张图完全没有文字时才输出空;只要看得见字就必须抄出来。" +) + + +def sniff_image_mime(raw: bytes) -> str: + """ + 按文件头认图片类型。企微截图经常是 PNG,不能一律当 JPEG 送给千问。 + """ + blob = raw or b"" + if len(blob) < 12: + return "" + if blob.startswith(b"\x89PNG\r\n\x1a\n"): + return "image/png" + if blob.startswith(b"\xff\xd8\xff"): + return "image/jpeg" + if blob[:4] == b"RIFF" and blob[8:12] == b"WEBP": + return "image/webp" + if blob.startswith(b"GIF87a") or blob.startswith(b"GIF89a"): + return "image/gif" + if blob[4:8] == b"ftyp": + return "image/heic" + if blob.startswith(b"BM"): + return "image/bmp" + return "" + def mode_meta() -> dict[str, Any]: return { "mode": MODE.value, "provider": "qwen", "purpose": "vision", - "implemented": False, + "implemented": True, + } + + +def _pack_text(text: str, *, stub: bool = False) -> dict[str, Any]: + body = (text or "").strip() + return { + "ok": True, + "stub": stub, + "mode": MODE.value, + "text": body, + "empty": not bool(body), + "error": "", } def invoke_shell(*, payload: dict[str, Any]) -> dict[str, Any]: + """ + 无网/单测入口:assistant_text 当作图上可见原文。 + + 没有该键时保持失败码,避免假装已经读图。 + """ + if "assistant_text" in (payload or {}): + packed = _pack_text(str(payload.get("assistant_text") or ""), stub=True) + return packed return { "ok": False, "stub": True, "mode": MODE.value, + "text": "", + "empty": True, "error": "vision_shell_not_implemented", - "echo_keys": sorted(payload.keys()), + "echo_keys": sorted((payload or {}).keys()), } + + +def _vision_messages(*, image_b64: str, mime: str, image_url: str) -> list[dict[str, Any]]: + """拼兼容 OpenAI 的图文消息。优先 data URL,其次 http 图链。""" + url = (image_url or "").strip() + if not url and image_b64: + kind = (mime or "image/jpeg").strip() or "image/jpeg" + url = f"data:{kind};base64,{image_b64}" + user_content: list[dict[str, Any]] = [] + if url: + user_content.append({"type": "image_url", "image_url": {"url": url}}) + user_content.append({"type": "text", "text": "请只抄写图上可见文字。"}) + return [ + {"role": "system", "content": VISION_SYSTEM}, + {"role": "user", "content": user_content}, + ] + + +def invoke_vision( + *, + payload: dict[str, Any], + chat_fn: Optional[Any] = None, + allow_network: bool = False, +) -> dict[str, Any]: + """ + 读一张图,返回可见原文。 + + chat_fn:单测注入,签名同 LlmHttpClient.chat。 + allow_network=False 且无 chat_fn:走 invoke_shell,不打真网。 + """ + data = dict(payload or {}) + image_b64 = str(data.get("image_b64") or "") + image_url = str(data.get("image_url") or "") + mime = str(data.get("mime") or "").strip() or "image/jpeg" + if chat_fn is None and not allow_network: + return invoke_shell(payload=data) + if not image_b64 and not image_url: + logger.info("vision 拒绝:没有图") + return { + "ok": False, + "stub": False, + "mode": MODE.value, + "text": "", + "empty": True, + "error": "no_image", + } + if chat_fn is None: + from agent.llm.http_client import LlmHttpClient + + client = LlmHttpClient.from_settings() + chat_fn = client.chat + messages = _vision_messages( + image_b64=image_b64, + mime=mime, + image_url=image_url, + ) + result = chat_fn(mode=MODE, messages=messages) + if not getattr(result, "ok", False): + err = str(getattr(result, "error", "") or "vision_failed") + logger.info("vision 失败 err=%s", err) + return { + "ok": False, + "stub": False, + "mode": MODE.value, + "text": "", + "empty": True, + "error": err, + } + packed = _pack_text(str(getattr(result, "content", "") or "")) + packed["stub"] = False + logger.info( + "vision 读完 ok=%s empty=%s text_len=%s mime=%s", + True, + packed.get("empty"), + len(str(packed.get("text") or "")), + mime, + ) + return packed diff --git a/inquiry-agent/agent/policy/air_text_flow.py b/inquiry-agent/agent/policy/air_text_flow.py index c756162..45320f9 100644 --- a/inquiry-agent/agent/policy/air_text_flow.py +++ b/inquiry-agent/agent/policy/air_text_flow.py @@ -12,7 +12,7 @@ from __future__ import annotations import logging import threading import uuid -from dataclasses import dataclass, field +from dataclasses import asdict, dataclass, field from typing import Any, Callable, Optional from agent.ledger.http_ledger import HttpLedger @@ -55,6 +55,43 @@ class FlowSession: pending_product_quote: dict[str, Any] = field(default_factory=dict) +def _session_payload(sess: FlowSession) -> dict[str, Any]: + """书签转 JSON。tuple 改 list,方便 HTTP/Worker 共用。""" + data = asdict(sess) + data["allowed"] = list(sess.allowed or ()) + data["invited_userids"] = list(sess.invited_userids or ()) + data["invited_product_names"] = list(sess.invited_product_names or ()) + data["facts"] = {str(k): str(v) for k, v in (sess.facts or {}).items()} + return data + + +def _session_from_payload(raw: dict[str, Any]) -> FlowSession: + """Redis JSON → 书签。缺字段用默认,避免旧键打崩。""" + return FlowSession( + sender_id=str(raw.get("sender_id") or ""), + thread_id=str(raw.get("thread_id") or raw.get("sender_id") or ""), + wait_version=int(raw.get("wait_version") or 0), + phase=str(raw.get("phase") or "idle"), + business_line=str(raw.get("business_line") or ""), + facts={str(k): str(v) for k, v in dict(raw.get("facts") or {}).items()}, + work_order_no=str(raw.get("work_order_no") or ""), + first_or_same=str(raw.get("first_or_same") or "first"), + history_work_order_no=str(raw.get("history_work_order_no") or ""), + quote=dict(raw.get("quote") or {}), + status=str(raw.get("status") or ""), + allowed=tuple(str(x) for x in (raw.get("allowed") or ())), + immutable_text=str(raw.get("immutable_text") or ""), + collab_chat_id=str(raw.get("collab_chat_id") or ""), + collab_facts={str(k): str(v) for k, v in dict(raw.get("collab_facts") or {}).items()}, + invited_userids=tuple(str(x) for x in (raw.get("invited_userids") or ())), + invited_sales_name=str(raw.get("invited_sales_name") or ""), + invited_product_names=tuple(str(x) for x in (raw.get("invited_product_names") or ())), + deal_version=int(raw.get("deal_version") or 0), + collab_ended=bool(raw.get("collab_ended")), + pending_product_quote=dict(raw.get("pending_product_quote") or {}), + ) + + class AirTextInquiryFlow: """ 空运文字询价 Owner。海运走 SeaTextInquiryFlow,陆运走 LandTextInquiryFlow。 @@ -63,12 +100,15 @@ class AirTextInquiryFlow: 价格卡按工单号另存一份当时会话:点哪张卡就生成哪张卡当时的报价单。 """ - def __init__(self, ledger: Any = None, group_client: Any = None) -> None: + bookmark_kind = "air" + + def __init__(self, ledger: Any = None, group_client: Any = None, *, bookmark_store: Any = None) -> None: self._ledger = ledger or build_ledger() self._lock = threading.Lock() self._by_sender: dict[str, FlowSession] = {} self._by_ticket: dict[str, FlowSession] = {} self._group_client = group_client + self._bookmark_store = bookmark_store @property def ledger(self) -> Any: @@ -83,8 +123,16 @@ class AirTextInquiryFlow: return self._group_client def session_of(self, sender_id: str) -> Optional[FlowSession]: + sid = (sender_id or "").strip() with self._lock: - return self._by_sender.get(sender_id) + hit = self._by_sender.get(sid) + if hit is not None: + return hit + loaded = self._load_bookmark(sid) + if loaded is None: + return None + self._save(loaded) + return loaded def session_by_work_order(self, work_order_no: str) -> Optional[FlowSession]: """用工单号找回书签,不校验 sender。群成交卡校验用。""" @@ -126,17 +174,14 @@ class AirTextInquiryFlow: 必须校验 sender_id,禁止用别人的工单号串会话。 """ no = copy.work_order_from_card_meta(card_meta, text=text) + if not no: + return self.session_of(sender_id) with self._lock: - if no: - hit = self._by_ticket.get(no) - if hit and hit.sender_id == sender_id: - return hit - elif not no: - return self._by_sender.get(sender_id) - if no: - # 卡上有工单号就只动这一单;进程刚重启时从主账收回,禁止落到别人/最新单。 - return self._load_ticket_session(sender_id, no) - return None + hit = self._by_ticket.get(no) + if hit and hit.sender_id == sender_id: + return hit + # 卡上有工单号就只动这一单;进程刚重启时从主账收回,禁止落到别人/最新单。 + return self._load_ticket_session(sender_id, no) def _save(self, sess: FlowSession) -> None: with self._lock: @@ -144,6 +189,53 @@ class AirTextInquiryFlow: no = (sess.work_order_no or "").strip() if no: self._by_ticket[no] = sess + self._persist_bookmark(sess) + + def drop_current_bookmark(self, sender_id: str) -> Optional[FlowSession]: + """ + 发图新开:丢掉当前书签,已出号会话仍留在 _by_ticket。 + + 未出号的补问/核对草稿从此不再被文字当「当前单」。 + HTTP 与 Worker 必须同时丢掉 Redis 里那份,否则补「托盘」会问运输方式。 + """ + sid = (sender_id or "").strip() + if not sid: + return None + with self._lock: + old = self._by_sender.pop(sid, None) + self._delete_bookmark(sid) + return old + + def _persist_bookmark(self, sess: FlowSession) -> None: + from agent.redis_coord.flow_bookmark import save_flow_bookmark + + save_flow_bookmark( + self.bookmark_kind, + sess.sender_id, + _session_payload(sess), + store=self._bookmark_store, + ) + + def _load_bookmark(self, sender_id: str) -> Optional[FlowSession]: + from agent.redis_coord.flow_bookmark import load_flow_bookmark + + raw = load_flow_bookmark( + self.bookmark_kind, + sender_id, + store=self._bookmark_store, + ) + if not raw: + return None + return _session_from_payload(raw) + + def _delete_bookmark(self, sender_id: str) -> None: + from agent.redis_coord.flow_bookmark import delete_flow_bookmark + + delete_flow_bookmark( + self.bookmark_kind, + sender_id, + store=self._bookmark_store, + ) def _load_ticket_session(self, sender_id: str, work_order_no: str) -> Optional[FlowSession]: """ diff --git a/inquiry-agent/agent/policy/image_wave.py b/inquiry-agent/agent/policy/image_wave.py new file mode 100644 index 0000000..988fdea --- /dev/null +++ b/inquiry-agent/agent/policy/image_wave.py @@ -0,0 +1,386 @@ +""" +私聊图片波次账。 + +本文件职责:记住「这一波」图和紧挨的字,供图片询价拼口播。 +属于 policy 包。生产波次内容落 Redis(HTTP 与 Worker 共用);单测可纯内存。 +禁止:调模型、建单、写六态、在这里等几秒;禁止一把大锁串全部销售。 + +10 秒只约束发图前那一句。发图之后只看业务结果出没出。 +""" + +from __future__ import annotations + +import json +import logging +import threading +import uuid +from dataclasses import asdict, dataclass, field +from typing import Any, Optional, Protocol + +# 规格:发图前紧挨的一句字,间隔不超过 10 秒才带到新单。 +PRECEDE_SECONDS = 10.0 +# 波次短协调:超时后当新开,不当六态。 +_WAVE_TTL_SEC = 2 * 3600 +logger = logging.getLogger(__name__) + + +class WaveKV(Protocol): + """HTTP 与 Worker 共用的短存。单测可注入内存假实现。""" + + def get_json(self, kind: str, sender_id: str) -> Optional[dict[str, Any]]: ... + + def set_json(self, kind: str, sender_id: str, data: dict[str, Any]) -> None: ... + + def delete(self, kind: str, sender_id: str) -> None: ... + + +@dataclass +class LastPrivateText: + """同一发送人最近一句私聊文字。""" + + sender_id: str + text: str + ts: float + + +@dataclass +class ImageWave: + """ + 一张新询价对应的图片波次。 + + pending:还在认的图张数;为 0 且未 business_done 时才能出业务结果。 + """ + + sender_id: str + wave_id: str + ack_sent: bool = False + business_done: bool = False + media_ids: list[str] = field(default_factory=list) + vision_texts: list[str] = field(default_factory=list) + pending: int = 0 + precede_text: str = "" + follow_text: str = "" + object_keys: list[str] = field(default_factory=list) + filenames: list[str] = field(default_factory=list) + errors: list[str] = field(default_factory=list) + + +class ImageWaveBook: + """ + 图片波次账。按 sender 一把锁,不同销售互不等。 + + 生产必须挂 Redis:收图在 HTTP 进程、识图在 Worker,两边不能各记一份内存。 + 单测不传 kv,仍走进程内字典。 + """ + + def __init__(self, kv: Optional[WaveKV] = None) -> None: + self._lock = threading.Lock() + self._kv = kv + self._last_text: dict[str, LastPrivateText] = {} + self._waves: dict[str, ImageWave] = {} + + def _hydrate(self, sender_id: str) -> None: + """从 Redis 覆盖本进程缓存。无 kv 时不做事。""" + sid = (sender_id or "").strip() + if not sid or self._kv is None: + return + try: + raw_wave = self._kv.get_json("wave", sid) + if raw_wave: + self._waves[sid] = _wave_from_dict(raw_wave) + else: + self._waves.pop(sid, None) + raw_text = self._kv.get_json("text", sid) + if raw_text and raw_text.get("text"): + self._last_text[sid] = LastPrivateText( + sender_id=sid, + text=str(raw_text.get("text") or ""), + ts=float(raw_text.get("ts") or 0.0), + ) + else: + self._last_text.pop(sid, None) + except Exception: + logger.exception("读图片波次失败 sender=%s", sid) + + def _persist(self, sender_id: str) -> None: + """把当前 sender 的波次和前一句写回 Redis。""" + sid = (sender_id or "").strip() + if not sid or self._kv is None: + return + try: + wave = self._waves.get(sid) + if wave is None: + self._kv.delete("wave", sid) + else: + self._kv.set_json("wave", sid, asdict(wave)) + row = self._last_text.get(sid) + if row is None: + self._kv.delete("text", sid) + else: + self._kv.set_json("text", sid, {"text": row.text, "ts": row.ts}) + except Exception: + logger.exception("写图片波次失败 sender=%s", sid) + + def note_text(self, sender_id: str, text: str, now: float) -> None: + """记下最近一句私聊文字,供 10 秒内发图带走。空句不记。""" + body = (text or "").strip() + sid = (sender_id or "").strip() + if not sid or not body: + return + with self._lock: + self._hydrate(sid) + self._last_text[sid] = LastPrivateText(sender_id=sid, text=body, ts=float(now)) + self._persist(sid) + + def forget_text(self, sender_id: str) -> None: + """ + 图片口播不是销售刚打的字,认完后清掉,避免 10 秒内再发图把旧口播带走。 + """ + sid = (sender_id or "").strip() + if not sid: + return + with self._lock: + self._hydrate(sid) + self._last_text.pop(sid, None) + self._persist(sid) + + def precede_text(self, sender_id: str, now: float) -> str: + """发图前那一句:超过 10 秒返回空。""" + sid = (sender_id or "").strip() + with self._lock: + self._hydrate(sid) + row = self._last_text.get(sid) + if row is None: + return "" + if float(now) - row.ts > PRECEDE_SECONDS: + return "" + return row.text + + def active(self, sender_id: str) -> Optional[ImageWave]: + """未出业务结果的当前波;没有则 None。""" + sid = (sender_id or "").strip() + with self._lock: + self._hydrate(sid) + wave = self._waves.get(sid) + if wave is None or wave.business_done: + return None + return wave + + def open_or_join( + self, sender_id: str, media_id: str, *, now: float + ) -> tuple[ImageWave, str]: + """ + 第一张图开波;业务发出前的后续图并入;发出后再发图开新波。 + + 返回 how:opened / joined / new_after_done。 + 副作用:pending +1。不发送企微。 + """ + sid = (sender_id or "").strip() + mid = (media_id or "").strip() + precede = self.precede_text(sid, now) + with self._lock: + self._hydrate(sid) + cur = self._waves.get(sid) + if cur is not None and not cur.business_done: + if mid: + cur.media_ids.append(mid) + cur.pending += 1 + self._persist(sid) + return cur, "joined" + how = "new_after_done" if cur is not None and cur.business_done else "opened" + wave = ImageWave( + sender_id=sid, + wave_id=uuid.uuid4().hex, + media_ids=[mid] if mid else [], + pending=1, + precede_text=precede, + ) + self._waves[sid] = wave + self._persist(sid) + return wave, how + + def mark_ack(self, sender_id: str) -> None: + """已经回过「正在识别图片」,连发第二张不得再 ack。""" + sid = (sender_id or "").strip() + with self._lock: + self._hydrate(sid) + wave = self._waves.get(sid) + if wave is not None: + wave.ack_sent = True + self._persist(sid) + + def append_follow_text(self, sender_id: str, text: str) -> bool: + """业务结果发出前的字并进本波。发出后返回 False,交给文字询价。""" + body = (text or "").strip() + sid = (sender_id or "").strip() + if not sid or not body: + return False + with self._lock: + self._hydrate(sid) + wave = self._waves.get(sid) + if wave is None or wave.business_done: + return False + if wave.follow_text: + wave.follow_text = f"{wave.follow_text}\n{body}" + else: + wave.follow_text = body + self._persist(sid) + return True + + def add_vision( + self, + sender_id: str, + text: str, + *, + object_key: str = "", + filename: str = "", + error: str = "", + ) -> ImageWave: + """ + 一张图认完。pending 减一,可见原文追加。 + + 调用:Worker 识别槽。不得在本方法出发业务回复。 + """ + sid = (sender_id or "").strip() + with self._lock: + self._hydrate(sid) + wave = self._waves.get(sid) + if wave is None: + wave = ImageWave(sender_id=sid, wave_id=uuid.uuid4().hex, pending=1) + self._waves[sid] = wave + body = (text or "").strip() + if body: + wave.vision_texts.append(body) + if object_key: + wave.object_keys.append(object_key) + if filename: + wave.filenames.append(filename) + if error: + wave.errors.append(error) + if wave.pending > 0: + wave.pending -= 1 + self._persist(sid) + return wave + + def compose_oral(self, sender_id: str) -> str: + """precede + 各张可见原文 + 后到的字,换行拼接。""" + sid = (sender_id or "").strip() + with self._lock: + self._hydrate(sid) + wave = self._waves.get(sid) + if wave is None: + return "" + parts = [] + if wave.precede_text.strip(): + parts.append(wave.precede_text.strip()) + for chunk in wave.vision_texts: + if chunk.strip(): + parts.append(chunk.strip()) + if wave.follow_text.strip(): + parts.append(wave.follow_text.strip()) + return "\n".join(parts) + + def mark_business_done(self, sender_id: str) -> None: + """业务补问/卡片已发出。之后再发图必须新开。""" + sid = (sender_id or "").strip() + with self._lock: + self._hydrate(sid) + wave = self._waves.get(sid) + if wave is not None: + wave.business_done = True + wave.pending = 0 + self._persist(sid) + + +_BOOK: Optional[ImageWaveBook] = None +_BOOK_GUARD = threading.Lock() + + +def _wave_from_dict(raw: dict[str, Any]) -> ImageWave: + """Redis JSON → 波次对象。缺字段用默认,避免旧键打崩。""" + return ImageWave( + sender_id=str(raw.get("sender_id") or ""), + wave_id=str(raw.get("wave_id") or uuid.uuid4().hex), + ack_sent=bool(raw.get("ack_sent")), + business_done=bool(raw.get("business_done")), + media_ids=[str(x) for x in (raw.get("media_ids") or [])], + vision_texts=[str(x) for x in (raw.get("vision_texts") or [])], + pending=int(raw.get("pending") or 0), + precede_text=str(raw.get("precede_text") or ""), + follow_text=str(raw.get("follow_text") or ""), + object_keys=[str(x) for x in (raw.get("object_keys") or [])], + filenames=[str(x) for x in (raw.get("filenames") or [])], + errors=[str(x) for x in (raw.get("errors") or [])], + ) + + +class RedisWaveKV: + """ + 把图片波次落到本前缀 Redis,TTL 2 小时。 + + HTTP 收图、Worker 认图必须读同一把键,否则第二张图会丢。 + """ + + def __init__(self, client: Any) -> None: + self._client = client + + def get_json(self, kind: str, sender_id: str) -> Optional[dict[str, Any]]: + from agent.redis_coord.keys import IMAGE_WAVE, IMAGE_WAVE_TEXT + + suffix = IMAGE_WAVE if kind == "wave" else IMAGE_WAVE_TEXT + raw = self._client.raw.get(self._client.key(suffix, sender_id)) + if not raw: + return None + if isinstance(raw, bytes): + raw = raw.decode("utf-8") + data = json.loads(raw) + return dict(data) if isinstance(data, dict) else None + + def set_json(self, kind: str, sender_id: str, data: dict[str, Any]) -> None: + from agent.redis_coord.keys import IMAGE_WAVE, IMAGE_WAVE_TEXT + + suffix = IMAGE_WAVE if kind == "wave" else IMAGE_WAVE_TEXT + self._client.raw.set( + self._client.key(suffix, sender_id), + json.dumps(data, ensure_ascii=False), + ex=_WAVE_TTL_SEC, + ) + + def delete(self, kind: str, sender_id: str) -> None: + from agent.redis_coord.keys import IMAGE_WAVE, IMAGE_WAVE_TEXT + + suffix = IMAGE_WAVE if kind == "wave" else IMAGE_WAVE_TEXT + self._client.raw.delete(self._client.key(suffix, sender_id)) + + +def _production_kv() -> Optional[WaveKV]: + """生产挂 Redis;连不上则退回内存并打日志。""" + try: + from agent.redis_coord.runtime import get_redis_runtime + + rt = get_redis_runtime() + if rt is None or rt.client is None: + return None + return RedisWaveKV(rt.client) + except Exception: + logger.exception("图片波次 Redis 不可用,本进程退回内存账(双进程会丢第二张图)") + return None + + +def get_image_wave_book() -> ImageWaveBook: + """ + 进程内一份句柄;生产波次内容在 Redis,HTTP 与 Worker 才能对上。 + """ + global _BOOK + with _BOOK_GUARD: + if _BOOK is None: + _BOOK = ImageWaveBook(kv=_production_kv()) + return _BOOK + + +def reset_image_wave_book_for_test(kv: Optional[WaveKV] = None) -> ImageWaveBook: + """单测重置。传入同一 kv 可模拟 HTTP/Worker 两份句柄。""" + global _BOOK + with _BOOK_GUARD: + _BOOK = ImageWaveBook(kv=kv) + return _BOOK diff --git a/inquiry-agent/agent/policy/inquiry_copy.py b/inquiry-agent/agent/policy/inquiry_copy.py index 233f10f..615d72b 100644 --- a/inquiry-agent/agent/policy/inquiry_copy.py +++ b/inquiry-agent/agent/policy/inquiry_copy.py @@ -100,6 +100,24 @@ _CONTINUE_HINT = ( ASK_TRANSPORT = "请补充运输方式:海运 / 空运 / 陆运。" +# 私聊图片询价:先回识别中;看不清/失败与补问拼成一条,销售侧不加 [test]。 +IMAGE_RECOGNIZING = "正在识别图片" +IMAGE_UNREADABLE = "图片没看清询价信息,请换一张更清楚的,或直接打字。" +IMAGE_TECH_FAIL = "图片暂时没认出来,请再发一张或直接打字。" + + +def image_fail_then_clarify(reason: str, clarify: str) -> str: + """ + 看不清或技术失败时,说明 + 文字询价补问写在同一条回复。 + + 禁止只回说明就把人丢在那里。两端空白去掉,中间空一行。 + """ + head = (reason or "").strip() + body = (clarify or "").strip() + if head and body: + return f"{head}\n\n{body}" + return head or body + # 旧一句只问运输类型;销售要一次看到类型、线路、分类。残留调用也走完整清单。 ASK_LAND_SUBTYPE = _land_option_prompt(current={}) @@ -195,6 +213,24 @@ def land_group_pair_rejected() -> str: return "拉群失败,请排查填写内容,并重新发起工单。" +def _land_optional_keys(land_type: str) -> tuple[str, ...]: + """当前类型询价非必填内部键。国内没有通关口岸。""" + if land_type in _LAND_CUSTOMS_TYPES: + return ("通关口岸",) + return () + + +def _land_field_line(key: str, value: str, *, required: bool, example: str = "") -> str: + """销售侧一行:名称(必填/非必填):值或示例。""" + label = LAND_DISPLAY.get(key, key) + mark = "必填" if required else "非必填" + if value: + return f"{label}({mark}):{value}" + if example: + return f"{label}({mark}):(如:{example})" + return f"{label}({mark}):" + + def ask_land_fields( *, facts: dict[str, str] | None = None, @@ -204,30 +240,37 @@ def ask_land_fields( """ 陆运货物缺项补问。不出核对卡、不出现确定句。 + 已识别和待补充都标明必填/非必填。跨境/中港空着的通关口岸也列进待补充。 报价日期不进待补充。 """ src = dict(facts or {}) skip = {"报价日期", "运输类型", "线路类别", "运输分类"} missing = [k for k in missing_keys if k and k not in skip] land_type = (land_subtype or src.get("运输类型") or "").strip() + optional = _land_optional_keys(land_type) lines = ["请确认并补充以下信息:", "运输方式:陆运"] if land_type: lines.append(f"运输类型:{land_type}") for key in _land_cargo_keys(land_type, src): val = (src.get(key) or "").strip() if val: - lines.append(f"{LAND_DISPLAY.get(key, key)}:{val}") + lines.append(_land_field_line(key, val, required=True)) + for key in optional: + val = (src.get(key) or "").strip() + if val: + lines.append(_land_field_line(key, val, required=False)) + pending = list(missing) + for key in optional: + if not (src.get(key) or "").strip() and key not in pending: + pending.append(key) + if not pending: + return "" lines.append("") lines.append("待补充:") - if not missing: - return "" - for key in missing: - label = display_name(key, "LAND") + for key in pending: + required = key not in optional example = _LAND_EXAMPLES.get(key, "") - if example: - lines.append(f"{label}:(如:{example})") - else: - lines.append(f"{label}:") + lines.append(_land_field_line(key, "", required=required, example=example)) return "\n".join(lines) @@ -367,7 +410,8 @@ def inquiry_field_lines( """ 询价确认字段行(销售可见)。 - 起运/目的展示销售原文(南京/北京),三字码只给出站 TMS,不对销售展示。 + 起运/目的展示销售原文:说南京就南京,图上/原话是 SZX 就展示 SZX。 + originCode 只给出站 TMS,不拿来改写成城市名再给销售看。 """ origin = (facts.get("起运港") or facts.get("起运地") or "-").strip() dest = (facts.get("目的港") or facts.get("目的地") or "-").strip() diff --git a/inquiry-agent/agent/policy/land_text_flow.py b/inquiry-agent/agent/policy/land_text_flow.py index db93440..a06fd3b 100644 --- a/inquiry-agent/agent/policy/land_text_flow.py +++ b/inquiry-agent/agent/policy/land_text_flow.py @@ -41,6 +41,8 @@ class LandTextInquiryFlow(SeaTextInquiryFlow): 前半段与海运不同:核对通过前不建单。后半段套海运建群,匹配改走线路类别。 """ + bookmark_kind = "land" + def on_text( self, *, diff --git a/inquiry-agent/agent/policy/sea_text_flow.py b/inquiry-agent/agent/policy/sea_text_flow.py index a240c71..03b680b 100644 --- a/inquiry-agent/agent/policy/sea_text_flow.py +++ b/inquiry-agent/agent/policy/sea_text_flow.py @@ -40,6 +40,8 @@ class SeaTextInquiryFlow(AirTextInquiryFlow): 复用空运后半段(跳过协同 / 出方案卡 / Excel PDF / 成交),只改查价前与卡片、拉群。 """ + bookmark_kind = "sea" + def __init__(self, ledger: Any = None, group_client: Any = None) -> None: super().__init__(ledger=ledger) self._group_client = group_client diff --git a/inquiry-agent/agent/redis_coord/flow_bookmark.py b/inquiry-agent/agent/redis_coord/flow_bookmark.py new file mode 100644 index 0000000..4136620 --- /dev/null +++ b/inquiry-agent/agent/redis_coord/flow_bookmark.py @@ -0,0 +1,133 @@ +""" +询价当前书签短协调:HTTP 收字、Worker 识图后补问必须读同一份。 + +不当六态、不当 checkpoint。只存当前 sender 的补问/核对草稿。 +禁止:写入现网 Redis 前缀;把已出号工单只放这里当主账。 +""" + +from __future__ import annotations + +import json +import logging +from typing import Any, Optional, Protocol + +from agent.redis_coord.keys import FLOW_BOOKMARK +from agent.redis_coord.runtime import get_redis_runtime + +logger = logging.getLogger(__name__) + +_TTL_SEC = 24 * 3600 + + +class BookmarkStore(Protocol): + """单测可注入;生产走 Redis。""" + + def save(self, kind: str, sender_id: str, payload: dict[str, Any]) -> None: ... + + def load(self, kind: str, sender_id: str) -> Optional[dict[str, Any]]: ... + + def delete(self, kind: str, sender_id: str) -> None: ... + + +class DictBookmarkStore: + """两份 Flow 实例共用同一字典,模拟 HTTP/Worker。""" + + def __init__(self) -> None: + self._d: dict[str, dict[str, Any]] = {} + + def save(self, kind: str, sender_id: str, payload: dict[str, Any]) -> None: + self._d[f"{kind}:{sender_id}"] = json.loads(json.dumps(payload)) + + def load(self, kind: str, sender_id: str) -> Optional[dict[str, Any]]: + row = self._d.get(f"{kind}:{sender_id}") + return json.loads(json.dumps(row)) if row is not None else None + + def delete(self, kind: str, sender_id: str) -> None: + self._d.pop(f"{kind}:{sender_id}", None) + + +def _redis_ok() -> bool: + try: + return get_redis_runtime().client.backend == "redis" + except Exception: + return False + + +def save_flow_bookmark( + kind: str, + sender_id: str, + payload: dict[str, Any], + *, + store: Optional[BookmarkStore] = None, +) -> None: + """写下当前书签。store 优先;否则仅真 Redis。""" + sid = (sender_id or "").strip() + name = (kind or "").strip() or "air" + if not sid or not isinstance(payload, dict): + return + if store is not None: + store.save(name, sid, payload) + return + if not _redis_ok(): + return + try: + client = get_redis_runtime().client + client.raw.set( + client.key(FLOW_BOOKMARK, name, sid), + json.dumps(payload, ensure_ascii=False), + ex=_TTL_SEC, + ) + except Exception: + logger.exception("写询价书签失败 kind=%s sender=%s", name, sid) + + +def load_flow_bookmark( + kind: str, + sender_id: str, + *, + store: Optional[BookmarkStore] = None, +) -> dict[str, Any]: + """取出当前书签。没有则空。""" + sid = (sender_id or "").strip() + name = (kind or "").strip() or "air" + if not sid: + return {} + if store is not None: + return dict(store.load(name, sid) or {}) + if not _redis_ok(): + return {} + try: + client = get_redis_runtime().client + raw = client.raw.get(client.key(FLOW_BOOKMARK, name, sid)) + if not raw: + return {} + if isinstance(raw, bytes): + raw = raw.decode("utf-8") + data = json.loads(raw) + return dict(data) if isinstance(data, dict) else {} + except Exception: + logger.exception("读书签失败 kind=%s sender=%s", name, sid) + return {} + + +def delete_flow_bookmark( + kind: str, + sender_id: str, + *, + store: Optional[BookmarkStore] = None, +) -> None: + """发图新开或结束当前草稿时删掉。已出号会话不靠这把键。""" + sid = (sender_id or "").strip() + name = (kind or "").strip() or "air" + if not sid: + return + if store is not None: + store.delete(name, sid) + return + if not _redis_ok(): + return + try: + client = get_redis_runtime().client + client.raw.delete(client.key(FLOW_BOOKMARK, name, sid)) + except Exception: + logger.exception("删书签失败 kind=%s sender=%s", name, sid) diff --git a/inquiry-agent/agent/redis_coord/keys.py b/inquiry-agent/agent/redis_coord/keys.py index 8845e64..26a4816 100644 --- a/inquiry-agent/agent/redis_coord/keys.py +++ b/inquiry-agent/agent/redis_coord/keys.py @@ -57,6 +57,11 @@ QUOTE_FILE_READY = "quote_file_ready" PENDING_PRODUCT_QUOTE = "pending_product_quote" # 陆运核对卡已出、等确定查价:短协调,不当六态、不当工单号 PENDING_LAND_CONFIRM = "pending_land_confirm" +# 空运/海运/陆运当前书签:HTTP 与 Worker 共用,不当六态 +FLOW_BOOKMARK = "flow_bookmark" +# 私聊图片波次:HTTP 与 Worker 必须共用,不当六态 +IMAGE_WAVE = "image_wave" +IMAGE_WAVE_TEXT = "image_wave_text" # 空运群同一句:BOT 与存档补听去重,不当六态 AIR_INBOUND_MSG = "air_inbound_msg" AIR_INBOUND_TEXT = "air_inbound_text" diff --git a/inquiry-agent/agent/routing/dispatch.py b/inquiry-agent/agent/routing/dispatch.py index 561ab9f..a00821c 100644 --- a/inquiry-agent/agent/routing/dispatch.py +++ b/inquiry-agent/agent/routing/dispatch.py @@ -14,6 +14,7 @@ from agent.channel.wecom.models import InboundMessage from agent.handlers import abnormal, card_action, close_deal, continue_thread, new_inquiry from agent.handlers.echo_text import handle_echo_text from agent.handlers.group_collab import handle_group_collab +from agent.handlers.image_inquiry import handle_image_inquiry from agent.handlers.text_inquiry import handle_text_inquiry, handle_text_other, is_land_confirm_action from agent.routing.decision import RouteDecision @@ -56,6 +57,11 @@ def _wrap_new(message: InboundMessage) -> None: new_inquiry.handle_new_inquiry(message) +def _wrap_image(message: InboundMessage) -> None: + """私聊图片询价:认图后当口播新开,复用文字流程。""" + handle_image_inquiry(message) + + def _wrap_continue(message: InboundMessage, *, thread_id: str = "") -> None: continue_thread.handle_continue_thread(message, thread_id=thread_id) @@ -112,7 +118,7 @@ HANDLER_REGISTRY: dict[str, HandlerFn] = { "ordinary_text_inquiry": _wrap_text_inquiry, "ordinary_text_other": _wrap_text_other, "attachment_inquiry": _wrap_new, - "image_inquiry": _wrap_new, + "image_inquiry": _wrap_image, "multi_segment_transport": _wrap_new, "tms_read_only_quote": _wrap_continue, "new_inquiry": _wrap_new, @@ -203,6 +209,18 @@ def dispatch_inbound(message: InboundMessage, decision: RouteDecision) -> str: HANDLER_REGISTRY["reject_no_sender"](message) return "reject_no_sender" + kind = (message.msg_type or "text").lower() + if kind == "image": + handle_image_inquiry(message) + return "image_inquiry" + if kind in {"text", ""}: + from agent.handlers.image_inquiry import append_wave_text, is_active_image_wave + + if is_active_image_wave(message.sender_id) and (message.content or "").strip(): + append_wave_text(message.sender_id, message.content or "") + logger.info("dispatch 图片波次收字 sender=%s", message.sender_id) + return "image_wave_follow_text" + fn = resolve_handler(decision.intent) if fn is None: logger.warning("dispatch 未知意图 intent=%s → noop", decision.intent) diff --git a/inquiry-agent/agent/schema/field_validate.py b/inquiry-agent/agent/schema/field_validate.py index 6588234..2798f2b 100644 --- a/inquiry-agent/agent/schema/field_validate.py +++ b/inquiry-agent/agent/schema/field_validate.py @@ -189,6 +189,8 @@ _ALIASES = { "客户名称": "客户名称", "货值": "货值", "是否为危险品": "是否为危险品", + "通关口岸": "通关口岸", + "通关口岸(非必填)": "通关口岸", } diff --git a/inquiry-agent/agent/schema/tms_air_query.py b/inquiry-agent/agent/schema/tms_air_query.py index 46e5d06..1b28c60 100644 --- a/inquiry-agent/agent/schema/tms_air_query.py +++ b/inquiry-agent/agent/schema/tms_air_query.py @@ -47,6 +47,35 @@ _PACKS = { } _IATA = re.compile(r"^[A-Za-z]{3}$") +_IATA_PAIR = re.compile( + r"(? str: def harvest_named_air_route(text: str, facts: dict[str, str] | None = None) -> dict[str, str]: """ - 当前原话里的「从南京飞吉隆坡 / 北京飞吉隆坡」覆盖起运地、目的地。 + 当前原话里的航线覆盖起运地、目的地。 - 只认对照表里的地名,不猜没写过的港口。换城市必须丢掉旧三字码。 + 图上/原话写了 SZX-CGK 这类三字码对,必须保持英文,禁止改写成深圳/雅加达。 + 「从南京飞吉隆坡」仍按对照表收中文地名。换城市必须丢掉旧三字码。 """ merged = dict(facts or {}) raw = text or "" + pair = _IATA_PAIR.search(raw) + if pair: + origin_code = pair.group(1).upper() + dest_code = pair.group(2).upper() + blocked = origin_code in _IATA_PAIR_BLOCK or dest_code in _IATA_PAIR_BLOCK + if origin_code != dest_code and not blocked: + merged["起运港"] = origin_code + merged["目的港"] = dest_code + merged.pop("originCode", None) + merged.pop("destinationCode", None) + return merged names = sorted(_AIRPORTS.keys(), key=len, reverse=True) for origin in names: for dest in names: diff --git a/inquiry-agent/agent/worker_app.py b/inquiry-agent/agent/worker_app.py index 22dbd15..d6349b4 100644 --- a/inquiry-agent/agent/worker_app.py +++ b/inquiry-agent/agent/worker_app.py @@ -30,6 +30,8 @@ def main() -> None: level=logging.INFO, format="%(asctime)s %(levelname)s [%(name)s] %(message)s", ) + # Worker 自己下素材,必须关掉 httpx INFO,否则 access_token 会进日志。 + logging.getLogger("httpx").setLevel(logging.WARNING) settings = get_settings() graph_runtime = get_graph_runtime(settings) redis_runtime = get_redis_runtime(settings) diff --git a/inquiry-agent/tests/test_air_text_inquiry.py b/inquiry-agent/tests/test_air_text_inquiry.py index e283163..fe22c86 100644 --- a/inquiry-agent/tests/test_air_text_inquiry.py +++ b/inquiry-agent/tests/test_air_text_inquiry.py @@ -149,6 +149,33 @@ class AirTextInquiryTests(unittest.TestCase): self.assertNotIn("请补充运输方式", joined) self.assertEqual(self.flow.session_of("u_fly").business_line, "AIR") + def test_worker_clarify_http_pallet_keeps_air(self) -> None: + """识图在 Worker 写下补问,HTTP 补「托盘」必须仍是空运,不能再问运输方式。""" + from agent.redis_coord.flow_bookmark import DictBookmarkStore + + store = DictBookmarkStore() + worker = AirTextInquiryFlow(MemoryLedger(), bookmark_store=store) + http = AirTextInquiryFlow(MemoryLedger(), bookmark_store=store) + replies: list[str] = [] + + def reply(text: str, extra=None) -> None: + replies.append(text) + + phase = worker.on_text( + sender_id="u_img", + text="空运 深圳到雅加达 发光二级管组件 11P 2400KG 20cbm", + reply=reply, + ) + self.assertEqual(phase, "clarify") + replies.clear() + phase = http.on_text(sender_id="u_img", text="托盘", reply=reply) + self.assertNotEqual(phase, "need_mode") + self.assertNotIn("请补充运输方式", "\n".join(replies)) + sess = http.session_of("u_img") + self.assertIsNotNone(sess) + self.assertEqual(sess.business_line, "AIR") + self.assertEqual(sess.facts.get("包装方式"), "托盘") + def test_missing_mode_asks_transport(self) -> None: phase = self.flow.on_text( sender_id="u1", diff --git a/inquiry-agent/tests/test_extract_oral.py b/inquiry-agent/tests/test_extract_oral.py index b1ed629..96e5aae 100644 --- a/inquiry-agent/tests/test_extract_oral.py +++ b/inquiry-agent/tests/test_extract_oral.py @@ -15,7 +15,7 @@ if ROOT not in sys.path: os.environ.setdefault("LEDGER_BACKEND", "memory") -from agent.llm.extract_text import extract_inquiry_snapshot +from agent.llm.extract_text import extract_inquiry_snapshot, parse_labeled_facts from agent.llm.mode_extract_fields import invoke_shell, parse_extract_payload from agent.policy.air_text_flow import AirTextInquiryFlow from agent.ledger.memory_ledger import MemoryLedger @@ -65,6 +65,22 @@ class OralExtractTests(unittest.TestCase): self.assertEqual(snap["business_line"], "AIR") self.assertFalse(snap["facts"].get("起运港")) + def test_labeled_customs_port(self) -> None: + self.assertEqual(parse_labeled_facts("通关口岸:皇岗").get("通关口岸"), "皇岗") + snap = extract_inquiry_snapshot("通关口岸(非必填):皇岗") + self.assertEqual(snap["facts"].get("通关口岸"), "皇岗") + snap2 = extract_inquiry_snapshot("通关口岸皇岗") + self.assertEqual(snap2["facts"].get("通关口岸"), "皇岗") + + def test_b_keeps_land_customs_port(self) -> None: + parsed = parse_extract_payload( + { + "transport_mode": "LAND", + "fields": {"起运港": "广州", "通关口岸": "皇岗"}, + } + ) + self.assertEqual(parsed["facts"].get("通关口岸"), "皇岗") + def test_assemble_air_query_zhuhai_clark(self) -> None: from agent.schema.tms_air_query import assemble_air_query @@ -134,6 +150,46 @@ class OralExtractTests(unittest.TestCase): self.assertNotIn("起运地:NKG", shown) self.assertNotIn("起运地:南京", shown) + def test_iata_pair_kept_not_rewritten_to_city(self) -> None: + from agent.schema.field_validate import harvest_oral_measures + from agent.schema.tms_air_query import assemble_air_query + from agent.policy import inquiry_copy as copy + + oral = "空运 SZX-CGK 发光二级管组件 11P 2400KG 20cbm" + facts = harvest_oral_measures( + oral, + {"起运港": "深圳", "目的港": "雅加达", "品名": "发光二级管组件"}, + ) + self.assertEqual(facts["起运港"], "SZX") + self.assertEqual(facts["目的港"], "CGK") + assembled = assemble_air_query( + { + **facts, + "件数": "11P", + "毛重": "2400KG", + "体积": "20cbm", + "包装方式": "托盘", + "报价日期": "2026-09-18", + } + ) + self.assertEqual(assembled["facts"]["originCode"], "SZX") + self.assertEqual(assembled["facts"]["destinationCode"], "CGK") + card = copy.inquiry_card(work_order_no="WO1", facts=assembled["facts"]) + self.assertIn("起运地:SZX", card) + self.assertIn("目的地:CGK", card) + self.assertNotIn("深圳", card) + self.assertNotIn("雅加达", card) + + def test_currency_pair_is_not_air_route(self) -> None: + from agent.schema.field_validate import harvest_oral_measures + + facts = harvest_oral_measures( + "空运 货值 USD-CNY 茶叶", + {"起运港": "上海", "目的港": "洛杉矶"}, + ) + self.assertEqual(facts["起运港"], "上海") + self.assertEqual(facts["目的港"], "洛杉矶") + def test_oral_pack_and_quote_date_from_current_sentence(self) -> None: from agent.schema.field_validate import harvest_oral_measures @@ -225,6 +281,13 @@ class OralExtractTests(unittest.TestCase): self.assertEqual(norm["毛重"], "180") self.assertEqual(norm["体积"], "2.5") + def test_extract_prompt_keeps_iata(self) -> None: + from agent.llm.mode_extract_fields import build_extract_messages + + blob = str(build_extract_messages("空运 SZX-CGK")) + self.assertIn("SZX-CGK", blob) + self.assertIn("禁止翻译成城市名", blob) + def test_chat_fn_fixture_drives_flow(self) -> None: def chat_fn(**kwargs): return SimpleNamespace( diff --git a/inquiry-agent/tests/test_image_inquiry.py b/inquiry-agent/tests/test_image_inquiry.py new file mode 100644 index 0000000..15293e0 --- /dev/null +++ b/inquiry-agent/tests/test_image_inquiry.py @@ -0,0 +1,272 @@ +""" +私聊图片询价:识图当口播、新开一票、复用文字流程。 +""" + +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 +""" + +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() diff --git a/inquiry-agent/tests/test_image_wave.py b/inquiry-agent/tests/test_image_wave.py new file mode 100644 index 0000000..bbe7e26 --- /dev/null +++ b/inquiry-agent/tests/test_image_wave.py @@ -0,0 +1,103 @@ +""" +私聊图片波次:前一句 10 秒、连发图、业务发出前收字。 +""" + +from __future__ import annotations + +import json +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) + +from agent.policy.image_wave import ( + ImageWaveBook, + PRECEDE_SECONDS, + reset_image_wave_book_for_test, +) + + +class ImageWaveTests(unittest.TestCase): + def setUp(self) -> None: + self.book = reset_image_wave_book_for_test() + + def test_precede_window_is_ten_seconds(self) -> None: + self.assertEqual(PRECEDE_SECONDS, 10.0) + + def test_precede_within_10s(self) -> None: + self.book.note_text("s1", "空运 上海到纽约", now=100.0) + self.assertEqual(self.book.precede_text("s1", now=109.9), "空运 上海到纽约") + + def test_precede_after_10s_dropped(self) -> None: + self.book.note_text("s1", "空运 上海到纽约", now=100.0) + self.assertEqual(self.book.precede_text("s1", now=110.1), "") + + def test_forget_text_drops_precede(self) -> None: + self.book.note_text("s1", "空运 上海到纽约", now=100.0) + self.book.forget_text("s1") + self.assertEqual(self.book.precede_text("s1", now=100.5), "") + + def test_second_image_joins_before_business(self) -> None: + w1, how = self.book.open_or_join("s1", "mid-a", now=1.0) + self.assertEqual(how, "opened") + w2, how2 = self.book.open_or_join("s1", "mid-b", now=1.2) + self.assertEqual(how2, "joined") + self.assertEqual(w1.wave_id, w2.wave_id) + self.assertEqual(w2.media_ids, ["mid-a", "mid-b"]) + + def test_image_after_business_is_new_wave(self) -> None: + self.book.open_or_join("s1", "mid-a", now=1.0) + self.book.mark_business_done("s1") + w2, how = self.book.open_or_join("s1", "mid-b", now=2.0) + self.assertEqual(how, "new_after_done") + self.assertEqual(w2.media_ids, ["mid-b"]) + + def test_follow_text_before_business(self) -> None: + self.book.open_or_join("s1", "mid-a", now=1.0) + self.assertTrue(self.book.append_follow_text("s1", "海运")) + self.book.add_vision("s1", "宁波到汉堡") + oral = self.book.compose_oral("s1") + self.assertIn("海运", oral) + self.assertIn("宁波到汉堡", oral) + + def test_follow_text_after_business_rejected(self) -> None: + self.book.open_or_join("s1", "mid-a", now=1.0) + self.book.mark_business_done("s1") + self.assertFalse(self.book.append_follow_text("s1", "海运")) + + def test_http_and_worker_share_kv(self) -> None: + """复现测服事故:HTTP 收图、Worker 认完必须看到同一份结束标记。""" + kv = _MemoryWaveKV() + http = reset_image_wave_book_for_test(kv=kv) + worker = ImageWaveBook(kv=kv) + http.open_or_join("s1", "mid-a", now=1.0) + self.assertIsNotNone(worker.active("s1")) + worker.add_vision("s1", "SZX-CGK 货交深圳") + worker.mark_business_done("s1") + self.assertIsNone(http.active("s1")) + _w2, how = http.open_or_join("s1", "mid-b", now=2.0) + self.assertEqual(how, "new_after_done") + + +class _MemoryWaveKV: + """单测假 Redis:两份 ImageWaveBook 共用同一字典。""" + + def __init__(self) -> None: + self._d: dict[str, dict] = {} + + def get_json(self, kind: str, sender_id: str): + row = self._d.get(f"{kind}:{sender_id}") + return json.loads(json.dumps(row)) if row is not None else None + + def set_json(self, kind: str, sender_id: str, data: dict) -> None: + self._d[f"{kind}:{sender_id}"] = json.loads(json.dumps(data)) + + def delete(self, kind: str, sender_id: str) -> None: + self._d.pop(f"{kind}:{sender_id}", None) + + +if __name__ == "__main__": + unittest.main() diff --git a/inquiry-agent/tests/test_land_copy.py b/inquiry-agent/tests/test_land_copy.py index ec313fd..4dd207d 100644 --- a/inquiry-agent/tests/test_land_copy.py +++ b/inquiry-agent/tests/test_land_copy.py @@ -106,6 +106,57 @@ class LandCopyTest(unittest.TestCase): def test_no_staff_copy(self): self.assertIn("陆运产品岗", copy.land_group_no_staff()) + def test_ask_land_fields_marks_required_and_lists_optional(self): + text = copy.ask_land_fields( + facts={ + "运输类型": "跨境集拼", + "线路类别": "东南亚", + "运输分类": "急件", + "起运港": "广州", + "目的港": "河内", + "品名": "衣服", + "毛重": "50kg", + }, + missing_keys=["体积"], + land_subtype="跨境集拼", + ) + self.assertIn("始发站(必填):广州", text) + self.assertIn("待补充:", text) + self.assertIn("体积(必填)", text) + self.assertIn("通关口岸(非必填)", text) + self.assertNotIn(copy.LAND_CONFIRM_TAIL, text) + + def test_ask_land_fields_domestic_has_no_customs(self): + text = copy.ask_land_fields( + facts={ + "运输类型": "国内运输拼车", + "起运港": "广州", + "目的港": "深圳", + "品名": "衣服", + "毛重": "50kg", + }, + missing_keys=["体积"], + land_subtype="国内运输拼车", + ) + self.assertIn("体积(必填)", text) + self.assertNotIn("通关口岸", text) + + def test_ask_land_fields_filled_customs_not_pending(self): + text = copy.ask_land_fields( + facts={ + "运输类型": "中港整车", + "起运港": "广州", + "目的港": "香港", + "品名": "衣服", + "通关口岸": "皇岗", + }, + missing_keys=["起运港"], + land_subtype="中港整车", + ) + self.assertIn("通关口岸(非必填):皇岗", text) + pending = text.split("待补充:", 1)[1] + self.assertNotIn("通关口岸", pending) + def test_inquiry_card_lists_land_fields_and_wo(self): text = copy.inquiry_card( work_order_no="WO202609170006", @@ -261,6 +312,14 @@ class LandCopyTest(unittest.TestCase): self.assertIn("客户名称", text) self.assertNotIn("货好时间", text) + def test_image_copy(self): + self.assertEqual(copy.IMAGE_RECOGNIZING, "正在识别图片") + self.assertIn("没看清", copy.IMAGE_UNREADABLE) + self.assertIn("暂时没认出来", copy.IMAGE_TECH_FAIL) + merged = copy.image_fail_then_clarify(copy.IMAGE_UNREADABLE, copy.ASK_TRANSPORT) + self.assertTrue(merged.startswith(copy.IMAGE_UNREADABLE)) + self.assertIn(copy.ASK_TRANSPORT, merged) + if __name__ == "__main__": unittest.main() diff --git a/inquiry-agent/tests/test_land_text_inquiry.py b/inquiry-agent/tests/test_land_text_inquiry.py index 97c7c81..b3c3007 100644 --- a/inquiry-agent/tests/test_land_text_inquiry.py +++ b/inquiry-agent/tests/test_land_text_inquiry.py @@ -122,7 +122,8 @@ class LandTextInquiryFrontTests(unittest.TestCase): injected_mode="LAND", ) self.assertEqual(phase, "clarify") - self.assertIn("车型/数量", self._joined()) + self.assertIn("车型/数量(必填)", self._joined()) + self.assertIn("通关口岸(非必填)", self._joined()) self.assertNotIn("\n数量:", self._joined()) self.assertNotIn(copy.LAND_CONFIRM_TAIL, self._joined()) self.assertEqual(self.ledger.create_calls, 0) @@ -183,6 +184,53 @@ class LandTextInquiryFrontTests(unittest.TestCase): self.assertEqual(phase, "wait_confirm") self.assertIn("通关口岸(非必填):", self._joined()) + def test_wait_confirm_can_fill_customs_port(self) -> None: + """核对卡空着通关口岸时,再发「通关口岸:皇岗」必须写进下一张卡。""" + self.flow.on_text( + sender_id="u1", + text="跨境集拼 东南亚 急件", + reply=self._reply, + injected_facts={ + "起运港": "广州", + "目的港": "河内", + "品名": "衣服", + "体积": "100CBM", + "毛重": "50kg", + }, + injected_mode="LAND", + ) + self.replies.clear() + phase = self.flow.on_text( + sender_id="u1", + text="通关口岸:皇岗", + reply=self._reply, + ) + self.assertEqual(phase, "wait_confirm") + self.assertIn("通关口岸(非必填):皇岗", self._joined()) + self.assertEqual(self.flow.session_of("u1").facts.get("通关口岸"), "皇岗") + + def test_wait_confirm_fills_customs_from_card_label(self) -> None: + """销售按卡片原样写「通关口岸(非必填):皇岗」也要收下。""" + self.flow.on_text( + sender_id="u1", + text="中港整车 中港/中亚/中欧 急件", + reply=self._reply, + injected_facts={ + "起运港": "广州", + "目的港": "香港", + "品名": "衣服", + }, + injected_mode="LAND", + ) + self.replies.clear() + phase = self.flow.on_text( + sender_id="u1", + text="通关口岸(非必填):皇岗", + reply=self._reply, + ) + self.assertEqual(phase, "wait_confirm") + self.assertEqual(self.flow.session_of("u1").facts.get("通关口岸"), "皇岗") + def test_missing_volume_no_confirm(self) -> None: phase = self.flow.on_text( sender_id="u1", @@ -198,6 +246,8 @@ class LandTextInquiryFrontTests(unittest.TestCase): ) self.assertEqual(phase, "clarify") self.assertNotIn(copy.LAND_CONFIRM_TAIL, self._joined()) + self.assertIn("体积(必填)", self._joined()) + self.assertNotIn("通关口岸", self._joined()) class LandTextConfirmTests(unittest.TestCase): diff --git a/inquiry-agent/tests/test_mode_vision.py b/inquiry-agent/tests/test_mode_vision.py new file mode 100644 index 0000000..55ac754 --- /dev/null +++ b/inquiry-agent/tests/test_mode_vision.py @@ -0,0 +1,102 @@ +""" +千问视觉:只抄图上可见原文,不算缺项。 +""" + +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) + +from agent.llm.mode_vision import invoke_shell, invoke_vision +from agent.llm.modes import LlmMode + + +class ModeVisionTests(unittest.TestCase): + def test_shell_reads_assistant_text(self) -> None: + got = invoke_shell(payload={"assistant_text": "海运 宁波到汉堡 服装 1x40HQ"}) + self.assertTrue(got["ok"]) + self.assertIn("宁波到汉堡", got["text"]) + self.assertFalse(got["empty"]) + + def test_blank_is_empty(self) -> None: + got = invoke_shell(payload={"assistant_text": " "}) + self.assertTrue(got["ok"]) + self.assertTrue(got["empty"]) + self.assertEqual(got["text"], "") + + def test_injected_chat_builds_image_message(self) -> None: + seen: dict = {} + + class _R: + ok = True + content = "起运港上海" + error = "" + + def chat(**kwargs): + seen.update(kwargs) + return _R() + + got = invoke_vision( + payload={"image_b64": "abc", "mime": "image/jpeg"}, chat_fn=chat + ) + self.assertEqual(got["text"], "起运港上海") + mode = seen.get("mode") + mode_val = mode.value if hasattr(mode, "value") else str(mode) + self.assertEqual(mode_val, LlmMode.VISION.value) + messages = seen.get("messages") or [] + blob = str(messages) + self.assertIn("abc", blob) + self.assertIn("image", blob.lower()) + + def test_chat_fail(self) -> None: + class _R: + ok = False + content = "" + error = "http_timeout" + + got = invoke_vision( + payload={"image_b64": "x"}, chat_fn=lambda **k: _R() + ) + self.assertFalse(got["ok"]) + self.assertEqual(got["error"], "http_timeout") + + def test_sniff_png_not_jpeg(self) -> None: + from agent.llm.mode_vision import sniff_image_mime + + png = b"\x89PNG\r\n\x1a\n" + b"\x00" * 16 + jpg = b"\xff\xd8\xff\xe0" + b"\x00" * 16 + self.assertEqual(sniff_image_mime(png), "image/png") + self.assertEqual(sniff_image_mime(jpg), "image/jpeg") + self.assertEqual(sniff_image_mime(b"short"), "") + + def test_network_without_image_fails(self) -> None: + got = invoke_vision(payload={}, chat_fn=lambda **k: None, allow_network=True) + self.assertFalse(got["ok"]) + self.assertEqual(got["error"], "no_image") + + def test_plain_text_from_list_and_reasoning(self) -> None: + from agent.llm.http_client import message_plain_text + + self.assertEqual( + message_plain_text({"content": [{"type": "text", "text": "起运地深圳"}]}), + "起运地深圳", + ) + self.assertEqual( + message_plain_text({"content": "", "reasoning_content": "目的地雅加达"}), + "目的地雅加达", + ) + + def test_vision_prompt_keeps_iata(self) -> None: + from agent.llm.mode_vision import VISION_SYSTEM + + self.assertIn("SZX-CGK", VISION_SYSTEM) + self.assertIn("禁止翻译", VISION_SYSTEM) + + +if __name__ == "__main__": + unittest.main()