Files
inquiry_robot/inquiry-agent/tests/test_sheet_map_mode.py

335 lines
14 KiB
Python

"""
报价模板 sheet_map:JSON 剥离、缺项检测、补跑单测(不打外网)。
"""
from __future__ import annotations
import unittest
from types import SimpleNamespace
from agent.llm.mode_sheet_map import (
assess_mapping_gaps,
build_user_message,
invoke_sheet_map,
parse_mapping_json,
strip_json_fence,
)
class TestSheetMapParse(unittest.TestCase):
def test_strip_fence(self) -> None:
raw = '```json\n{"schemaVersion":"quote-template-mapping-v1","fields":{}}\n```'
self.assertIn("schemaVersion", strip_json_fence(raw))
def test_parse_ok(self) -> None:
raw = '{"schemaVersion":"quote-template-mapping-v1","fields":{"quote_no":"S!A1"},"feeRegions":[]}'
out = parse_mapping_json(raw)
self.assertTrue(out["ok"])
self.assertEqual(out["mapping"]["fields"]["quote_no"], "S!A1")
def test_parse_invalid(self) -> None:
out = parse_mapping_json("not-json")
self.assertFalse(out["ok"])
def test_user_message_contains_inventory(self) -> None:
msg = build_user_message(
{
"transportMode": "sea",
"fileName": "海运报价单模板.xlsx",
"sheet": "Sheet1",
"cellInventory": ["Sheet1!A1 | 报价日期 | - | FFFF00"],
}
)
self.assertIn("海运", msg)
self.assertIn("Sheet1!A1", msg)
self.assertIn("自检", msg)
def test_user_message_gaps_block(self) -> None:
msg = build_user_message(
{
"transportMode": "sea",
"fileName": "x.xlsx",
"sheet": "S",
"cellInventory": ["S!A1 | x | - | -"],
},
gaps=["费用区缺失:始发港费用"],
)
self.assertIn("补跑缺口", msg)
self.assertIn("始发港费用", msg)
def test_assess_gaps_sea_fee_and_fields(self) -> None:
payload = {
"transportMode": "sea",
"cellInventory": [
"海运询价单!A18 | 始发港费用/Origin Charges | - | -",
"海运询价单!A30 | 海运费用/Ocean Freight | - | -",
"海运询价单!A35 | 目的港费用/Destination Charges | - | -",
"海运询价单!E6 | 提货地址/Place of Receipt | - | -",
"海运询价单!A11 | 货值/Cargo Value | - | -",
],
}
mapping = {
"fields": {"quote_no": "海运询价单!B3"},
"feeRegions": [],
}
gaps = assess_mapping_gaps(payload, mapping)
self.assertTrue(any("始发港费用" in g for g in gaps))
self.assertTrue(any("海运费用" in g for g in gaps))
self.assertTrue(any("目的港费用" in g for g in gaps))
self.assertTrue(any("pickup_address" in g for g in gaps))
self.assertTrue(any("cargo_value" in g for g in gaps))
def test_assess_gaps_complete_empty(self) -> None:
payload = {
"transportMode": "sea",
"cellInventory": [
"S!A18 | 始发港费用/Origin Charges | - | -",
"S!E6 | 提货地址 | - | -",
],
}
mapping = {
"fields": {"pickup_address": "S!H6"},
"feeRegions": [
{
"id": "origin_charges",
"title": "始发港费用",
"tmsSource": "polCostItems",
"stopAtText": ["x"],
"anchor": "S!A19",
}
],
}
self.assertEqual(assess_mapping_gaps(payload, mapping), [])
def test_assess_gaps_air_ignores_sea_labels(self) -> None:
"""空运不按海运词典报「提货地址/货值」缺口。"""
payload = {
"transportMode": "air",
"cellInventory": [
"A!B2 | 起运地/Origin Airport | - | -",
"A!B3 | 毛重/Gross Weight | - | -",
"A!Z99 | 提货地址/Place of Receipt | - | -",
],
}
mapping = {"fields": {"origin": "A!C2", "weight_kg": "A!C3"}, "feeRegions": []}
gaps = assess_mapping_gaps(payload, mapping)
self.assertFalse(any("pickup_address" in g for g in gaps))
self.assertEqual(gaps, [])
def test_assess_gaps_short_en_no_false_positive(self) -> None:
"""短英文(如 POL)单独出现不误报;须中文或 ≥4 字符英文标签。"""
payload = {
"transportMode": "sea",
"cellInventory": ["S!A1 | POL code table | - | -"],
}
mapping = {"fields": {}, "feeRegions": []}
gaps = assess_mapping_gaps(payload, mapping)
self.assertFalse(any("→ 应映射 pol" in g for g in gaps))
def test_assess_gaps_land_needs_route_regions(self) -> None:
payload = {
"transportMode": "land",
"cellInventory": [
"L!A1 | 项目 | - | -",
"L!B1 | 3T | - | -",
"L!C1 | 拼车 | - | -",
],
}
mapping = {"fields": {}, "feeRegions": [], "routeRegions": []}
gaps = assess_mapping_gaps(payload, mapping)
self.assertTrue(any("routeRegions" in g for g in gaps))
def test_assess_gaps_land_route_skips_matrix_fields(self) -> None:
"""陆运已有 routeRegions 时,不把矩阵列头始发地/目的地当 fields 缺口。"""
payload = {
"transportMode": "land",
"cellInventory": [
"L!A1 | 项目 | - | -",
"L!B1 | 始发地 | - | -",
"L!C1 | 目的地 | - | -",
"L!D1 | 运输时效 | - | -",
"L!E1 | 特殊说明 | - | -",
],
}
mapping = {
"fields": {},
"feeRegions": [],
"routeRegions": [{"id": "main", "title": "线路", "anchor": "L!A2"}],
}
gaps = assess_mapping_gaps(payload, mapping)
self.assertFalse(any("origin" in g for g in gaps))
self.assertFalse(any("destination" in g for g in gaps))
self.assertFalse(any("transit_time" in g for g in gaps))
self.assertTrue(any("special_remarks" in g for g in gaps))
def test_merge_prefer_existing_keeps_old_field(self) -> None:
from agent.llm.mode_sheet_map import merge_mapping_prefer_existing
old = {
"fields": {"quote_no": "S!B3", "pol": "S!B5"},
"feeRegions": [{"id": "origin_charges", "title": "旧始发港"}],
"routeRegions": [],
}
new = {
"fields": {"quote_no": "S!Z99", "pod": "S!B6"},
"feeRegions": [
{"id": "origin_charges", "title": "新始发港"},
{"id": "ocean_freight", "title": "海运费用"},
],
"routeRegions": [],
"mappingGaps": [],
}
# 无清单:保守旧优先
merged = merge_mapping_prefer_existing(old, new)
self.assertEqual(merged["fields"]["quote_no"], "S!B3")
self.assertEqual(merged["fields"]["pol"], "S!B5")
self.assertEqual(merged["fields"]["pod"], "S!B6")
ids = [r["id"] for r in merged["feeRegions"]]
self.assertEqual(ids[0], "origin_charges")
self.assertEqual(merged["feeRegions"][0]["title"], "旧始发港")
self.assertIn("ocean_freight", ids)
def test_merge_replaces_misaligned_quote_no(self) -> None:
"""旧 quote_no 落在单选 B3 → 改用新坐标 F3。"""
from agent.llm.mode_sheet_map import merge_mapping_prefer_existing
inventory = [
"S!A3 | RFQ No: | A3:A3 | -",
"S!B3 | ○ Port to Port / 其它,3工作天回复 | B3:E3 | -",
"S!F3 | | F3:G3 | -",
]
old = {"fields": {"quote_no": "S!B3"}, "feeRegions": [], "routeRegions": []}
new = {"fields": {"quote_no": "S!F3"}, "feeRegions": [], "routeRegions": []}
merged = merge_mapping_prefer_existing(old, new, cell_inventory=inventory)
self.assertEqual(merged["fields"]["quote_no"], "S!F3")
def test_merge_keeps_aligned_quote_no_when_ai_drifts(self) -> None:
"""旧已在 RFQ 旁空格、新一轮漂到远处 → 仍保留旧坐标。"""
from agent.llm.mode_sheet_map import merge_mapping_prefer_existing
inventory = [
"S!A3 | RFQ No: | - | -",
"S!B3 | | - | -",
]
old = {"fields": {"quote_no": "S!B3"}, "feeRegions": [], "routeRegions": []}
new = {"fields": {"quote_no": "S!Z99"}, "feeRegions": [], "routeRegions": []}
merged = merge_mapping_prefer_existing(old, new, cell_inventory=inventory)
self.assertEqual(merged["fields"]["quote_no"], "S!B3")
def test_b3_radio_not_aligned(self) -> None:
from agent.llm.mode_sheet_map import _field_mapping_still_aligned
inventory = [
"S!A3 | RFQ No: | - | -",
"S!B3 | ○ Port to Port | - | -",
"S!F3 | WO202609190002 | - | -",
]
self.assertFalse(_field_mapping_still_aligned("quote_no", "S!B3", inventory))
self.assertTrue(_field_mapping_still_aligned("quote_no", "S!F3", inventory))
def test_invoke_retries_when_gaps(self) -> None:
calls: list[str] = []
def fake_chat(*, mode, messages, temperature=0.0): # noqa: ARG001
calls.append(messages[-1]["content"])
if len(calls) == 1:
content = (
'{"schemaVersion":"quote-template-mapping-v1",'
'"fields":{"quote_no":"S!B3"},"feeRegions":[]}'
)
else:
content = (
'{"schemaVersion":"quote-template-mapping-v1",'
'"fields":{"quote_no":"S!B3","pickup_address":"S!H6","cargo_value":"S!B11"},'
'"feeRegions":[{"id":"origin_charges","title":"始发港费用",'
'"tmsSource":"polCostItems","anchor":"S!A19","stopAtText":["海运费用"],'
'"grow":"insert_before_stop","styleCopyRows":1}]}'
)
return SimpleNamespace(ok=True, content=content, provider="stub", error="")
payload = {
"versionId": "v1",
"transportMode": "sea",
"fileName": "sea.xlsx",
"sheet": "S",
"cellInventory": [
"S!A18 | 始发港费用/Origin Charges | - | -",
"S!E6 | 提货地址/Place of Receipt | - | -",
"S!A11 | 货值/Cargo Value | - | -",
],
}
out = invoke_sheet_map(payload, chat_fn=fake_chat, allow_network=False)
self.assertTrue(out["ok"])
self.assertEqual(out.get("rounds"), 2)
self.assertEqual(len(calls), 2)
self.assertIn("补跑缺口", calls[1])
self.assertIn("pickup_address", out["mapping"]["fields"])
self.assertEqual(len(out["mapping"]["feeRegions"]), 1)
self.assertIn("mappingGaps", out["mapping"])
self.assertIsInstance(out["mapping"]["mappingGaps"], list)
# 第二轮已补齐清单内缺口 → mappingGaps 为空且与 gaps_after 一致
self.assertEqual(out.get("gaps_after"), out["mapping"]["mappingGaps"])
self.assertEqual(out["mapping"]["mappingGaps"], [])
def test_invoke_multi_round_until_complete(self) -> None:
"""最多 3 轮:第 1 残缺 → 第 2 仍缺 → 第 3 齐。"""
calls: list[int] = []
def fake_chat(*, mode, messages, temperature=0.0): # noqa: ARG001
calls.append(1)
n = len(calls)
if n == 1:
content = (
'{"schemaVersion":"quote-template-mapping-v1",'
'"fields":{"quote_no":"S!B3"},"feeRegions":[]}'
)
elif n == 2:
content = (
'{"schemaVersion":"quote-template-mapping-v1",'
'"fields":{"quote_no":"S!B3","pickup_address":"S!H6"},'
'"feeRegions":[{"id":"origin_charges","title":"始发港费用",'
'"tmsSource":"polCostItems","anchor":"S!A19","stopAtText":["海运费用"],'
'"grow":"insert_before_stop","styleCopyRows":1}]}'
)
else:
content = (
'{"schemaVersion":"quote-template-mapping-v1",'
'"fields":{"quote_no":"S!B3","pickup_address":"S!H6","cargo_value":"S!B11"},'
'"feeRegions":['
'{"id":"origin_charges","title":"始发港费用","tmsSource":"polCostItems",'
'"anchor":"S!A19","stopAtText":["海运费用"],'
'"grow":"insert_before_stop","styleCopyRows":1},'
'{"id":"ocean_freight","title":"海运费用","tmsSource":"oceanCostItems",'
'"anchor":"S!A31","stopAtText":["目的港费用"],'
'"grow":"insert_before_stop","styleCopyRows":1},'
'{"id":"destination_charges","title":"目的港费用","tmsSource":"podCostItems",'
'"anchor":"S!A36","stopAtText":["备注"],'
'"grow":"insert_before_stop","styleCopyRows":1}'
']}'
)
return SimpleNamespace(ok=True, content=content, provider="stub", error="")
payload = {
"versionId": "v2",
"transportMode": "sea",
"fileName": "sea.xlsx",
"sheet": "S",
"cellInventory": [
"S!A18 | 始发港费用/Origin Charges | - | -",
"S!A30 | 海运费用/Ocean Freight | - | -",
"S!A35 | 目的港费用/Destination Charges | - | -",
"S!E6 | 提货地址/Place of Receipt | - | -",
"S!A11 | 货值/Cargo Value | - | -",
],
}
out = invoke_sheet_map(payload, chat_fn=fake_chat, allow_network=False, max_rounds=3)
self.assertTrue(out["ok"])
self.assertEqual(out.get("rounds"), 3)
self.assertEqual(len(calls), 3)
self.assertEqual(out["mapping"]["mappingGaps"], [])
self.assertEqual(len(out["mapping"]["feeRegions"]), 3)
if __name__ == "__main__":
unittest.main()