Compare commits

..
30 changed files with 2153 additions and 45 deletions
@@ -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(
+5
View File
@@ -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 "",
+8 -2
View File
@@ -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
+22 -2
View File
@@ -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(
+1
View File
@@ -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:
+2 -2
View File
@@ -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)
+16 -1
View File
@@ -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,不读附件、不调模型。
+38 -4
View File
@@ -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]):
+41 -2
View File
@@ -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。"
),
+140 -4
View File
@@ -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
+105 -13
View File
@@ -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]:
"""
+386
View File
@@ -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
+54 -10
View File
@@ -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)
+5
View File
@@ -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"
+19 -1
View File
@@ -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 = {
"客户名称": "客户名称",
"货值": "货值",
"是否为危险品": "是否为危险品",
"通关口岸": "通关口岸",
"通关口岸(非必填)": "通关口岸",
}
+43 -2
View File
@@ -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:
+2
View File
@@ -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",
+64 -1
View File
@@ -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(
+272
View File
@@ -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()
+103
View File
@@ -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()
+59
View File
@@ -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()
+51 -1
View File
@@ -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):
+102
View File
@@ -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()