Oracle / tests /test_backend.py
spacedout-bits's picture
Run a local model on ZeroGPU instead of paid inference
ad4bf51 verified
Raw History Blame Contribute Delete
7.84 kB
"""Backend selection and the local GPU client.
Real model inference is not exercised here -- downloading weights in CI would
be absurd. What is pinned down is the decision logic, the prompt/JSON handling
around generation, and that every failure path degrades instead of raising into
the request handler.
"""
import pytest
from finbot import local_llm
from finbot.backend import NullLLM, make_llm, resolve
from finbot.config import Settings
from finbot.llm import LLMUnavailable
def settings(**overrides):
base = dict(
telegram_bot_token="",
telegram_webhook_secret="",
allowed_telegram_user_ids="",
hf_token="",
hf_dataset_repo="",
llm_backend="auto",
)
base.update(overrides)
return Settings(**base)
class TestResolve:
def test_auto_prefers_local_when_a_gpu_is_present(self):
assert resolve(settings(), gpu_present=True) == "local"
def test_auto_falls_back_to_remote_with_a_token(self):
assert resolve(settings(hf_token="hf_x"), gpu_present=False) == "remote"
def test_auto_is_off_with_neither(self):
assert resolve(settings(), gpu_present=False) == "off"
def test_explicit_choices_are_honoured(self):
assert resolve(settings(llm_backend="local"), gpu_present=False) == "local"
assert resolve(settings(llm_backend="remote"), gpu_present=True) == "remote"
assert resolve(settings(llm_backend="off"), gpu_present=True) == "off"
def test_unknown_value_falls_back_to_auto(self):
assert resolve(settings(llm_backend="banana"), gpu_present=True) == "local"
def test_case_and_whitespace_tolerant(self):
assert resolve(settings(llm_backend=" LOCAL "), gpu_present=False) == "local"
class TestNullLLM:
def test_not_enabled(self):
assert NullLLM().enabled is False
def test_describes_itself(self):
assert "no GPU" in NullLLM().describe()
@pytest.mark.parametrize("method", ["chat", "chat_json", "vision_json"])
def test_every_call_raises_the_documented_error(self, method):
with pytest.raises(LLMUnavailable):
getattr(NullLLM(), method)("a", "b")
def test_close_is_safe(self):
assert NullLLM().close() is None
class TestMakeLLM:
def test_nothing_configured_yields_null(self):
assert isinstance(make_llm(settings(llm_backend="off")), NullLLM)
def test_remote_without_a_token_degrades_to_null(self):
assert isinstance(make_llm(settings(llm_backend="remote")), NullLLM)
def test_remote_with_a_token(self):
from finbot.llm import LLMClient
client = make_llm(settings(llm_backend="remote", hf_token="hf_x"))
assert isinstance(client, LLMClient)
assert client.enabled
def test_local_without_torch_degrades_rather_than_crashing(self, monkeypatch):
# A Space misconfigured to "local" on a CPU box must still boot.
monkeypatch.setattr(local_llm, "_TORCH_AVAILABLE", False)
client = make_llm(settings(llm_backend="local"))
assert isinstance(client, NullLLM)
def test_local_without_torch_falls_back_to_remote_when_possible(self, monkeypatch):
from finbot.llm import LLMClient
monkeypatch.setattr(local_llm, "_TORCH_AVAILABLE", False)
client = make_llm(settings(llm_backend="local", hf_token="hf_x"))
assert isinstance(client, LLMClient)
class TestLocalClient:
def _client(self, monkeypatch, output):
monkeypatch.setattr(local_llm, "_TORCH_AVAILABLE", True)
calls = []
def fake_generate(model_id, messages, max_new_tokens):
calls.append((model_id, messages, max_new_tokens))
return output.pop(0) if isinstance(output, list) else output
monkeypatch.setattr(local_llm, "generate_text", fake_generate)
client = local_llm.LocalLLMClient(settings(llm_backend="local"))
return client, calls
def test_chat_returns_the_completion(self, monkeypatch):
client, calls = self._client(monkeypatch, "hello there")
assert client.chat([{"role": "user", "content": "hi"}]) == "hello there"
assert calls[0][0] == "Qwen/Qwen2.5-3B-Instruct"
def test_chat_json_parses(self, monkeypatch):
client, _ = self._client(monkeypatch, '{"amount": 250}')
assert client.chat_json("sys", "user") == {"amount": 250}
def test_chat_json_strips_code_fences(self, monkeypatch):
client, _ = self._client(monkeypatch, '```json\n{"amount": 250}\n```')
assert client.chat_json("sys", "user") == {"amount": 250}
def test_chat_json_pulls_json_out_of_prose(self, monkeypatch):
client, _ = self._client(
monkeypatch, 'Sure! Here you go: {"amount": 7} Hope that helps.'
)
assert client.chat_json("sys", "user") == {"amount": 7}
def test_chat_json_repairs_on_a_second_pass(self, monkeypatch):
# Small models drift out of JSON; one repair attempt is made.
client, calls = self._client(monkeypatch, ["not json at all", '{"ok": true}'])
assert client.chat_json("sys", "user") == {"ok": True}
assert len(calls) == 2
def test_repeats_the_json_instruction_in_the_user_turn(self, monkeypatch):
client, calls = self._client(monkeypatch, "{}")
client.chat_json("sys", "extract this")
assert "valid JSON only" in calls[0][1][-1]["content"]
def test_generation_failure_becomes_llm_unavailable(self, monkeypatch):
monkeypatch.setattr(local_llm, "_TORCH_AVAILABLE", True)
def boom(*_a, **_k):
raise RuntimeError("CUDA out of memory")
monkeypatch.setattr(local_llm, "generate_text", boom)
client = local_llm.LocalLLMClient(settings())
with pytest.raises(LLMUnavailable, match="local generation failed"):
client.chat([{"role": "user", "content": "x"}])
def test_disabled_client_raises_cleanly(self, monkeypatch):
monkeypatch.setattr(local_llm, "_TORCH_AVAILABLE", False)
client = local_llm.LocalLLMClient(settings())
assert client.enabled is False
with pytest.raises(LLMUnavailable):
client.chat([{"role": "user", "content": "x"}])
def test_vision_requires_a_configured_model(self, monkeypatch):
monkeypatch.setattr(local_llm, "_TORCH_AVAILABLE", True)
client = local_llm.LocalLLMClient(settings(local_vision_model=""))
with pytest.raises(LLMUnavailable, match="no local vision model"):
client.vision_json("s", "u", b"jpeg")
def test_vision_parses_json(self, monkeypatch):
monkeypatch.setattr(local_llm, "_TORCH_AVAILABLE", True)
monkeypatch.setattr(
local_llm,
"generate_vision",
lambda *a, **k: '{"amount": 480, "currency": "INR"}',
)
client = local_llm.LocalLLMClient(settings())
assert client.vision_json("s", "u", b"jpeg")["amount"] == 480
def test_describe_mentions_the_model(self, monkeypatch):
monkeypatch.setattr(local_llm, "_TORCH_AVAILABLE", True)
assert "Qwen2.5-3B" in local_llm.LocalLLMClient(settings()).describe()
class TestGpuDecorator:
def test_decorator_is_a_noop_without_the_spaces_package(self):
# Off ZeroGPU the wrapper must return the function untouched, so the
# same code path runs locally.
if local_llm._SPACES_AVAILABLE:
pytest.skip("spaces installed; no-op path not applicable")
def fn(x):
return x * 2
assert local_llm._gpu(5)(fn)(3) == 6
def test_gpu_entrypoints_exist_at_module_level(self):
# ZeroGPU scans for these at startup; if they stop being module-level
# functions the Space silently fails to boot.
assert callable(local_llm.generate_text)
assert callable(local_llm.generate_vision)