"""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