Spaces:
Sleeping
Sleeping
| """Deploy-POC tests: model allowlist parsing/validation, OpenAI-compatible SSE streaming, | |
| and the public rate limiter. Same conventions as the rest of the suite — torch/FlagEmbedding | |
| stubbed by conftest, no network (urllib monkeypatched), no models.""" | |
| import importlib | |
| import pytest | |
| # ---------- model allowlist (env parsed at import -> reload under monkeypatched env) ---------- | |
| def _reload_gen(monkeypatch, **env): | |
| for k in ("PROBRAG_MODEL", "PROBRAG_MODELS", "PROBRAG_DEFAULT_MODEL", "ZAI_MODEL"): | |
| monkeypatch.delenv(k, raising=False) | |
| for k, v in env.items(): | |
| monkeypatch.setenv(k, v) | |
| import src.generate as gen | |
| return importlib.reload(gen) | |
| def test_default_is_qwen_only(monkeypatch): | |
| gen = _reload_gen(monkeypatch) | |
| assert gen.ENABLED_MODELS == ["qwen"] | |
| assert gen.DEFAULT_MODEL == "qwen" | |
| assert gen.MODELS["qwen"]["model"] == "qwen2.5:3b-instruct" | |
| def test_legacy_probrag_model_overrides_qwen_id(monkeypatch): | |
| gen = _reload_gen(monkeypatch, PROBRAG_MODEL="deepseek-r1:7b") | |
| assert gen.MODELS["qwen"]["model"] == "deepseek-r1:7b" | |
| assert gen.ENABLED_MODELS == ["qwen"] # legacy override doesn't add an entry | |
| def test_enable_list_filters_unknown_and_picks_default(monkeypatch): | |
| gen = _reload_gen(monkeypatch, PROBRAG_MODELS="glm-flash,bogus,qwen") | |
| assert gen.ENABLED_MODELS == ["glm-flash", "qwen"] # bogus dropped, order preserved | |
| assert gen.DEFAULT_MODEL == "glm-flash" # first enabled | |
| def test_explicit_default_model(monkeypatch): | |
| gen = _reload_gen(monkeypatch, PROBRAG_MODELS="glm-flash,qwen", | |
| PROBRAG_DEFAULT_MODEL="qwen") | |
| assert gen.DEFAULT_MODEL == "qwen" | |
| def test_bad_default_falls_back_to_first_enabled(monkeypatch): | |
| gen = _reload_gen(monkeypatch, PROBRAG_MODELS="qwen", PROBRAG_DEFAULT_MODEL="glm-flash") | |
| assert gen.DEFAULT_MODEL == "qwen" # glm-flash not enabled -> fall back | |
| def test_zai_model_env_flows_into_registry(monkeypatch): | |
| gen = _reload_gen(monkeypatch, ZAI_MODEL="glm-4.5-flash") | |
| assert gen.MODELS["glm-flash"]["model"] == "glm-4.5-flash" | |
| def test_resolve_rejects_not_enabled(monkeypatch): | |
| gen = _reload_gen(monkeypatch, PROBRAG_MODELS="qwen") | |
| with pytest.raises(ValueError): | |
| gen._resolve("glm-flash") # a real registry name, but not enabled | |
| def test_resolve_none_is_default(monkeypatch): | |
| gen = _reload_gen(monkeypatch) | |
| assert gen._resolve(None) is gen.MODELS["qwen"] | |
| # ---------- OpenAI-compatible SSE streaming ---------- | |
| class _FakeResp: | |
| """Context-manager stand-in for urllib's response: supports .read() (non-stream) and | |
| iteration over byte lines (stream).""" | |
| def __init__(self, lines): | |
| self._lines = [l if isinstance(l, bytes) else l.encode() for l in lines] | |
| def __enter__(self): | |
| return self | |
| def __exit__(self, *a): | |
| return False | |
| def __iter__(self): | |
| return iter(self._lines) | |
| def read(self): | |
| return b"".join(self._lines) | |
| def _openai_entry(): | |
| return {"model": "glm-4-flash", "url": "https://x/v1/chat/completions", | |
| "api": "openai", "api_key_env": "ZAI_API_KEY"} | |
| def test_openai_stream_parses_sse_deltas(monkeypatch): | |
| import src.generate as gen | |
| monkeypatch.setenv("ZAI_API_KEY", "sk-test") | |
| lines = [ | |
| 'data: {"choices":[{"delta":{"content":"Hel"}}]}\n', | |
| 'data: {"choices":[{"delta":{"content":"lo"}}]}\n', | |
| 'data: {"choices":[{"delta":{}}]}\n', # keep-alive w/ no content | |
| ': comment line\n', # non-data line ignored | |
| 'data: [DONE]\n', | |
| 'data: {"choices":[{"delta":{"content":"IGNORED"}}]}\n', # after DONE -> not reached | |
| ] | |
| monkeypatch.setattr(gen.urllib.request, "urlopen", lambda req, timeout=600: _FakeResp(lines)) | |
| out = list(gen._chat_stream([{"role": "user", "content": "hi"}], _openai_entry())) | |
| assert out == ["Hel", "lo"] | |
| def test_openai_missing_key_raises(monkeypatch): | |
| import src.generate as gen | |
| monkeypatch.delenv("ZAI_API_KEY", raising=False) | |
| with pytest.raises(RuntimeError): | |
| list(gen._chat_stream([{"role": "user", "content": "hi"}], _openai_entry())) | |
| def test_openai_nonstream_parses_content(monkeypatch): | |
| import src.generate as gen | |
| monkeypatch.setenv("ZAI_API_KEY", "sk-test") | |
| body = '{"choices":[{"message":{"content":"the answer"}}]}' | |
| monkeypatch.setattr(gen.urllib.request, "urlopen", | |
| lambda req, timeout=600: _FakeResp([body])) | |
| assert gen._chat([{"role": "user", "content": "q"}], _openai_entry()) == "the answer" | |
| def test_openai_sets_bearer_header(monkeypatch): | |
| import src.generate as gen | |
| monkeypatch.setenv("ZAI_API_KEY", "sk-secret") | |
| captured = {} | |
| def fake_urlopen(req, timeout=600): | |
| captured["auth"] = req.headers.get("Authorization") | |
| return _FakeResp(['{"choices":[{"message":{"content":"x"}}]}']) | |
| monkeypatch.setattr(gen.urllib.request, "urlopen", fake_urlopen) | |
| gen._chat([{"role": "user", "content": "q"}], _openai_entry()) | |
| assert captured["auth"] == "Bearer sk-secret" | |
| # ---------- rate limiter ---------- | |
| def _fresh_limiter(monkeypatch, tmp_path, ip_limit=5, global_limit=200): | |
| monkeypatch.setenv("PROBRAG_LIMIT_DB", str(tmp_path / "rl.sqlite3")) | |
| monkeypatch.setenv("PROBRAG_IP_LIMIT", str(ip_limit)) | |
| monkeypatch.setenv("PROBRAG_GLOBAL_LIMIT", str(global_limit)) | |
| import app.limiter as lim | |
| importlib.reload(lim) | |
| return lim | |
| class _Req: | |
| def __init__(self, ip): | |
| self.headers = {"x-forwarded-for": ip} | |
| self.client = None | |
| def test_ip_limit_blocks_after_cap(monkeypatch, tmp_path): | |
| from fastapi import HTTPException | |
| lim = _fresh_limiter(monkeypatch, tmp_path, ip_limit=2) | |
| r = _Req("9.9.9.9") | |
| lim.check_and_count(r) | |
| lim.check_and_count(r) | |
| with pytest.raises(HTTPException) as e: | |
| lim.check_and_count(r) | |
| assert e.value.status_code == 429 and "clone the repo" in e.value.detail | |
| def test_new_day_resets_ip_counter(monkeypatch, tmp_path): | |
| lim = _fresh_limiter(monkeypatch, tmp_path, ip_limit=2) | |
| r = _Req("5.5.5.5") | |
| lim.check_and_count(r) | |
| lim.check_and_count(r) # at cap for today | |
| # roll the clock forward a day -> the (day, ip_hash) key changes, counter starts fresh | |
| monkeypatch.setattr(lim, "_today", lambda: "2099-01-01") | |
| lim.check_and_count(r) | |
| lim.check_and_count(r) # would 429 if the day didn't roll — no raise means it reset | |
| def test_global_cap_blocks_across_ips(monkeypatch, tmp_path): | |
| from fastapi import HTTPException | |
| lim = _fresh_limiter(monkeypatch, tmp_path, ip_limit=100, global_limit=3) | |
| for i in range(3): | |
| lim.check_and_count(_Req(f"1.0.0.{i}")) # distinct IPs, under per-IP cap | |
| with pytest.raises(HTTPException) as e: | |
| lim.check_and_count(_Req("1.0.0.99")) | |
| assert e.value.status_code == 429 and "budget exhausted" in e.value.detail | |
| def test_raw_ip_never_stored(monkeypatch, tmp_path): | |
| lim = _fresh_limiter(monkeypatch, tmp_path) | |
| lim.check_and_count(_Req("203.0.113.7")) | |
| rows = lim._db().execute("SELECT ip_hash FROM ip_counts").fetchall() | |
| assert rows and all("203.0.113.7" not in row[0] for row in rows) | |
| def test_xff_first_value_used(monkeypatch, tmp_path): | |
| lim = _fresh_limiter(monkeypatch, tmp_path) | |
| class _Multi: | |
| headers = {"x-forwarded-for": "8.8.8.8, 10.0.0.1, 10.0.0.2"} | |
| client = None | |
| assert lim._client_ip(_Multi()) == "8.8.8.8" | |