215 lines
8.4 KiB
Python
215 lines
8.4 KiB
Python
import os
|
|
import sys
|
|
import json
|
|
import struct
|
|
import unittest
|
|
import warnings
|
|
|
|
warnings.filterwarnings("ignore")
|
|
|
|
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 Win_Client
|
|
import pgpy
|
|
from pgpy.constants import PubKeyAlgorithm, KeyFlags, HashAlgorithm, SymmetricKeyAlgorithm, CompressionAlgorithm
|
|
|
|
|
|
class TestWinClientComponent(unittest.TestCase):
|
|
def setUp(self):
|
|
self.dummy_config = "test_win_client_config.json"
|
|
# Generate dummy PGP key for testing
|
|
key = pgpy.PGPKey.new(PubKeyAlgorithm.RSAEncryptOrSign, 2048)
|
|
uid = pgpy.PGPUID.new("TestHub")
|
|
key.add_uid(
|
|
uid,
|
|
usage={KeyFlags.EncryptCommunications, KeyFlags.EncryptStorage},
|
|
hashes=[HashAlgorithm.SHA256],
|
|
ciphers=[SymmetricKeyAlgorithm.AES256],
|
|
compression=[CompressionAlgorithm.Uncompressed]
|
|
)
|
|
self.server_priv = key
|
|
self.server_pub = key.pubkey
|
|
self.fingerprint = str(key.pubkey.fingerprint)
|
|
|
|
with open(self.dummy_config, "w", encoding="utf-8") as f:
|
|
json.dump({
|
|
"server_host": "127.0.0.1",
|
|
"server_port": 9443,
|
|
"server_fingerprint": self.fingerprint,
|
|
"server_public_key": str(self.server_pub),
|
|
"auth_token": "secret-test-token"
|
|
}, f)
|
|
|
|
def tearDown(self):
|
|
if os.path.exists(self.dummy_config):
|
|
try:
|
|
os.remove(self.dummy_config)
|
|
except Exception:
|
|
pass
|
|
|
|
def test_client_config_anonymity(self):
|
|
config = Win_Client.load_config(self.dummy_config)
|
|
self.assertNotIn("server_name", config)
|
|
self.assertNotIn("name", config)
|
|
self.assertNotIn("site_name", config)
|
|
self.assertEqual(config["server_fingerprint"], self.fingerprint)
|
|
|
|
def test_get_machine_identifier(self):
|
|
machine_id = Win_Client.get_machine_identifier()
|
|
self.assertIsInstance(machine_id, str)
|
|
self.assertGreater(len(machine_id), 0)
|
|
self.assertNotEqual(machine_id, "localhost")
|
|
|
|
def test_encryption_and_envelope_creation(self):
|
|
config = Win_Client.load_config(self.dummy_config)
|
|
logs = [{
|
|
"server": Win_Client.get_machine_identifier(),
|
|
"signature": "TestWinSignature",
|
|
"severity": "WARNING",
|
|
"message": "Disk space threshold warning"
|
|
}]
|
|
|
|
pub_key, _ = pgpy.PGPKey.from_blob(config["server_public_key"])
|
|
payload = {
|
|
"server": Win_Client.get_machine_identifier(),
|
|
"logs": logs
|
|
}
|
|
msg = pgpy.PGPMessage.new(json.dumps(payload))
|
|
enc = pub_key.encrypt(msg)
|
|
self.assertTrue(str(enc).startswith("-----BEGIN PGP MESSAGE-----"))
|
|
|
|
# Decrypt with private key to verify end-to-end payload integrity
|
|
dec = self.server_priv.decrypt(enc)
|
|
restored = json.loads(dec.message)
|
|
self.assertEqual(restored["logs"][0]["signature"], "TestWinSignature")
|
|
|
|
def test_framing_protocol(self):
|
|
envelope_data = json.dumps({"test": "data"}).encode("utf-8")
|
|
frame = struct.pack(">I", len(envelope_data)) + envelope_data
|
|
self.assertEqual(len(frame), 4 + len(envelope_data))
|
|
length = struct.unpack(">I", frame[:4])[0]
|
|
self.assertEqual(length, len(envelope_data))
|
|
|
|
def test_windows_event_filtering_and_severity_map(self):
|
|
# sev_map: 1 -> ERROR, 2 -> WARNING, 4 -> INFO
|
|
sev_map = {1: "ERROR", 2: "WARNING", 4: "INFO"}
|
|
raw_event_types = [1, 2, 4, 8, 16] # 8 is Audit Success, 16 is Audit Failure
|
|
filtered = [sev_map[et] for et in raw_event_types if et in sev_map]
|
|
self.assertEqual(filtered, ["ERROR", "WARNING", "INFO"])
|
|
|
|
def test_state_lifecycle(self):
|
|
state_path = "test_win_state.json"
|
|
try:
|
|
# 1. Load non-existent returns empty dict
|
|
state = Win_Client.load_state(state_path)
|
|
self.assertEqual(state, {})
|
|
|
|
# 2. Stage new record number and sent IDs
|
|
state["new_last_record_number"] = 42
|
|
state["new_sent_record_ids"] = ["42:2026-09-04T12:00:00"]
|
|
|
|
# 3. Commit state moves staged keys to permanent and writes atomically
|
|
Win_Client.commit_state(state, state_path)
|
|
self.assertNotIn("new_last_record_number", state)
|
|
self.assertEqual(state.get("last_record_number"), 42)
|
|
self.assertEqual(state.get("sent_record_ids"), ["42:2026-09-04T12:00:00"])
|
|
|
|
# 4. Reload from disk
|
|
reloaded = Win_Client.load_state(state_path)
|
|
self.assertEqual(reloaded.get("last_record_number"), 42)
|
|
self.assertEqual(reloaded.get("sent_record_ids"), ["42:2026-09-04T12:00:00"])
|
|
finally:
|
|
if os.path.exists(state_path):
|
|
os.remove(state_path)
|
|
|
|
def test_duplicate_suppression_and_lookback_logic(self):
|
|
from datetime import datetime, timezone, timedelta
|
|
now = datetime.now()
|
|
cutoff_time = now - timedelta(hours=24)
|
|
|
|
# Mock event object
|
|
class MockEvent:
|
|
def __init__(self, rec_num, time_gen, event_type, source="TestApp", inserts=None):
|
|
self.RecordNumber = rec_num
|
|
self.TimeGenerated = time_gen
|
|
self.EventType = event_type
|
|
self.SourceName = source
|
|
self.StringInserts = inserts or ["Test"]
|
|
|
|
# Events read backwards: newest (rec 103) down to older (rec 99)
|
|
mock_events = [
|
|
# 1. New error within last 24h
|
|
MockEvent(103, now - timedelta(hours=1), 1),
|
|
# 2. New warning within last 24h
|
|
MockEvent(102, now - timedelta(hours=2), 2),
|
|
# 3. Already sent event (rec 101)
|
|
MockEvent(101, now - timedelta(hours=3), 4),
|
|
# 4. Event at or before last_record_number (rec 100) -> should stop backwards scan
|
|
MockEvent(100, now - timedelta(hours=4), 1),
|
|
# 5. Old event (> 24h)
|
|
MockEvent(99, now - timedelta(hours=26), 1),
|
|
]
|
|
|
|
state = {
|
|
"last_record_number": 100,
|
|
"sent_record_ids": ["101:" + (now - timedelta(hours=3)).isoformat()]
|
|
}
|
|
|
|
sev_map = {1: "ERROR", 2: "WARNING", 4: "INFO"}
|
|
logs = []
|
|
last_record_number = int(state.get("last_record_number", 0))
|
|
sent_record_ids = set(state.get("sent_record_ids", []))
|
|
newest_record_number = 0
|
|
|
|
for event in mock_events:
|
|
rec_num = int(event.RecordNumber)
|
|
if newest_record_number == 0:
|
|
newest_record_number = rec_num
|
|
|
|
if event.TimeGenerated < cutoff_time:
|
|
break
|
|
|
|
if last_record_number > 0 and newest_record_number >= last_record_number:
|
|
if rec_num <= last_record_number:
|
|
break
|
|
|
|
rec_id = f"{rec_num}:{event.TimeGenerated.isoformat()}"
|
|
if rec_id in sent_record_ids:
|
|
continue
|
|
|
|
if event.EventType in sev_map:
|
|
logs.append(rec_num)
|
|
|
|
# Only rec 103 and 102 should be processed (101 is already sent, <= 100 breaks early)
|
|
self.assertEqual(logs, [103, 102])
|
|
|
|
def test_mtls_client_certificate_handling(self):
|
|
import shutil
|
|
test_dir = "test_win_mtls_certs"
|
|
os.makedirs(test_dir, exist_ok=True)
|
|
try:
|
|
from src import server_enrollment as se
|
|
ca_cert, ca_key, ca_pem, _ = se.generate_ca_if_needed(test_dir)
|
|
client_cert_pem, client_key_pem = se.issue_client_cert("win-client-test", ca_cert, ca_key)
|
|
|
|
with open(os.path.join(test_dir, "ca.crt"), "w") as f:
|
|
f.write(ca_pem)
|
|
with open(os.path.join(test_dir, "client.crt"), "w") as f:
|
|
f.write(client_cert_pem)
|
|
with open(os.path.join(test_dir, "client.key"), "w") as f:
|
|
f.write(client_key_pem)
|
|
|
|
# Test missing certs exception
|
|
empty_dir = "test_empty_certs"
|
|
os.makedirs(empty_dir, exist_ok=True)
|
|
with self.assertRaises(FileNotFoundError):
|
|
Win_Client.get_tls_socket("127.0.0.1", 9443, empty_dir)
|
|
shutil.rmtree(empty_dir, ignore_errors=True)
|
|
finally:
|
|
shutil.rmtree(test_dir, ignore_errors=True)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|