diff --git a/.github/workflows/linux-release.yml b/.github/workflows/linux-release.yml index 3f118d3..cb136ff 100644 --- a/.github/workflows/linux-release.yml +++ b/.github/workflows/linux-release.yml @@ -40,7 +40,10 @@ jobs: run: | python3 -m pip install pillow fastapi python-multipart python3 -m unittest discover -v - python3 -m compileall -q app.py backend_gemma.py backend_mlx.py backend_vlm.py image_preprocessing.py + python3 -m compileall -q app.py backend_gemma.py backend_gemma_cloud.py backend_mlx.py backend_vlm.py provider_config.py image_preprocessing.py snip.py + bash -n packaging/build-deb.sh packaging/install-user.sh packaging/localtex-launcher run.sh + shellcheck packaging/build-deb.sh packaging/install-user.sh packaging/localtex-launcher packaging/macos/build.sh run.sh + sed -n '/ -KevinTex — Snip & Get, offline +KevinTex — Snip & Get @@ -342,6 +342,19 @@ .dlg-actions { display: flex; justify-content: flex-end; gap: 8px; } .confirm-danger { background: var(--danger); border-color: var(--danger); color: #fff; } .confirm-danger:hover { background: var(--danger-hover-fg); } + .provider-group[hidden] { display: none; } + .provider-status { margin: -2px 0 12px; font-size: 12px; color: var(--muted); } + .provider-controls { display: grid; gap: 8px; } + .provider-controls button { width: 100%; text-align: left; } + .provider-key { display: grid; grid-template-columns: minmax(0, 1fr); gap: 8px; margin-top: 8px; } + .provider-key input { + width: 100%; min-width: 0; box-sizing: border-box; min-height: 40px; + padding: 9px 10px; border: 1px solid var(--border); border-radius: 8px; + background: var(--card); color: var(--text); font: inherit; + } + .provider-key input::placeholder { color: var(--muted); } + .provider-help { margin-top: 8px; font-size: 11px; color: var(--muted); } + .provider-help a { color: var(--accent); } .set-group { margin-bottom: 18px; } .set-group > label.title { display: block; font-size: 13px; font-weight: 600; margin-bottom: 8px; } .radio-row { display: flex; flex-wrap: wrap; gap: 6px 16px; } @@ -424,6 +437,17 @@ transition: width .4s ease; } .splash-sub { font-size: 12px; color: #8a94a8; min-height: 1em; } + .splash-setup { margin-top: 18px; text-align: left; } + .splash-setup[hidden] { display: none; } + .splash-setup > p { margin: 0 0 12px; color: #b9c3d6; font-size: 13px; line-height: 1.45; } + .setup-option { padding: 12px; border: 1px solid rgba(255,255,255,.15); border-radius: 10px; background: rgba(255,255,255,.05); } + .setup-option + .setup-option { margin-top: 9px; } + .setup-option button { width: 100%; padding: 9px 12px; border: 0; border-radius: 7px; cursor: pointer; background: #4b8bff; color: #fff; font-weight: 600; font-size: 13px; } + .setup-option button:hover { background: #6ea8ff; } + .setup-option small { display: block; margin-top: 7px; color: #9ca8bd; font-size: 11px; line-height: 1.4; } + .setup-option input { width: 100%; box-sizing: border-box; margin-bottom: 7px; padding: 9px 10px; border: 1px solid rgba(255,255,255,.22); border-radius: 7px; background: rgba(0,0,0,.22); color: #fff; font: inherit; } + .setup-option input::placeholder { color: #8a94a8; } + .setup-link { display: inline-block; margin-top: 10px; color: #8eb8ff; font-size: 11px; } .splash-retry { margin-top: 16px; padding: 8px 18px; border-radius: 99px; cursor: pointer; background: #4b8bff; color: #fff; border: none; font-weight: 600; font-size: 13px; @@ -435,11 +459,24 @@
- +
KevinTex
Starting up…
+
@@ -541,7 +578,7 @@

KevinTex — Snip & Get
☁︎
Drag, paste (Ctrl+V) a formula image, or click here to browse - PNG / JPG / BMP / WEBP — runs 100% on this machine, free & unlimited + PNG / JPG / BMP / WEBP — local mode runs 100% on this machine
+
@@ -675,6 +724,7 @@

Clear conversion history?

