atlas-iac/testing/tests/test_hermes_model_gate_lan.py

140 lines
5.5 KiB
Python

"""Exercise the LAN gateway's authentication and local-only HTTP boundary."""
import http.client
import io
import json
from pathlib import Path
import threading
from urllib.error import URLError
import pytest
import yaml
from services.hermes.scripts import batch_api
import sys
sys.modules.setdefault("batch_api", batch_api)
@pytest.fixture
def gateway(tmp_path):
"""Run the shipped handler with an isolated token and intercepted upstream."""
path = Path(__file__).parents[2] / "services/hermes/model-gate-configmap.yaml"
namespace = {"__name__": "test_gateway"}
exec(compile(yaml.safe_load(path.read_text())["data"]["model_gate.py"], str(path), "exec"), namespace)
token = tmp_path / "token"
token.write_text("a" * 64)
namespace["LAN_API_TOKEN_FILE"] = token
calls = []
def upstream(request, timeout):
calls.append(request)
response = io.BytesIO(b'{"response":"LAN_OK","done":true}')
response.status = 200
return response
namespace["urlopen"] = upstream
server = namespace["ThreadingHTTPServer"](("127.0.0.1", 0), namespace["LanHandler"])
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
def request(method="GET", path="/healthz", body=None, **headers):
defaults = {"Authorization": "Bearer " + "a" * 64, "X-Forwarded-For": "192.168.22.8"}
defaults.update(headers)
connection = http.client.HTTPConnection(*server.server_address, timeout=3)
connection.request(method, path, body=body, headers=defaults)
response = connection.getresponse()
result = response.status, json.loads(response.read())
connection.close()
return result
yield namespace, calls, request
server.shutdown()
server.server_close()
thread.join(timeout=2)
@pytest.mark.parametrize("address", ["", "203.0.113.1", "192.168.22.evil", "192.168.22.8, 203.0.113.1"])
def test_non_lan_and_spoofed_addresses_are_rejected(gateway, address):
_, calls, request = gateway
assert request(**{"X-Forwarded-For": address})[0] == 403
assert not calls
@pytest.mark.parametrize("token", ["", "Bearer wrong", "Basic abc", "Bearer é"])
def test_token_is_required(gateway, token):
_, calls, request = gateway
assert request(Authorization=token)[0] == 401
assert not calls
def test_health_requires_a_readable_credential(gateway):
namespace, calls, request = gateway
assert request()[0] == 200
namespace["LAN_API_TOKEN_FILE"].unlink()
assert request()[0] == 503
assert not calls
def test_generation_uses_only_local_upstream_without_forwarding_token(gateway):
namespace, calls, request = gateway
body = json.dumps({"model": namespace["LAN_MODEL"], "prompt": "Say LAN_OK", "stream": False})
assert request("POST", "/api/generate", body) == (200, {"response": "LAN_OK", "done": True})
assert calls[0].full_url == "http://ollama.ai.svc.cluster.local:11434/api/generate"
assert "Authorization" not in calls[0].headers
assert json.loads(calls[0].data)["options"]["num_predict"] == 256
@pytest.mark.parametrize("changes", [
{"model": "external/model"}, {"stream": True}, {"tools": []}, {"images": []},
{"context": []}, {"keep_alive": 0}, {"options": {"num_predict": -1}},
{"options": {"num_predict": 2049}}, {"options": {"num_gpu": 100}},
{"options": {"temperature": float("nan")}}, {"prompt": []},
])
def test_unsupported_payloads_never_reach_ollama(gateway, changes):
namespace, calls, request = gateway
payload = {"model": namespace["LAN_MODEL"], "prompt": "hello", "stream": False, **changes}
assert request("POST", "/api/generate", json.dumps(payload))[0] == 400
assert not calls
@pytest.mark.parametrize("length,status", [("-1", 413), ("131073", 413), ("invalid", 400)])
def test_invalid_request_lengths_are_rejected_before_read(gateway, length, status):
_, calls, request = gateway
assert request("POST", "/api/generate", "{}", **{"Content-Length": length})[0] == status
assert not calls
def test_outage_and_concurrent_request_fail_without_fallback(gateway):
namespace, calls, request = gateway
body = json.dumps({"model": namespace["LAN_MODEL"], "prompt": "hello", "stream": False})
namespace["_lan_inference"].acquire()
assert request("POST", "/api/generate", body)[0] == 429
namespace["_lan_inference"].release()
def unavailable(*args, **kwargs):
raise URLError("local service unavailable")
namespace["urlopen"] = unavailable
assert request("POST", "/api/generate", body)[0] == 503
assert namespace["_lan_inference"].acquire(timeout=1)
namespace["_lan_inference"].release()
assert not calls
@pytest.mark.parametrize("path", ["/api/pull", "/api/delete", "/v1/chat/completions", "/api/chat"])
def test_model_management_and_agent_routes_are_not_exposed(gateway, path):
_, calls, request = gateway
assert request("POST", path, "{}")[0] == 404
assert not calls
def test_batch_routes_share_lan_authentication(gateway, monkeypatch):
_, calls, request = gateway
monkeypatch.setattr(batch_api, "models_response", lambda: (200, {"models": []}))
monkeypatch.setattr(batch_api, "generate", lambda body: (200, {"batch": True}))
assert request(path="/api/batch/models", Authorization="")[0] == 401
assert request("POST", "/api/batch/generate", "{}", Authorization="")[0] == 401
assert request(path="/api/batch/models") == (200, {"models": []})
assert request("POST", "/api/batch/generate", "{}") == (200, {"batch": True})
assert not calls