Files
LOGAR/tests/test_server.py
T

334 lines
14 KiB
Python

import os
import sys
import json
import socket
import struct
import sqlite3
import unittest
import urllib.request
import warnings
from datetime import datetime, timezone, timedelta
warnings.filterwarnings("ignore")
# Ensure parent directory and src directory are in path to import Server
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "src")))
import Server
import pgpy
class TestServerComponent(unittest.TestCase):
def setUp(self):
self.test_db = "test_server_state.db"
self.test_config = "test_server_config.json"
if os.path.exists(self.test_db):
os.remove(self.test_db)
if os.path.exists(self.test_config):
os.remove(self.test_config)
def tearDown(self):
if os.path.exists(self.test_db):
try:
os.remove(self.test_db)
except Exception:
pass
if os.path.exists(self.test_config):
try:
os.remove(self.test_config)
except Exception:
pass
def test_first_run_config_and_keypair_generation(self):
config = Server.load_or_init_config(self.test_config)
self.assertTrue(os.path.exists(self.test_config))
self.assertIn("server_fingerprint", config)
self.assertIn("public_key", config)
self.assertIn("private_key", config)
self.assertIn("auth_token", config)
self.assertNotIn("site_name", config)
# Verify keypair
priv_key, _ = pgpy.PGPKey.from_blob(config["private_key"])
pub_key, _ = pgpy.PGPKey.from_blob(config["public_key"])
self.assertEqual(str(pub_key.fingerprint), config["server_fingerprint"])
def test_create_client_config(self):
Server.load_or_init_config(self.test_config)
client_out = "test_client_out.json"
try:
client_conf = Server.create_client_config(
server_host="10.0.0.1",
server_port=9443,
output_path=client_out,
config_path=self.test_config
)
self.assertTrue(os.path.exists(client_out))
self.assertEqual(client_conf["server_host"], "10.0.0.1")
self.assertEqual(client_conf["server_port"], 9443)
# Verify no machine name or site_name is included
self.assertNotIn("server_name", client_conf)
self.assertNotIn("name", client_conf)
self.assertNotIn("site_name", client_conf)
finally:
if os.path.exists(client_out):
os.remove(client_out)
def test_4_run_rule_and_12h_window(self):
Server.init_db(self.test_db)
log_entry = {
"server": "app-worker-01.corp.local",
"signature": "PostgresConnWarning",
"severity": "WARNING",
"message": "Connection to database pool near capacity: 85%",
"os_type": "linux"
}
payload = {
"server": "app-worker-01.corp.local",
"logs": [log_entry]
}
# Runs 1 to 3: WARNING should remain TRANSIENT
for run_idx in range(1, 4):
res = Server.process_ingested_logs(payload, self.test_db, window_hours=12, min_runs=4)
self.assertEqual(res["status"], "success")
self.assertEqual(res["promoted_verified"], 0)
conn = sqlite3.connect(self.test_db)
c = conn.cursor()
c.execute("SELECT run_count, status FROM active_issues WHERE signature = ?", ("PostgresConnWarning",))
row = c.fetchone()
conn.close()
self.assertEqual(row[0], 3)
self.assertEqual(row[1], "TRANSIENT")
# Run 4: promotes WARNING to VERIFIED!
res4 = Server.process_ingested_logs(payload, self.test_db, window_hours=12, min_runs=4)
self.assertEqual(res4["promoted_verified"], 1)
conn = sqlite3.connect(self.test_db)
c = conn.cursor()
c.execute("SELECT run_count, status FROM active_issues WHERE signature = ?", ("PostgresConnWarning",))
row = c.fetchone()
conn.close()
self.assertEqual(row[0], 4)
self.assertEqual(row[1], "VERIFIED")
def test_error_immediate_pass(self):
Server.init_db(self.test_db)
log_entry = {
"server": "app-worker-01.corp.local",
"signature": "KernelPanicCritical",
"severity": "ERROR",
"message": "Kernel panic - not syncing: Fatal hardware error",
"os_type": "linux"
}
payload = {
"server": "app-worker-01.corp.local",
"logs": [log_entry]
}
# Run 1: ERROR must immediately promote to VERIFIED
res = Server.process_ingested_logs(payload, self.test_db, window_hours=12, min_runs=4)
self.assertEqual(res["status"], "success")
self.assertEqual(res["promoted_verified"], 1)
conn = sqlite3.connect(self.test_db)
c = conn.cursor()
c.execute("SELECT run_count, status, severity FROM active_issues WHERE signature = ?", ("KernelPanicCritical",))
row = c.fetchone()
conn.close()
self.assertIsNotNone(row)
self.assertEqual(row[0], 1)
self.assertEqual(row[1], "VERIFIED")
self.assertEqual(row[2], "ERROR")
def test_server_severity_filtering(self):
Server.init_db(self.test_db)
payload = {
"server": "app-worker-01.corp.local",
"logs": [
{"server": "app-worker-01", "signature": "SigInfo", "severity": "INFO", "message": "Info msg", "os_type": "linux"},
{"server": "app-worker-01", "signature": "SigWarn", "severity": "WARNING", "message": "Warn msg", "os_type": "linux"},
{"server": "app-worker-01", "signature": "SigErr", "severity": "ERROR", "message": "Err msg", "os_type": "linux"},
{"server": "app-worker-01", "signature": "SigDebug", "severity": "DEBUG", "message": "Debug msg", "os_type": "linux"},
{"server": "app-worker-01", "signature": "SigTrace", "severity": "TRACE", "message": "Trace msg", "os_type": "linux"}
]
}
res = Server.process_ingested_logs(payload, self.test_db, window_hours=12, min_runs=4)
self.assertEqual(res["status"], "success")
conn = sqlite3.connect(self.test_db)
c = conn.cursor()
c.execute("SELECT signature, status FROM active_issues ORDER BY signature")
rows = dict(c.fetchall())
conn.close()
self.assertIn("SigInfo", rows)
self.assertIn("SigWarn", rows)
self.assertIn("SigErr", rows)
self.assertNotIn("SigDebug", rows)
self.assertNotIn("SigTrace", rows)
# SigErr is immediately VERIFIED; SigWarn and SigInfo are TRANSIENT on run 1
self.assertEqual(rows["SigErr"], "VERIFIED")
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)
def test_cert_validity_and_hub_pki_renewal(self):
test_cert_dir = "test_certs_renew"
try:
ca_cert, ca_key, ca_pem, _ = Server.enrollment.generate_ca_if_needed(cert_dir=test_cert_dir)
srv_cert, srv_key, srv_pem, _ = Server.enrollment.generate_server_cert_if_needed(
ca_cert, ca_key, hostnames=["127.0.0.1"], cert_dir=test_cert_dir
)
# 1. Freshly generated certificates should NOT be expiring soon with standard 30-day threshold
self.assertFalse(Server.enrollment.is_cert_expiring_soon(ca_pem, threshold_days=30))
self.assertFalse(Server.enrollment.is_cert_expiring_soon(srv_pem, threshold_days=30))
# 2. Huge threshold (e.g. 5000 days) should flag expiration
self.assertTrue(Server.enrollment.is_cert_expiring_soon(srv_pem, threshold_days=5000))
# 3. check_and_renew_hub_pki with standard threshold should report no renewal needed
ca_renewed, srv_renewed = Server.enrollment.check_and_renew_hub_pki(cert_dir=test_cert_dir, threshold_days=30)
self.assertFalse(ca_renewed)
self.assertFalse(srv_renewed)
# 4. In-flight reload of SSLContext
ssl_ctx = Server.init_mtls_server_context(cert_dir=test_cert_dir)
Server.SERVER_STATE["ssl_ctx"] = ssl_ctx
Server.SERVER_STATE["config"] = {"db_path": self.test_db, "cert_dir": test_cert_dir, "tcp_host": "127.0.0.1"}
# Trigger rotation using high threshold
rotated = Server.check_and_rotate_server_certs(cert_dir=test_cert_dir, hostnames=["127.0.0.1"], threshold_days=5000)
self.assertTrue(rotated)
# Check that backup files were generated
bak_files = [f for f in os.listdir(test_cert_dir) if f.endswith(".bak")]
self.assertGreater(len(bak_files), 0)
finally:
import shutil
if os.path.exists(test_cert_dir):
shutil.rmtree(test_cert_dir, ignore_errors=True)
if __name__ == "__main__":
unittest.main()