Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8219dce84b |
@@ -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(
|
||||
|
||||
@@ -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 "",
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,不读附件、不调模型。
|
||||
|
||||
@@ -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]):
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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。"
|
||||
),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]:
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -41,6 +41,8 @@ class LandTextInquiryFlow(SeaTextInquiryFlow):
|
||||
前半段与海运不同:核对通过前不建单。后半段套海运建群,匹配改走线路类别。
|
||||
"""
|
||||
|
||||
bookmark_kind = "land"
|
||||
|
||||
def on_text(
|
||||
self,
|
||||
*,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -189,6 +189,8 @@ _ALIASES = {
|
||||
"客户名称": "客户名称",
|
||||
"货值": "货值",
|
||||
"是否为危险品": "是否为危险品",
|
||||
"通关口岸": "通关口岸",
|
||||
"通关口岸(非必填)": "通关口岸",
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -47,6 +47,35 @@ _PACKS = {
|
||||
}
|
||||
|
||||
_IATA = re.compile(r"^[A-Za-z]{3}$")
|
||||
_IATA_PAIR = re.compile(
|
||||
r"(?<![A-Za-z])([A-Za-z]{3})\s*[-–—/]\s*([A-Za-z]{3})(?![A-Za-z])"
|
||||
)
|
||||
# 原话里的 AAA-BBB 不一定是航线,货币/单位不能当成起运目的。
|
||||
_IATA_PAIR_BLOCK = frozenset(
|
||||
{
|
||||
"USD",
|
||||
"CNY",
|
||||
"EUR",
|
||||
"GBP",
|
||||
"HKD",
|
||||
"JPY",
|
||||
"CBM",
|
||||
"KGS",
|
||||
"PCS",
|
||||
"PKG",
|
||||
"PDF",
|
||||
"JPG",
|
||||
"PNG",
|
||||
"AND",
|
||||
"THE",
|
||||
"FOR",
|
||||
"NOT",
|
||||
"MAX",
|
||||
"MIN",
|
||||
"GMT",
|
||||
"UTC",
|
||||
}
|
||||
)
|
||||
_NUMBER = re.compile(r"(\d+(?:\.\d+)?)")
|
||||
|
||||
|
||||
@@ -67,12 +96,24 @@ def resolve_airport_code(raw: str) -> 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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user