Appearance
@@ -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()