processedImage: null, sourceView: "original", recognizing: false, + backend: null, settings: Object.assign({}, defaultSettings, loadLocalJson("localtex-settings", defaultSettings)), history: loadLocalJson("localtex-history", []), historyQuery: "", @@ -905,8 +955,9 @@

Clear conversion history?

state.latex = latex; $("#latex-out").value = latex; const mode = preprocessing?.mode; + const runtime = state.backend === "gemma-cloud" ? "via AI Studio" : "on-device"; $("#elapsed").textContent = elapsed != null - ? `${elapsed}s on-device${mode ? ` · sharpener ${mode}` : ""}` + ? `${elapsed}s ${runtime}${mode ? ` · sharpener ${mode}` : ""}` : ""; renderPreview(); } @@ -1508,6 +1559,83 @@

Clear conversion history?

// Settings dialog const dlg = $("#settings-dlg"); +const providerSettings = $("#provider-settings"); +const providerStatus = $("#provider-status"); +let providerSwitching = false; + +function describeProvider(p) { + if (!p || !p.supported) return "Provider selection is unavailable for this backend."; + if (p.mode === "cloud") { + const audioModel = p.cloud_audio_model ? ` · voice ${p.cloud_audio_model}` : ""; + return p.api_key_configured + ? `Google AI Studio API · ${p.cloud_model}${audioModel} · key ${p.api_key_hint}` + : "Google AI Studio API · API key required"; + } + if (p.mode === "local") { + return p.local_model_available + ? "Local Gemma 4 weights · already downloaded" + : "Local Gemma 4 weights · downloads on first use"; + } + return "Choose local weights or Google AI Studio"; +} + +async function refreshProviderSettings() { + try { + const response = await fetch("/api/provider"); + if (!response.ok) return; + const provider = await response.json(); + providerSettings.hidden = !provider.supported; + providerStatus.textContent = describeProvider(provider); + } catch {} +} + +async function chooseProvider(mode, apiKey = "") { + if (providerSwitching) return; + providerSwitching = true; + const payload = { mode }; + if (apiKey.trim()) payload.api_key = apiKey.trim(); + const localButton = $("#settings-local"); + const cloudButton = $("#settings-cloud"); + const setupLocal = $("#setup-local"); + const setupCloud = $("#setup-cloud"); + [localButton, cloudButton, setupLocal, setupCloud].forEach((button) => { if (button) button.disabled = true; }); + try { + const response = await fetch("/api/provider", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(payload), + }); + const data = await response.json().catch(() => ({})); + if (!response.ok) throw new Error(data.error || response.statusText); + if (dlg.open) dlg.close(); + _ready = false; + splashSetup.hidden = true; + $("#splash-spinner").hidden = false; + splashRetry.hidden = true; + showSplash(mode === "cloud" ? "Starting cloud model…" : "Preparing local model…"); + await pollStatus(); + startStatusPolling(); + } catch (error) { + toast(error.message || "Could not change provider"); + } finally { + providerSwitching = false; + [localButton, cloudButton, setupLocal, setupCloud].forEach((button) => { if (button) button.disabled = false; }); + $("#setup-api-key").value = ""; + $("#settings-api-key").value = ""; + refreshProviderSettings(); + } +} + +$("#setup-local").addEventListener("click", () => chooseProvider("local")); +$("#setup-cloud").addEventListener("click", () => chooseProvider("cloud", $("#setup-api-key").value)); +$("#settings-local").addEventListener("click", () => chooseProvider("local")); +$("#settings-cloud").addEventListener("click", () => chooseProvider("cloud", $("#settings-api-key").value)); +$("#setup-api-key").addEventListener("keydown", (event) => { + if (event.key === "Enter") chooseProvider("cloud", event.currentTarget.value); +}); +$("#settings-api-key").addEventListener("keydown", (event) => { + if (event.key === "Enter") chooseProvider("cloud", event.currentTarget.value); +}); applyTheme(); matchMedia("(prefers-color-scheme: dark)").addEventListener("change", () => { if ((state.settings.theme || "system") === "system") applyTheme(); @@ -1524,6 +1652,7 @@

Clear conversion history?

dlg.querySelector(`input[name=autocopy][value="${state.settings.autocopy}"]`).checked = true; dlg.querySelector(`input[name=theme][value="${state.settings.theme}"]`).checked = true; dlg.showModal(); + refreshProviderSettings(); }); $("#dlg-close").addEventListener("click", () => dlg.close()); dlg.addEventListener("click", (e) => { if (e.target === dlg) dlg.close(); }); @@ -1544,6 +1673,8 @@

