test(server): add PKI, enrollment, and license quota tests in test_server.py
This commit is contained in:
@@ -176,6 +176,121 @@ class TestServerComponent(unittest.TestCase):
|
||||
self.assertEqual(rows["SigWarn"], "TRANSIENT")
|
||||
self.assertEqual(rows["SigInfo"], "TRANSIENT")
|
||||
|
||||
def test_license_schema_and_pki_generation(self):
|
||||
secret = "test-secret-12345"
|
||||
Server.init_db(self.test_db, enrollment_secret=secret, max_seats=5)
|
||||
|
||||
conn = sqlite3.connect(self.test_db)
|
||||
c = conn.cursor()
|
||||
c.execute("SELECT max_seats, enrollment_secret FROM license_config WHERE id = 1")
|
||||
row = c.fetchone()
|
||||
conn.close()
|
||||
|
||||
self.assertIsNotNone(row)
|
||||
self.assertEqual(row[0], 5)
|
||||
self.assertEqual(row[1], secret)
|
||||
|
||||
# Test Dynamic PKI
|
||||
test_cert_dir = "test_certs_pki"
|
||||
try:
|
||||
ca_cert, ca_key, ca_pem, ca_key_pem = Server.enrollment.generate_ca_if_needed(cert_dir=test_cert_dir)
|
||||
self.assertIn("BEGIN CERTIFICATE", ca_pem)
|
||||
self.assertIn("BEGIN RSA PRIVATE KEY", ca_key_pem)
|
||||
|
||||
srv_cert, srv_key, srv_pem, srv_key_pem = Server.enrollment.generate_server_cert_if_needed(
|
||||
ca_cert, ca_key, hostnames=["127.0.0.1", "localhost"], cert_dir=test_cert_dir
|
||||
)
|
||||
self.assertIn("BEGIN CERTIFICATE", srv_pem)
|
||||
|
||||
client_cert_pem, client_key_pem = Server.enrollment.issue_client_cert("node-test-1", ca_cert, ca_key)
|
||||
self.assertIn("BEGIN CERTIFICATE", client_cert_pem)
|
||||
self.assertIn("BEGIN RSA PRIVATE KEY", client_key_pem)
|
||||
|
||||
fp = Server.enrollment.calculate_cert_fingerprint(client_cert_pem)
|
||||
self.assertEqual(len(fp), 64)
|
||||
finally:
|
||||
import shutil
|
||||
if os.path.exists(test_cert_dir):
|
||||
shutil.rmtree(test_cert_dir, ignore_errors=True)
|
||||
|
||||
def test_enrollment_endpoint_and_seat_quota(self):
|
||||
from fastapi import HTTPException
|
||||
secret = "super-secret-enrollment"
|
||||
max_seats = 2
|
||||
Server.init_db(self.test_db, enrollment_secret=secret, max_seats=max_seats)
|
||||
|
||||
test_cert_dir = "test_certs_enroll"
|
||||
try:
|
||||
ca_cert, ca_key, ca_pem, _ = Server.enrollment.generate_ca_if_needed(cert_dir=test_cert_dir)
|
||||
Server.SERVER_STATE["config"] = {"db_path": self.test_db}
|
||||
Server.SERVER_STATE["ca_cert"] = ca_cert
|
||||
Server.SERVER_STATE["ca_key"] = ca_key
|
||||
Server.SERVER_STATE["ca_cert_pem"] = ca_pem
|
||||
|
||||
# 1. Invalid secret should raise 403
|
||||
bad_req = Server.ClientEnrollRequest(
|
||||
client_id="client-1",
|
||||
hostname="host-1",
|
||||
os="linux",
|
||||
enrollment_secret="wrong-secret"
|
||||
)
|
||||
with self.assertRaises(HTTPException) as cm:
|
||||
Server.enroll_client(bad_req)
|
||||
self.assertEqual(cm.exception.status_code, 403)
|
||||
|
||||
# 2. Valid enrollment for client 1
|
||||
req1 = Server.ClientEnrollRequest(
|
||||
client_id="client-1",
|
||||
hostname="host-1",
|
||||
os="linux",
|
||||
enrollment_secret=secret
|
||||
)
|
||||
resp1 = Server.enroll_client(req1)
|
||||
self.assertIn("client_cert", resp1)
|
||||
self.assertIn("client_key", resp1)
|
||||
self.assertEqual(resp1["ca_cert"], ca_pem)
|
||||
|
||||
# 3. Valid enrollment for client 2
|
||||
req2 = Server.ClientEnrollRequest(
|
||||
client_id="client-2",
|
||||
hostname="host-2",
|
||||
os="windows",
|
||||
enrollment_secret=secret
|
||||
)
|
||||
resp2 = Server.enroll_client(req2)
|
||||
self.assertIn("client_cert", resp2)
|
||||
|
||||
# 4. Seat quota exhausted: client 3 should raise 403
|
||||
req3 = Server.ClientEnrollRequest(
|
||||
client_id="client-3",
|
||||
hostname="host-3",
|
||||
os="linux",
|
||||
enrollment_secret=secret
|
||||
)
|
||||
with self.assertRaises(HTTPException) as cm:
|
||||
Server.enroll_client(req3)
|
||||
self.assertEqual(cm.exception.status_code, 403)
|
||||
self.assertIn("License seat limit reached", cm.exception.detail)
|
||||
|
||||
# 5. Re-enrollment for existing client 1 should succeed
|
||||
resp1_re = Server.enroll_client(req1)
|
||||
self.assertIn("client_cert", resp1_re)
|
||||
|
||||
# 6. Revoked client should be rejected
|
||||
conn = sqlite3.connect(self.test_db)
|
||||
conn.execute("UPDATE clients SET status = 'revoked' WHERE client_id = 'client-1'")
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
with self.assertRaises(HTTPException) as cm:
|
||||
Server.enroll_client(req1)
|
||||
self.assertEqual(cm.exception.status_code, 403)
|
||||
self.assertIn("revoked", cm.exception.detail)
|
||||
finally:
|
||||
import shutil
|
||||
if os.path.exists(test_cert_dir):
|
||||
shutil.rmtree(test_cert_dir, ignore_errors=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user