atlas-iac/testing/tests/test_hermes_batch_client.py

150 lines
5.7 KiB
Python

"""Check that WSL batch transport and cache cannot silently change providers."""
import copy
import json
from pathlib import Path
import socket
import pytest
from scripts.ops import hermes_batch_client as client
@pytest.fixture
def envelope():
return {"campaign_id": "SYNTHETIC", "suite_id": "S1", "source_sha256": "a" * 64,
"prompt_version": "profile-v1", "request": {
"model": "qwen3.5:9b", "prompt": "Synthetic case only", "stream": False,
"think": False, "format": "json", "options": {
"num_ctx": 16384, "num_predict": 1024, "temperature": 0,
"top_p": 1, "top_k": 20, "seed": 42}}}
@pytest.fixture
def transport(tmp_path, monkeypatch):
token = tmp_path / "token"
token.write_text("b" * 64)
token.chmod(0o600)
instance = client.BatchClient(token, tmp_path / "cache")
calls = []
def respond(suffix, payload=None):
calls.append(suffix)
if suffix == "/models":
return {"runtime": client.RUNTIME, "fallback": None,
"protocol_version": client.PROTOCOL_VERSION, "backend_api": client.BACKEND_API, "models": [
{"model": n, "digest": d} for n, d in client.MODELS.items()]}
return {"model": payload["model"], "done": True, "done_reason": "stop", "response": '{}',
"batch_provenance": {"model_digest": client.MODELS[payload["model"]],
"runtime": client.RUNTIME, "think": payload["think"],
"protocol_version": client.PROTOCOL_VERSION, "backend_api": client.BACKEND_API,
"options": {**payload["options"], "num_thread": 16, "num_gpu": 0}}}
monkeypatch.setattr(instance, "_request", respond)
return instance, calls
def test_cache_resume_reuses_unchanged_results_without_resending_source(transport, envelope):
instance, calls = transport
first = instance.run(envelope)
second = instance.run(envelope)
assert first["cache_hit"] is False and second["cache_hit"] is True
assert calls == ["/models", "/generate"]
path = Path(first["path"])
assert path.stat().st_mode & 0o777 == 0o600
assert path.parent.stat().st_mode & 0o777 == 0o700
assert "Synthetic case only" not in path.read_text()
assert not list(path.parent.glob(".pending-*"))
@pytest.mark.parametrize("dimension", ["source", "prompt", "model", "schema", "options"])
def test_cache_invalidates_on_every_reproducibility_dimension(transport, envelope, dimension):
instance, calls = transport
first = instance.run(envelope)
changed = copy.deepcopy(envelope)
if dimension == "source":
changed["source_sha256"] = "c" * 64
elif dimension == "prompt":
changed["prompt_version"] = "profile-v2"
elif dimension == "model":
changed["request"]["model"] = "qwen3.6:27b"
elif dimension == "schema":
changed["request"]["format"] = {"type": "object"}
else:
changed["request"]["options"]["seed"] = 43
second = instance.run(changed)
assert first["path"] != second["path"]
assert calls == ["/models", "/generate", "/models", "/generate"]
def test_unknown_scope_or_substituted_model_never_reaches_transport(transport, envelope):
instance, calls = transport
envelope["campaign_id"] = "FA02"
with pytest.raises(ValueError):
instance.run(envelope)
assert not calls
envelope["campaign_id"] = "FA01"
envelope["request"]["model"] = "cloud"
with pytest.raises(ValueError):
instance.run(envelope)
assert not calls
def test_identity_drift_stops_before_generation_and_failed_calls_are_not_cached(transport, envelope, monkeypatch):
instance, calls = transport
def changed(suffix, payload=None):
calls.append(suffix)
return {"runtime": "new", "models": [], "fallback": None}
monkeypatch.setattr(instance, "_request", changed)
with pytest.raises(ValueError):
instance.run(envelope)
assert calls == ["/models"]
assert list(instance.cache_dir.glob("*.json")) == []
attempt = json.loads((instance.cache_dir / "attempts.jsonl").read_text())
assert attempt["outcome"] == "failed"
assert attempt["wall_seconds"] >= 0
def test_client_uses_fixed_lan_address_and_worker_sni(monkeypatch):
seen = []
sock = object()
monkeypatch.setattr(socket, "create_connection", lambda address, timeout: (seen.append(address), sock)[1])
connection = client.LanConnection(client.HOST)
class Context:
def wrap_socket(self, raw, server_hostname):
seen.append(server_hostname)
assert raw is sock
return sock
connection._context = Context()
connection.connect()
assert seen == [(client.ADDRESS, 443), client.HOST]
with pytest.raises(ValueError):
client.LanConnection("external.invalid").connect()
def test_redirects_and_shared_token_files_are_rejected(tmp_path):
assert client.NoRedirect().redirect_request(None, None, 307, "", {}, "https://external.invalid") is None
token = tmp_path / "token"
token.write_text("b" * 64)
token.chmod(0o644)
with pytest.raises(ValueError):
client.BatchClient(token, tmp_path / "cache")
def test_application_validation_failure_keeps_diagnostics_but_is_not_reused(transport, envelope):
instance, calls = transport
def reject(content):
raise ValueError("missing case IDs")
for _ in range(2):
with pytest.raises(ValueError, match="application validation failed"):
instance.run(envelope, validator=reject)
assert calls.count("/generate") == 2
assert len(list(instance.cache_dir.glob("*.rejected.json"))) == 1
assert len(list(instance.cache_dir.glob("*.json"))) == 1