Clear conversion history?

const splashMsg = $("#splash-message"); const splashBar = $("#splash-bar"); const splashSub = $("#splash-sub"); +const splashSpinner = $("#splash-spinner"); +const splashSetup = $("#splash-setup"); const splashRetry = $("#splash-retry"); let _ready = false; let _statusTimer = null; @@ -1552,20 +1683,51 @@

Clear conversion history?

function hideSplash() { splash.classList.add("hidden"); } function fmtBackendLabel(s) { - return s === "pix2tex" ? "pix2tex" : (s === "lfm-vl" ? "LFM2.5-VL" : (s === "mlx" ? "Gemma 4 E2B · MLX" : "Gemma 4 E2B")); + return s === "pix2tex" ? "pix2tex" : (s === "lfm-vl" ? "LFM2.5-VL" : (s === "mlx" ? "Gemma 4 E2B · MLX" : (s === "gemma-cloud" ? "Gemma 4 · AI Studio" : "Gemma 4 E2B"))); } function applyStatus(s) { if (!s) return; - $("#voice-btn").hidden = !s.audio_supported; - const dev = s.device === "cuda" ? "GPU" : (s.device === "metal" ? "Apple Silicon" : "CPU"); + state.backend = s.backend || null; + const voiceButton = $("#voice-btn"); + voiceButton.hidden = !s.audio_supported; + voiceButton.title = s.backend === "gemma-cloud" + ? "Dictate a formula — the recording is sent to Google AI Studio" + : "Dictate a formula"; + const deviceLabels = { + cuda: "NVIDIA GPU", + rocm: "AMD GPU · ROCm", + vulkan: "GPU · Vulkan", + sycl: "Intel GPU · SYCL", + metal: "Apple Silicon", + cloud: "Google AI Studio", + cpu: "CPU", + }; + const dev = deviceLabels[s.device] || "CPU"; $("#device-badge").textContent = `${fmtBackendLabel(s.backend)} · ${dev}`; + document.title = s.backend === "gemma-cloud" ? "KevinTex — Snip & Get, AI Studio" : "KevinTex — Snip & Get"; + $("#drop-help").textContent = s.backend === "gemma-cloud" + ? "PNG / JPG / BMP / WEBP — images and voice recordings are sent to Google AI Studio" + : "PNG / JPG / BMP / WEBP — local mode runs 100% on this machine"; + if (s.phase === "setup" || s.setup_required) { + _ready = false; + showSplash("Choose how to run Gemma 4"); + splashSpinner.hidden = true; + splashSetup.hidden = false; + splashRetry.hidden = true; + splashBar.style.width = "0%"; + splashSub.textContent = "This choice can be changed later in Settings"; + return; + } + splashSpinner.hidden = false; + splashSetup.hidden = true; if (s.phase === "ready") { _ready = true; hideSplash(); return; } if (s.phase === "error") { splashMsg.textContent = "Model failed to load"; splashSub.textContent = s.error || ""; splashBar.style.width = "0%"; splashRetry.hidden = false; + splashSetup.hidden = !s.provider?.supported; return; } showSplash(); @@ -1573,14 +1735,18 @@

Clear conversion history?

splashBar.style.width = (s.progress || 0) + "%"; splashSub.textContent = s.phase === "downloading" ? "One-time setup — this won't repeat on future launches" - : ((s.backend === "gemma" || s.backend === "mlx") ? "Gemma 4 E2B-it · " + dev : dev); + : (s.backend === "gemma-cloud" ? "Hosted models · images and voice recordings are sent to Google" : ((s.backend === "gemma" || s.backend === "mlx") ? "Gemma 4 E2B-it · " + dev : dev)); splashRetry.hidden = true; } async function pollStatus() { try { const r = await fetch("/api/status"); - if (r.ok) { const s = await r.json(); applyStatus(s); if (s.phase === "ready") return true; } + if (r.ok) { + const s = await r.json(); + applyStatus(s); + if (s.phase === "ready" || s.phase === "error") return true; + } } catch {} return false; } @@ -1596,8 +1762,10 @@

Clear conversion history?

