103 lines
3.3 KiB
Python
103 lines
3.3 KiB
Python
"""
|
|
千问视觉:只抄图上可见原文,不算缺项。
|
|
"""
|
|
|
|
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()
|