140 lines
5.5 KiB
Python
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
|