114 lines
4.2 KiB
Python
114 lines
4.2 KiB
Python
"""Tests for multimodal (vision) support and ad-hoc context (#81)."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import base64
|
|
|
|
from backend.ai_chat import _content_to_gemini_parts, _gemini_contents
|
|
from backend.bookslm import (
|
|
collect_adhoc_context,
|
|
is_image_path,
|
|
load_vault_image_data_url,
|
|
merge_contexts,
|
|
)
|
|
|
|
|
|
class TestGeminiParts:
|
|
def test_plain_string(self):
|
|
assert _content_to_gemini_parts("hello") == [{"text": "hello"}]
|
|
|
|
def test_none_becomes_empty_text(self):
|
|
assert _content_to_gemini_parts(None) == [{"text": ""}]
|
|
|
|
def test_text_and_image_data_url(self):
|
|
raw = base64.b64encode(b"\x89PNG").decode()
|
|
parts = _content_to_gemini_parts([
|
|
{"type": "text", "text": "décris"},
|
|
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{raw}"}},
|
|
])
|
|
assert parts[0] == {"text": "décris"}
|
|
assert parts[1] == {"inlineData": {"mimeType": "image/png", "data": raw}}
|
|
|
|
def test_remote_image_url_becomes_file_data(self):
|
|
parts = _content_to_gemini_parts([
|
|
{"type": "image_url", "image_url": {"url": "https://x.test/a.png"}},
|
|
])
|
|
assert parts == [{"fileData": {"fileUri": "https://x.test/a.png"}}]
|
|
|
|
def test_gemini_contents_handles_multimodal_user_message(self):
|
|
raw = base64.b64encode(b"img").decode()
|
|
system, contents = _gemini_contents([
|
|
{"role": "system", "content": "sys"},
|
|
{"role": "user", "content": [
|
|
{"type": "text", "text": "hi"},
|
|
{"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{raw}"}},
|
|
]},
|
|
])
|
|
assert system == "sys"
|
|
assert contents[0]["role"] == "user"
|
|
assert contents[0]["parts"][1]["inlineData"]["mimeType"] == "image/jpeg"
|
|
|
|
|
|
class TestImageHelpers:
|
|
def test_is_image_path(self):
|
|
assert is_image_path("a/b.png")
|
|
assert is_image_path("photo.JPEG")
|
|
assert not is_image_path("note.md")
|
|
|
|
def test_load_vault_image_data_url(self, tmp_path):
|
|
vault = tmp_path / "vault"
|
|
vault.mkdir()
|
|
(vault / "pic.png").write_bytes(b"PNGDATA")
|
|
url = load_vault_image_data_url(vault, "pic.png")
|
|
assert url is not None
|
|
assert url.startswith("data:image/png;base64,")
|
|
assert base64.b64decode(url.split(",", 1)[1]) == b"PNGDATA"
|
|
|
|
def test_load_rejects_path_traversal(self, tmp_path):
|
|
vault = tmp_path / "vault"
|
|
vault.mkdir()
|
|
outside = tmp_path / "outside.png"
|
|
outside.write_bytes(b"x")
|
|
assert load_vault_image_data_url(vault, "../outside.png") is None
|
|
|
|
def test_load_rejects_non_image(self, tmp_path):
|
|
vault = tmp_path / "vault"
|
|
vault.mkdir()
|
|
(vault / "note.md").write_text("hi", encoding="utf-8")
|
|
assert load_vault_image_data_url(vault, "note.md") is None
|
|
|
|
|
|
class TestAdhocContext:
|
|
def test_files_and_directories(self, tmp_path):
|
|
vault = tmp_path / "vault"
|
|
(vault / "sub").mkdir(parents=True)
|
|
(vault / "a.md").write_text("# A", encoding="utf-8")
|
|
(vault / "sub" / "b.md").write_text("# B", encoding="utf-8")
|
|
|
|
ctx = collect_adhoc_context(vault, files=["a.md"], directories=["sub"])
|
|
paths = {f["path"] for f in ctx["files"]}
|
|
assert paths == {"a.md", "sub/b.md"}
|
|
|
|
def test_deduplicates(self, tmp_path):
|
|
vault = tmp_path / "vault"
|
|
vault.mkdir()
|
|
(vault / "a.md").write_text("# A", encoding="utf-8")
|
|
ctx = collect_adhoc_context(vault, files=["a.md"], directories=["."])
|
|
assert ctx["file_count"] == 1
|
|
|
|
def test_merge_contexts(self):
|
|
base = {
|
|
"files": [{"path": "a.md", "content": "aaa"}],
|
|
"file_count": 1, "total_chars": 3, "directory_tree": "tree", "scope": "directory",
|
|
}
|
|
extra = {
|
|
"files": [
|
|
{"path": "a.md", "content": "aaa"},
|
|
{"path": "b.md", "content": "bb"},
|
|
],
|
|
"file_count": 2, "total_chars": 5, "directory_tree": "", "scope": "directory",
|
|
}
|
|
merged = merge_contexts(base, extra)
|
|
assert [f["path"] for f in merged["files"]] == ["a.md", "b.md"]
|
|
assert merged["total_chars"] == 5
|