"""Tests für den Schlüssel-Proxy. Ausführen: python3 -m unittest discover -s tools/ki-proxy""" from __future__ import annotations import base64 import http.client import json import logging import os import sys import threading import unittest from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer sys.path.insert(0, os.path.dirname(__file__)) import ki_proxy # noqa: E402 KEYS = {"ki-bauin/itm-ki": "test-ki-schluessel", "ki-bauin/dataloader": "test-dl-schluessel", "ki-bauin/ocr": None} class FakeUpstream(BaseHTTPRequestHandler): """Platzhalter für api-werk.de: merkt sich die letzte Anfrage.""" protocol_version = "HTTP/1.1" last: dict = {} release_stream = threading.Event() def log_message(self, format, *args): # noqa: A002 return def _record(self) -> bytes: length = int(self.headers.get("Content-Length") or 0) body = self.rfile.read(length) if length else b"" FakeUpstream.last = {"method": self.command, "path": self.path, "headers": dict(self.headers.items()), "body": body} return body def do_GET(self): self._record() if self.path.startswith("/v1/ki/stream"): self.send_response(200) self.send_header("Content-Type", "text/event-stream") self.send_header("Transfer-Encoding", "chunked") self.end_headers() self._chunk(b"data: erster\n\n") FakeUpstream.release_stream.wait(5) self._chunk(b"data: [DONE]\n\n") self.wfile.write(b"0\r\n\r\n") return status = 409 if self.path.startswith("/v1/dataloader/jobs/") else 200 payload = json.dumps({"data": [{"id": "chat"}, {"id": "embed"}]}).encode() self.send_response(status) self.send_header("Content-Type", "application/json") self.send_header("Content-Length", str(len(payload))) self.end_headers() self.wfile.write(payload) def do_POST(self): body = self._record() payload = json.dumps({"job_id": "abc", "bytes": len(body)}).encode() self.send_response(202) self.send_header("Content-Type", "application/json") self.send_header("Content-Length", str(len(payload))) self.end_headers() self.wfile.write(payload) def _chunk(self, data: bytes) -> None: self.wfile.write(f"{len(data):x}\r\n".encode() + data + b"\r\n") self.wfile.flush() class ProxyTest(unittest.TestCase): @classmethod def setUpClass(cls): cls.upstream = ThreadingHTTPServer(("127.0.0.1", 0), FakeUpstream) threading.Thread(target=cls.upstream.serve_forever, daemon=True).start() upstream_url = f"http://127.0.0.1:{cls.upstream.server_address[1]}" cls.proxy = ki_proxy.start_servers(["127.0.0.1"], 0, dict(KEYS), upstream_url)[0] cls.port = cls.proxy.server_address[1] @classmethod def tearDownClass(cls): cls.proxy.shutdown() cls.upstream.shutdown() def setUp(self): FakeUpstream.last = {} FakeUpstream.release_stream.clear() self.logs = [] handler = logging.Handler() handler.emit = lambda record: self.logs.append(record.getMessage()) ki_proxy.log.addHandler(handler) ki_proxy.log.setLevel(logging.INFO) self.addCleanup(ki_proxy.log.removeHandler, handler) def request(self, method, path, body=None, headers=None): conn = http.client.HTTPConnection("127.0.0.1", self.port, timeout=10) conn.request(method, path, body=body, headers=headers or {}) return conn, conn.getresponse() def test_setzt_den_schluessel_der_route_ein(self): conn, response = self.request("GET", "/v1/ki/models") self.assertEqual(200, response.status) self.assertEqual(["chat", "embed"], [m["id"] for m in json.loads(response.read())["data"]]) self.assertEqual("Bearer test-ki-schluessel", FakeUpstream.last["headers"]["Authorization"]) conn.close() def test_ueberschreibt_mitgeschickte_authorization(self): conn, response = self.request("GET", "/v1/ki/models", headers={"Authorization": "Bearer fremd"}) response.read() self.assertEqual("Bearer test-ki-schluessel", FakeUpstream.last["headers"]["Authorization"]) conn.close() def test_upload_kommt_unveraendert_an_mit_eigenem_schluessel(self): body = os.urandom(300_000) headers = {"Content-Type": "multipart/form-data; boundary=x"} conn, response = self.request("POST", "/v1/dataloader/jobs?x=1", body=body, headers=headers) self.assertEqual(202, response.status) self.assertEqual(len(body), json.loads(response.read())["bytes"]) self.assertEqual(body, FakeUpstream.last["body"]) self.assertEqual("/v1/dataloader/jobs?x=1", FakeUpstream.last["path"]) self.assertEqual("Bearer test-dl-schluessel", FakeUpstream.last["headers"]["Authorization"]) conn.close() def test_reicht_fehlerstatus_durch(self): conn, response = self.request("GET", "/v1/dataloader/jobs/abc/result?format=markdown") self.assertEqual(409, response.status) response.read() conn.close() def test_fehlender_schluessel_ergibt_503_ohne_weiterleitung(self): conn, response = self.request("GET", "/v1/ocr/jobs/abc") self.assertEqual(503, response.status) self.assertIn("ki-bauin/ocr", json.loads(response.read())["error"]) self.assertEqual({}, FakeUpstream.last) conn.close() def test_unbekannter_pfad_ergibt_404(self): conn, response = self.request("GET", "/admin") self.assertEqual(404, response.status) response.read() self.assertEqual({}, FakeUpstream.last) conn.close() def test_streaming_kommt_sofort_an(self): conn, response = self.request("GET", "/v1/ki/stream") self.assertEqual(200, response.status) first = response.read1(1024) self.assertIn(b"erster", first) # kommt an, bevor der Platzhalter den Rest freigibt self.assertNotIn(b"[DONE]", first) FakeUpstream.release_stream.set() rest = response.read() self.assertIn(b"[DONE]", rest) conn.close() def test_protokoll_enthaelt_keine_schluessel(self): for path in ("/v1/ki/models", "/v1/ocr/x", "/unbekannt"): conn, response = self.request("GET", path, headers={"Authorization": "Bearer fremd"}) response.read() conn.close() self.assertTrue(self.logs) for line in self.logs: for secret in ("test-ki-schluessel", "test-dl-schluessel", "fremd", "Bearer"): self.assertNotIn(secret, line) class TresorAusgabeTest(unittest.TestCase): def test_wertet_gefundene_und_fehlende_eintraege_aus(self): wert = base64.b64encode("zki_Äbc123".encode()).decode() output = f"ki-bauin/itm-ki\t{wert}\r\nki-bauin/ocr\t-\r\n" self.assertEqual({"ki-bauin/itm-ki": "zki_Äbc123", "ki-bauin/ocr": None}, ki_proxy.parse_tresor_output(output)) def test_check_meldet_fehlende_schluessel_ohne_werte_auszugeben(self): logs = [] handler = logging.Handler() handler.emit = lambda record: logs.append(record.getMessage()) ki_proxy.log.addHandler(handler) try: code = ki_proxy.main(["--check"], reader=lambda targets: dict(KEYS)) finally: ki_proxy.log.removeHandler(handler) self.assertEqual(1, code) joined = "\n".join(logs) self.assertIn("FEHLT", joined) self.assertNotIn("test-ki-schluessel", joined) if __name__ == "__main__": unittest.main()