async function ensureReady() { if (_ready) return true; showSplash("Warming up…"); - while (!await pollStatus()) { await new Promise((r) => setTimeout(r, 1000)); } - return _ready; + while (true) { + if (await pollStatus()) return _ready; + await new Promise((r) => setTimeout(r, 1000)); + } } splashRetry.addEventListener("click", () => { diff --git a/test_backend_cleanup.py b/test_backend_cleanup.py index 250246f..f75755b 100644 --- a/test_backend_cleanup.py +++ b/test_backend_cleanup.py @@ -2,7 +2,7 @@ import unittest -from backend_vlm import _extract_markdown, _strip_meta_junk, _to_markdown_math +from backend_vlm import _extract_markdown, _to_markdown_math class BackendCleanupTests(unittest.TestCase): diff --git a/test_cuda_runtime.py b/test_cuda_runtime.py new file mode 100644 index 0000000..1a27ed7 --- /dev/null +++ b/test_cuda_runtime.py @@ -0,0 +1,45 @@ +"""CUDA runtime discovery tests.""" + +import os +import sys +import tempfile +import types +import unittest +from pathlib import Path +from unittest.mock import patch + +import backend_gemma + + +class CudaRuntimeTests(unittest.TestCase): + def test_windows_torch_lib_is_kept_on_dll_search_path(self): + with tempfile.TemporaryDirectory() as root: + torch_dir = Path(root, "torch") + torch_lib = torch_dir / "lib" + torch_lib.mkdir(parents=True) + fake_torch = types.SimpleNamespace(__file__=str(torch_dir / "__init__.py")) + dll_handle = object() + + backend_gemma._CUDA_DLL_DIR_HANDLES.clear() + with ( + patch.dict(sys.modules, {"torch": fake_torch}), + patch.object(backend_gemma.os, "name", "nt"), + patch.object( + backend_gemma.os, + "add_dll_directory", + return_value=dll_handle, + create=True, + ) as add_dll_directory, + patch.dict(os.environ, {"PATH": "existing"}, clear=True), + ): + backend_gemma._preload_cuda_libs() + + add_dll_directory.assert_called_once_with(str(torch_lib)) + self.assertEqual(backend_gemma._CUDA_DLL_DIR_HANDLES, [dll_handle]) + self.assertEqual( + os.environ["PATH"], f"{torch_lib}{os.pathsep}existing" + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test_desktop_acceleration.py b/test_desktop_acceleration.py new file mode 100644 index 0000000..4ce4f0b --- /dev/null +++ b/test_desktop_acceleration.py @@ -0,0 +1,91 @@ +"""Windows launcher acceleration-selection tests.""" + +import os +import sys +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch +from xml.etree import ElementTree + +import desktop + + +class DesktopAccelerationTests(unittest.TestCase): + def test_windows_bundle_allows_web_marked_managed_assemblies(self): + config = Path(__file__).parent / "packaging/windows/KevinTex.exe.config" + setting = ElementTree.parse(config).find("./runtime/loadFromRemoteSources") + + self.assertIsNotNone(setting) + self.assertEqual(setting.attrib.get("enabled"), "true") + + def test_windows_interop_dependencies_are_reproducibly_pinned(self): + requirements = (Path(__file__).parent / "requirements-windows.txt").read_text( + encoding="utf-8" + ) + + for dependency in ( + "pyinstaller==6.21.0", + "pywebview==6.2.1", + "pythonnet==3.0.5", + "clr-loader==0.2.7.post0", + ): + self.assertIn(dependency, requirements) + + def test_windows_smoke_check_imports_winforms_backend(self): + with patch.object(desktop.sys, "platform", "win32"), patch( + "desktop.importlib.import_module" + ) as import_module: + desktop._verify_native_window_backend() + + import_module.assert_called_once_with("webview.platforms.winforms") + + def test_non_windows_smoke_check_skips_winforms_backend(self): + with patch.object(desktop.sys, "platform", "linux"), patch( + "desktop.importlib.import_module" + ) as import_module: + desktop._verify_native_window_backend() + + import_module.assert_not_called() + + def test_cpu_is_safe_default(self): + with tempfile.TemporaryDirectory() as bundle, patch.object( + sys, "_MEIPASS", bundle, create=True + ), patch.dict(os.environ, {}, clear=True): + self.assertEqual(desktop._bundled_acceleration(), "cpu") + + def test_packaged_acceleration_marker_is_used(self): + with tempfile.TemporaryDirectory() as bundle: + Path(bundle, "acceleration.txt").write_text("vulkan\n", encoding="utf-8") + with patch.object(sys, "_MEIPASS", bundle, create=True), patch.dict( + os.environ, {}, clear=True + ): + self.assertEqual(desktop._bundled_acceleration(), "vulkan") + + def test_packaged_marker_wins_over_environment_override(self): + with tempfile.TemporaryDirectory() as bundle: + Path(bundle, "acceleration.txt").write_text("cpu\n", encoding="utf-8") + with patch.object(sys, "_MEIPASS", bundle, create=True), patch.dict( + os.environ, {"KEVINTEX_ACCELERATION": "ROCM"}, clear=True + ): + self.assertEqual(desktop._bundled_acceleration(), "cpu") + + def test_environment_selects_a_source_build_without_marker(self): + with tempfile.TemporaryDirectory() as bundle, patch.object( + sys, "_MEIPASS", bundle, create=True + ), patch.dict( + os.environ, {"KEVINTEX_ACCELERATION": "ROCM"}, clear=True + ): + self.assertEqual(desktop._bundled_acceleration(), "rocm") + + def test_legacy_cuda_marker_remains_compatible(self): + with tempfile.TemporaryDirectory() as bundle: + Path(bundle, "cuda_enabled.txt").write_text("CUDA build", encoding="utf-8") + with patch.object(sys, "_MEIPASS", bundle, create=True), patch.dict( + os.environ, {}, clear=True + ): + self.assertEqual(desktop._bundled_acceleration(), "cuda") + + +if __name__ == "__main__": + unittest.main() diff --git a/test_gemma_cloud.py b/test_gemma_cloud.py new file mode 100644 index 0000000..ac136a4 --- /dev/null +++ b/test_gemma_cloud.py @@ -0,0 +1,201 @@ +"""Tests for the hosted Gemma provider and provider selection state.""" + +import asyncio +import os +import stat +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +from PIL import Image + +import app +import backend_gemma_cloud +import provider_config + + +class ProviderConfigTests(unittest.TestCase): + def test_api_key_is_persisted_without_being_returned_by_provider_state(self): + with tempfile.TemporaryDirectory() as directory: + config = Path(directory) / "provider.json" + with patch.dict(os.environ, {"LOCALTEX_PROVIDER_CONFIG": str(config)}, clear=False): + provider_config.save_config("cloud", "AIza-test-key") + saved = provider_config.read_config() + self.assertEqual(saved["mode"], "cloud") + self.assertEqual(saved["api_key"], "AIza-test-key") + if os.name != "nt": + self.assertEqual(stat.S_IMODE(config.stat().st_mode), 0o600) + + with patch.object(app, "BACKEND", "gemma"): + state = app._provider_state() + self.assertTrue(state["api_key_configured"]) + self.assertNotIn("api_key", state) + + def test_first_run_requires_a_provider(self): + with tempfile.TemporaryDirectory() as directory: + config = Path(directory) / "provider.json" + with patch.dict( + os.environ, + {"LOCALTEX_PROVIDER_CONFIG": str(config), "GEMINI_API_KEY": ""}, + clear=False, + ), patch.object(app, "BACKEND", "gemma"): + self.assertTrue(app._setup_required()) + self.assertEqual(app.status()["phase"], "setup") + + def test_provider_switch_is_rejected_during_inference(self): + class Request: + async def json(self): + return {"mode": "local"} + + acquired = [] + try: + for _ in range(app.INFERENCE_CONCURRENCY): + self.assertTrue(app._inference_slots.acquire(blocking=False)) + acquired.append(True) + with patch.object(app, "BACKEND", "gemma"): + response = asyncio.run(app.configure_provider(Request())) + self.assertEqual(response.status_code, 409) + self.assertIn("recognition", response.body.decode()) + finally: + for _ in acquired: + app._inference_slots.release() + + def test_backend_cleanup_calls_close(self): + class Backend: + closed = False + + def close(self): + self.closed = True + + backend = Backend() + app._dispose_model(backend) + self.assertTrue(backend.closed) + + +class CloudBackendTests(unittest.TestCase): + def test_image_request_uses_gemma_api_and_keeps_image_before_prompt(self): + calls = {} + + class FakePart: + @classmethod + def from_bytes(cls, **kwargs): + calls["image"] = kwargs + return "image-part" + + class FakeThinkingConfig: + def __init__(self, **kwargs): + self.kwargs = kwargs + + class FakeGenerateContentConfig: + def __init__(self, **kwargs): + self.kwargs = kwargs + + Types = type( + "Types", + (), + { + "Part": FakePart, + "ThinkingConfig": FakeThinkingConfig, + "GenerateContentConfig": FakeGenerateContentConfig, + }, + ) + + class Models: + def generate_content(self, **kwargs): + calls["request"] = kwargs + return type("Response", (), {"text": "LATEX: $x^2$"})() + + backend = backend_gemma_cloud.GemmaCloudBackend.__new__( + backend_gemma_cloud.GemmaCloudBackend + ) + backend.types = Types + backend.client = type("Client", (), {"models": Models()})() + backend.default_thinking = False + import threading + + backend._inference_lock = threading.Lock() + + image = Image.new("RGB", (3, 3), "white") + try: + self.assertEqual(backend.recognize(image), "$x^2$") + finally: + image.close() + + self.assertEqual(calls["request"]["model"], "gemma-4-26b-a4b-it") + self.assertEqual(calls["request"]["contents"], ["image-part", backend_gemma_cloud.PROMPT]) + self.assertEqual(calls["image"]["mime_type"], "image/png") + + def test_audio_request_uses_gemini_inline_wav_input(self): + calls = {} + + class FakePart: + @classmethod + def from_bytes(cls, **kwargs): + calls["audio"] = kwargs + return "audio-part" + + class FakeThinkingConfig: + def __init__(self, **kwargs): + self.kwargs = kwargs + + class FakeGenerateContentConfig: + def __init__(self, **kwargs): + self.kwargs = kwargs + + Types = type( + "Types", + (), + { + "Part": FakePart, + "ThinkingConfig": FakeThinkingConfig, + "GenerateContentConfig": FakeGenerateContentConfig, + }, + ) + + class Models: + def generate_content(self, **kwargs): + calls["request"] = kwargs + return type("Response", (), {"text": "LATEX: $\\frac{1}{2}$"})() + + backend = backend_gemma_cloud.GemmaCloudBackend.__new__( + backend_gemma_cloud.GemmaCloudBackend + ) + backend.types = Types + backend.client = type("Client", (), {"models": Models()})() + backend.default_thinking = False + import threading + + backend._inference_lock = threading.Lock() + + with tempfile.NamedTemporaryFile(suffix=".wav") as audio: + audio.write(b"RIFF" + b"\0" * 36 + b"WAVEaudio") + audio.flush() + with patch.object( + backend_gemma_cloud, "AUDIO_MODEL_ID", "gemini-audio-test" + ): + self.assertEqual(backend.recognize_audio(audio.name), "$\\frac{1}{2}$") + + self.assertEqual(calls["request"]["model"], "gemini-audio-test") + self.assertEqual( + calls["request"]["contents"], + [backend_gemma_cloud.AUDIO_PROMPT, "audio-part"], + ) + self.assertEqual(calls["audio"]["mime_type"], "audio/wav") + self.assertTrue(calls["audio"]["data"].startswith(b"RIFF")) + + def test_oversized_audio_is_rejected_before_reading(self): + backend = backend_gemma_cloud.GemmaCloudBackend.__new__( + backend_gemma_cloud.GemmaCloudBackend + ) + backend.default_thinking = False + with tempfile.NamedTemporaryFile(suffix=".wav") as audio: + audio.write(b"12345") + audio.flush() + with patch.object(backend_gemma_cloud, "MAX_INLINE_AUDIO_BYTES", 4): + with self.assertRaisesRegex(RuntimeError, "14 MiB maximum"): + backend.recognize_audio(audio.name) + + +if __name__ == "__main__": + unittest.main() diff --git a/test_image_preprocessing.py b/test_image_preprocessing.py index f635274..e3f329d 100644 --- a/test_image_preprocessing.py +++ b/test_image_preprocessing.py @@ -2,6 +2,7 @@ import asyncio import io +import tempfile import unittest from PIL import Image, ImageDraw, ImageStat @@ -10,6 +11,13 @@ from image_preprocessing import preprocess_image, validate_options +def upload_file(data: bytes, filename: str) -> UploadFile: + file = tempfile.SpooledTemporaryFile() + file.write(data) + file.seek(0) + return UploadFile(file, filename=filename) + + class PreprocessingTests(unittest.TestCase): def formula_image(self, size=(320, 100)): image = Image.new("RGB", size, (224, 224, 224)) @@ -58,7 +66,7 @@ def test_preprocess_api_runs_without_ocr_model(self): source = self.formula_image() encoded = io.BytesIO() source.save(encoded, "PNG") - upload = UploadFile(io.BytesIO(encoded.getvalue()), filename="generated.png") + upload = upload_file(encoded.getvalue(), "generated.png") response = asyncio.run( preprocess_preview(upload, preprocess="auto", rotation=90, invert=False) ) diff --git a/test_mlx_download.py b/test_mlx_download.py new file mode 100644 index 0000000..e695824 --- /dev/null +++ b/test_mlx_download.py @@ -0,0 +1,45 @@ +"""Regression tests for resilient macOS MLX model downloads.""" + +import os +import unittest +from unittest.mock import Mock, patch + +import backend_mlx + + +class MLXDownloadTests(unittest.TestCase): + def test_xet_is_disabled_before_huggingface_import(self): + self.assertEqual(os.environ.get("HF_HUB_DISABLE_XET"), "1") + + def test_xet_401_refreshes_download_once(self): + download = Mock( + side_effect=[ + RuntimeError("401 Unauthorized from cas-bridge.xethub.hf.co"), + "/tmp/model", + ] + ) + with patch.object(backend_mlx.time, "sleep") as sleep: + backend_mlx._download_model_snapshot(download) + + self.assertEqual(download.call_count, 2) + sleep.assert_called_once_with(1.0) + for call in download.call_args_list: + self.assertEqual(call.kwargs["max_workers"], 4) + self.assertEqual(call.kwargs["etag_timeout"], 30) + + def test_persistent_xet_401_has_short_actionable_error(self): + download = Mock( + side_effect=RuntimeError( + "401 Unauthorized https://cas-bridge.xethub.hf.co/very-long-url" + ) + ) + with patch.object(backend_mlx.time, "sleep"): + with self.assertRaisesRegex(RuntimeError, "Date & Time") as raised: + backend_mlx._download_model_snapshot(download) + + self.assertEqual(download.call_count, 2) + self.assertNotIn("https://", str(raised.exception)) + + +if __name__ == "__main__": + unittest.main() diff --git a/test_performance_safety.py b/test_performance_safety.py index 0080f60..f5898b5 100644 --- a/test_performance_safety.py +++ b/test_performance_safety.py @@ -1,7 +1,8 @@ """Concurrency, backpressure, and resource-limit tests for KevinTex.""" import asyncio -import io +import os +import tempfile import threading import time import unittest @@ -17,6 +18,13 @@ from pathlib import Path +def upload_file(data: bytes, filename: str) -> UploadFile: + file = tempfile.SpooledTemporaryFile() + file.write(data) + file.seek(0) + return UploadFile(file, filename=filename) + + class ModelInitializationTests(unittest.TestCase): def test_generic_backend_load_error_reaches_status(self): previous_backend = app.BACKEND @@ -76,6 +84,18 @@ def test_mlx_backend_is_selected_directly(self): app._model = previous_model app.BACKEND = previous_backend + def test_native_acceleration_label_uses_launcher_selection(self): + previous_backend = app.BACKEND + app.BACKEND = "gemma" + try: + for acceleration in ("cpu", "cuda", "rocm", "vulkan", "sycl"): + with self.subTest(acceleration=acceleration), patch.dict( + os.environ, {"LOCALTEX_ACCELERATION": acceleration} + ): + self.assertEqual(app._device_label(), acceleration) + finally: + app.BACKEND = previous_backend + def test_optiq_model_skips_audio_repair(self): with patch.object( backend_mlx, @@ -104,11 +124,13 @@ def create_chat_completion(self, **kwargs): backend.default_thinking = False backend._inference_lock = threading.Lock() backend.llm = FakeLlama() - import tempfile - with tempfile.NamedTemporaryFile(suffix=".wav") as audio: + with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as audio: audio.write(b"RIFF" + b"\0" * 40) - audio.flush() - self.assertEqual(backend.recognize_audio(audio.name), "$x^2$") + audio_path = audio.name + try: + self.assertEqual(backend.recognize_audio(audio_path), "$x^2$") + finally: + os.unlink(audio_path) media = backend.llm.kwargs["messages"][0]["content"][0] self.assertEqual(media["type"], "image_url") self.assertTrue(media["image_url"]["url"].startswith("data:audio/wav;base64,")) @@ -205,7 +227,7 @@ async def should_not_run(_request): def test_convert_rejects_oversized_upload_before_inference(self): previous_limit = app.MAX_UPLOAD_BYTES app.MAX_UPLOAD_BYTES = 8 - upload = UploadFile(io.BytesIO(b"x" * 9), filename="large.png") + upload = upload_file(b"x" * 9, "large.png") try: response = asyncio.run( app.convert( @@ -224,7 +246,7 @@ def test_convert_rejects_oversized_upload_before_inference(self): def test_reader_accepts_exact_limit(self): previous_limit = app.MAX_UPLOAD_BYTES app.MAX_UPLOAD_BYTES = 8 - upload = UploadFile(io.BytesIO(b"x" * 8), filename="exact.bin") + upload = upload_file(b"x" * 8, "exact.bin") try: raw = asyncio.run(app._read_upload(upload)) finally: @@ -252,7 +274,7 @@ def load(self): def test_voice_is_rejected_on_backend_without_audio(self): previous_backend = app.BACKEND app.BACKEND = "pix2tex" - upload = UploadFile(io.BytesIO(b"not audio"), filename="voice.wav") + upload = upload_file(b"not audio", "voice.wav") try: response = asyncio.run(app.voice(upload, thinking=False)) finally: @@ -262,13 +284,28 @@ def test_voice_is_rejected_on_backend_without_audio(self): def test_voice_validates_wav_before_inference(self): previous_backend = app.BACKEND app.BACKEND = "mlx" - upload = UploadFile(io.BytesIO(b"x" * 64), filename="voice.wav") + upload = upload_file(b"x" * 64, "voice.wav") try: response = asyncio.run(app.voice(upload, thinking=False)) finally: app.BACKEND = previous_backend self.assertEqual(response.status_code, 400) + def test_cloud_provider_accepts_voice_uploads(self): + wav = b"RIFF" + (40).to_bytes(4, "little") + b"WAVE" + b"\0" * 36 + upload = upload_file(wav, "voice.wav") + with ( + patch.object(app, "BACKEND", "gemma-cloud"), + patch.object(app, "model_ready", return_value=True), + patch.object(app, "_run_audio", return_value="$x^2$") as run_audio, + ): + response = asyncio.run(app.voice(upload, thinking=False)) + health = app.health() + + self.assertEqual(response["latex"], "$x^2$") + self.assertTrue(health["audio_supported"]) + run_audio.assert_called_once() + if __name__ == "__main__": unittest.main() diff --git a/test_snip.py b/test_snip.py new file mode 100644 index 0000000..f617afb --- /dev/null +++ b/test_snip.py @@ -0,0 +1,60 @@ +"""Screen-coordinate tests for the cross-platform snip overlay.""" + +import unittest +from unittest.mock import patch + +from PIL import Image +from snip import _capture_crop_box, _capture_screen, _display_size + + +class SnipCoordinateTests(unittest.TestCase): + def test_retina_capture_uses_logical_macos_screen_size(self): + self.assertEqual( + _display_size((3024, 1964), (1512, 982), platform="darwin"), + (1512, 982), + ) + + def test_retina_selection_maps_back_to_capture_pixels(self): + self.assertEqual( + _capture_crop_box( + (100, 50), + (300, 150), + display_size=(1512, 982), + capture_size=(3024, 1964), + ), + (200, 100, 600, 300), + ) + + def test_crop_coordinates_are_ordered_and_clamped(self): + self.assertEqual( + _capture_crop_box( + (900, 700), + (-10, -20), + display_size=(800, 600), + capture_size=(1600, 1200), + ), + (0, 0, 1600, 1200), + ) + + def test_non_macos_keeps_capture_dimensions(self): + self.assertEqual( + _display_size((1920, 1080), (1920, 1080), platform="linux"), + (1920, 1080), + ) + + def test_linux_capture_falls_back_when_pillow_cannot_access_desktop(self): + fallback = Image.new("RGB", (64, 32), "black") + with patch("snip.sys.platform", "linux"): + with patch( + "snip.ImageGrab.grab", side_effect=OSError("desktop unavailable") + ): + with patch("snip._capture_with_command", return_value=fallback): + captured = _capture_screen() + try: + self.assertEqual(captured.size, (64, 32)) + finally: + captured.close() + + +if __name__ == "__main__": + unittest.main()