- src/doc/ 拆分为 src/core/extraction/, matching/, validation/(核心业务逻辑) - src/bot/ 重命名为 src/infra/browser/(浏览器自动化基础设施) - fill_consumable_doc.py → src/infra/documents/consumable.py - 新增 Agent 调度模块:coordinator.py, events.py, session.py,重构 orchestrator.py - 更新 AGENTS.md、README.md 及所有子目录 README
203 lines
6.7 KiB
Python
203 lines
6.7 KiB
Python
"""LLM 信息提取模块单元测试
|
||
|
||
覆盖范围:
|
||
- `parse_json_response`:纯 JSON、Markdown 包裹、带前缀、解析失败
|
||
- _image_to_base64:图片转 base64
|
||
- extract_document:成功提取、LLM 失败
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import base64
|
||
import json
|
||
from pathlib import Path
|
||
|
||
import pytest
|
||
|
||
from src.core.extraction import (
|
||
_image_to_base64,
|
||
extract_document,
|
||
parse_json_response,
|
||
)
|
||
|
||
# 字段键名(与源码中的字符串字面量保持一致)
|
||
K_INVOICE_NUMBER = "invoice_number"
|
||
K_TOTAL_AMOUNT = "total_amount"
|
||
K_CARD_DATE = "card_date"
|
||
K_CARD_NO = "card_no"
|
||
K_CARD_AMOUNT = "card_amount"
|
||
|
||
# ------------------------------------------------------------------
|
||
# parse_json_response
|
||
# ------------------------------------------------------------------
|
||
|
||
|
||
class TestParseJsonResponse:
|
||
"""JSON 响应解析"""
|
||
|
||
def test_pure_json(self):
|
||
raw = json.dumps({K_INVOICE_NUMBER: "123456", K_TOTAL_AMOUNT: "100.00"})
|
||
result = parse_json_response(raw)
|
||
assert result[K_INVOICE_NUMBER] == "123456"
|
||
assert result[K_TOTAL_AMOUNT] == "100.00"
|
||
|
||
def test_markdown_json_block(self):
|
||
raw = f'```json\n{{"{K_INVOICE_NUMBER}": "789"}}\n```'
|
||
result = parse_json_response(raw)
|
||
assert result[K_INVOICE_NUMBER] == "789"
|
||
|
||
def test_markdown_block_without_lang(self):
|
||
raw = '```\n{"key": "value"}\n```'
|
||
result = parse_json_response(raw)
|
||
assert result["key"] == "value"
|
||
|
||
def test_json_prefix(self):
|
||
raw = f'json\n{{"{K_INVOICE_NUMBER}": "001"}}'
|
||
result = parse_json_response(raw)
|
||
assert result[K_INVOICE_NUMBER] == "001"
|
||
|
||
def test_json_prefix_with_whitespace(self):
|
||
raw = ' json \n{"a": 1}'
|
||
result = parse_json_response(raw)
|
||
assert result["a"] == 1
|
||
|
||
def test_nested_json(self):
|
||
raw = json.dumps({"outer": {"inner": [1, 2, 3]}})
|
||
result = parse_json_response(raw)
|
||
assert result["outer"]["inner"] == [1, 2, 3]
|
||
|
||
def test_whitespace_around_json(self):
|
||
raw = ' \n {"x": 42} \n '
|
||
result = parse_json_response(raw)
|
||
assert result["x"] == 42
|
||
|
||
def test_invalid_json_raises(self):
|
||
with pytest.raises(ValueError):
|
||
parse_json_response("not json at all")
|
||
|
||
def test_empty_string_raises(self):
|
||
with pytest.raises(ValueError):
|
||
parse_json_response("")
|
||
|
||
|
||
# ------------------------------------------------------------------
|
||
# _image_to_base64
|
||
# ------------------------------------------------------------------
|
||
|
||
|
||
class TestImageToBase64:
|
||
"""图片转 base64"""
|
||
|
||
def test_png_to_base64(self, tmp_path: Path):
|
||
content = b"\x89PNG\r\n\x1a\nfake_png_data"
|
||
img_path = tmp_path / "test.png"
|
||
img_path.write_bytes(content)
|
||
|
||
result = _image_to_base64(img_path)
|
||
assert isinstance(result, str)
|
||
assert base64.b64decode(result) == content
|
||
|
||
def test_jpg_to_base64(self, tmp_path: Path):
|
||
content = b"\xff\xd8\xff\xe0fake_jpg_data"
|
||
img_path = tmp_path / "test.jpg"
|
||
img_path.write_bytes(content)
|
||
|
||
result = _image_to_base64(img_path)
|
||
assert base64.b64decode(result) == content
|
||
|
||
def test_file_not_found_raises(self, tmp_path: Path):
|
||
img_path = tmp_path / "nonexistent.png"
|
||
with pytest.raises(FileNotFoundError):
|
||
_image_to_base64(img_path)
|
||
|
||
def test_returns_utf8_string(self, tmp_path: Path):
|
||
content = b"test_image_content"
|
||
img_path = tmp_path / "test.png"
|
||
img_path.write_bytes(content)
|
||
|
||
result = _image_to_base64(img_path)
|
||
assert type(result) is str
|
||
|
||
|
||
# ------------------------------------------------------------------
|
||
# extract_document (mock LLM)
|
||
# ------------------------------------------------------------------
|
||
|
||
|
||
class TestExtractDocument:
|
||
"""图片提取:通过 mock _llm_query_multimodal 避免真实 LLM 调用"""
|
||
|
||
def _mock_multimodal(self, monkeypatch, response_text: str):
|
||
def fake_query(system_prompt, text, image_b64, max_tokens=4096):
|
||
return response_text
|
||
|
||
monkeypatch.setattr("src.core.extraction.llm_extractor._llm_query_multimodal", fake_query)
|
||
|
||
def test_success(self, tmp_path: Path, monkeypatch):
|
||
img_path = tmp_path / "card.png"
|
||
img_path.write_bytes(b"fake_image")
|
||
|
||
mock_result = {
|
||
K_CARD_DATE: "2026-01-10",
|
||
K_CARD_NO: "6228480000000000",
|
||
K_CARD_AMOUNT: "500.00",
|
||
}
|
||
self._mock_multimodal(monkeypatch, json.dumps(mock_result))
|
||
|
||
result = extract_document(img_path)
|
||
assert result[K_CARD_DATE] == "2026-01-10"
|
||
assert result[K_CARD_NO] == "6228480000000000"
|
||
assert result[K_CARD_AMOUNT] == "500.00"
|
||
|
||
def test_llm_failure_propagates(self, tmp_path: Path, monkeypatch):
|
||
img_path = tmp_path / "card.png"
|
||
img_path.write_bytes(b"fake_image")
|
||
|
||
def fake_query(system_prompt, text, image_b64, max_tokens=4096):
|
||
raise RuntimeError("模型不可用")
|
||
|
||
monkeypatch.setattr("src.core.extraction.llm_extractor._llm_query_multimodal", fake_query)
|
||
with pytest.raises(RuntimeError, match="模型不可用"):
|
||
extract_document(img_path)
|
||
|
||
def test_file_not_found_propagates(self, tmp_path: Path, monkeypatch):
|
||
img_path = tmp_path / "missing.png"
|
||
self._mock_multimodal(monkeypatch, "{}")
|
||
with pytest.raises(FileNotFoundError):
|
||
extract_document(img_path)
|
||
|
||
def test_markdown_wrapped_json(self, tmp_path: Path, monkeypatch):
|
||
img_path = tmp_path / "card.png"
|
||
img_path.write_bytes(b"fake_image")
|
||
|
||
mock_result = {
|
||
K_CARD_DATE: "2026-02-01",
|
||
K_CARD_NO: "6228481111111111",
|
||
K_CARD_AMOUNT: "300.00",
|
||
}
|
||
wrapped = f"```json\n{json.dumps(mock_result)}\n```"
|
||
self._mock_multimodal(monkeypatch, wrapped)
|
||
|
||
result = extract_document(img_path)
|
||
assert result[K_CARD_DATE] == "2026-02-01"
|
||
|
||
def test_image_encoded_as_base64(self, tmp_path: Path, monkeypatch):
|
||
img_path = tmp_path / "card.png"
|
||
expected_content = b"test_image_data"
|
||
img_path.write_bytes(expected_content)
|
||
|
||
received_b64 = None
|
||
|
||
def capture_b64(system_prompt, text, image_b64s, max_tokens=4096):
|
||
nonlocal received_b64
|
||
received_b64 = image_b64s
|
||
return json.dumps({K_CARD_DATE: "2026-01-01", K_CARD_NO: "0000", K_CARD_AMOUNT: "100"})
|
||
|
||
monkeypatch.setattr("src.core.extraction.llm_extractor._llm_query_multimodal", capture_b64)
|
||
extract_document(img_path)
|
||
|
||
assert received_b64 is not None
|
||
assert isinstance(received_b64, list)
|
||
assert len(received_b64) >= 1
|
||
assert base64.b64decode(received_b64[0]) == expected_content
|