fix(mcp): decode_image_source 非图片输入 OSError 归一 ValueError,守 INVALID_PARAM 契约

评审 fix round 1:Image.open 两处(路径/base64 分支)包 try/except OSError,
合法 base64 或存在文件但非可识别图片时不再让 PIL.UnidentifiedImageError
逃逸出 upload_image 炸 MCP 会话;补 decode 层与工具层覆盖测试各一。

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
chenweiandClaude committed 2026-07-30 16:29:44 +00:00
1 parent 22402a508d
commit bb3dbcb06a
2 files changed
+21 -2

No files matched your search

+8 -2
View File
@@ -23,7 +23,10 @@ _NOTFOUND_HINT = "未发现 NRF- 前缀设备;确认设备在位且未处于
def decode_image_source(source: str) -> Image.Image:
"""source = 服务器本地路径,或 base64(可带 data URI 前缀)。"""
if os.path.exists(source):
return Image.open(source)
try:
return Image.open(source)
except OSError: # 文件存在但不是可识别图片(含截断)
raise ValueError("source 不是可识别的图片") from None
data = source
if data.startswith("data:"):
if "," not in data:
@@ -33,7 +36,10 @@ def decode_image_source(source: str) -> Image.Image:
raw = base64.b64decode(data, validate=True)
except (binascii.Error, ValueError):
raise ValueError("source 既不是存在的文件路径,也不是合法 base64") from None
return Image.open(io.BytesIO(raw))
try:
return Image.open(io.BytesIO(raw))
except OSError: # 合法 base64 但不是可识别图片(PIL.UnidentifiedImageError 属 OSError)
raise ValueError("source 不是可识别的图片") from None
def _err(tool: str, code: str, message: str, hint: str | None = None) -> dict:
+13
View File
@@ -101,6 +101,11 @@ class TestDecodeImageSource:
with pytest.raises(ValueError):
decode_image_source("!!!not-base64!!!")
def test_base64_not_image(self):
src = base64.b64encode(b"not an image").decode()
with pytest.raises(ValueError):
decode_image_source(src)
class TestTools:
@pytest.mark.asyncio
@@ -116,6 +121,14 @@ class TestTools:
assert r["data"]["bytes_bw"] == 15000
assert r["data"]["slot"] == 0
@pytest.mark.asyncio
async def test_upload_image_not_image(self):
src = base64.b64encode(b"not an image").decode()
r = await call(make_session_coro(FakeManager()), "upload_image",
{"source": src, "slot": 0})
assert r["ok"] is False
assert r["error"]["code"] == "INVALID_PARAM"
@pytest.mark.asyncio
async def test_set_mode_invalid_param(self):
r = await call(make_session_coro(FakeManager()), "set_mode",