import unittest import time from fgai.graylog_source import GraylogStreamSource class _Client: def __init__(self): self.arguments = None self.probes = 0 def probe(self): self.probes += 1 return {"status": "connected"} def call_tool(self, _name, arguments): self.arguments = arguments return {"result": {"content": [{"type": "text", "text": '{"schema":[{"field":"client"},{"field":"server"},{"field":"result"}],"datarows":[["10.0.0.1","8.8.8.8","deny"],["10.0.0.2","8.8.8.8","accept"]]}' }]}} class _ErrorClient: def probe(self): return {"status": "connected"} def call_tool(self, _name, _arguments): return {"result": {"isError": True, "content": [{"type": "text", "text": "Tool call failed: timeout"}]}} class _ProbeErrorClient: def probe(self): raise RuntimeError("connection refused") class GraylogSourceTests(unittest.TestCase): def test_applies_custom_mapping_to_generic_stream_message(self): client = _Client() events, status = GraylogStreamSource( client, "vpn", field_mapping='{"srcip":"client","dstip":"server","action":"result"}' ).fetch() self.assertEqual(status["events_fetched"], 2) self.assertEqual(events[0].src_ip, "10.0.0.1") self.assertEqual(events[0].dst_ip, "8.8.8.8") self.assertEqual(events[0].action, "deny") self.assertEqual(client.arguments["streams"], ["vpn"]) self.assertEqual(client.arguments["range_seconds"], 300) self.assertIn("client", client.arguments["fields"]) self.assertEqual(client.arguments["offset"], 0) def test_uses_requested_range_seconds(self): client = _Client() GraylogStreamSource(client, "vpn").fetch(range_seconds=86_400) self.assertEqual(client.arguments["range_seconds"], 86_400) def test_uses_existing_probe_status_without_reprobing(self): client = _Client() GraylogStreamSource(client, "vpn").fetch(probe_status={"status": "connected"}) self.assertEqual(client.probes, 0) def test_marks_stream_truncated_when_max_events_reached(self): client = _Client() events, status = GraylogStreamSource(client, "vpn").fetch(max_events=2) self.assertEqual(len(events), 2) self.assertEqual(client.arguments["limit"], 2) self.assertTrue(status["truncated"]) def test_uses_larger_single_page_for_raw_sample(self): client = _Client() GraylogStreamSource(client, "vpn").fetch(max_events=5000) self.assertEqual(client.arguments["limit"], 5000) def test_returns_partial_status_instead_of_raising_on_search_error(self): events, status = GraylogStreamSource(_ErrorClient(), "vpn").fetch() self.assertEqual(events, []) self.assertTrue(status["partial"]) self.assertIn("graylog_search_error", status["error"]) def test_deadline_stops_raw_fetch_before_search(self): client = _Client() events, status = GraylogStreamSource(client, "vpn").fetch(deadline_monotonic=time.monotonic() - 1) self.assertEqual(events, []) self.assertTrue(status["partial"]) self.assertEqual(status["error"], "poll_budget_exceeded") self.assertIsNone(client.arguments) def test_returns_partial_status_instead_of_raising_on_probe_error(self): events, status = GraylogStreamSource(_ProbeErrorClient(), "vpn").fetch() self.assertEqual(events, []) self.assertTrue(status["partial"]) self.assertIn("probe_error", status["error"]) def test_requests_selected_profile_fields(self): client = _Client() GraylogStreamSource(client, "windows", profile_fields=("TargetUserName", "EventID")).fetch() self.assertIn("targetusername", client.arguments["fields"]) self.assertIn("eventid", client.arguments["fields"]) def test_requests_common_firewall_alias_fields(self): client = _Client() GraylogStreamSource(client, "firewall").fetch() self.assertIn("src_addr", client.arguments["fields"]) self.assertIn("dest_port", client.arguments["fields"]) self.assertIn("fw_action", client.arguments["fields"]) self.assertIn("full_message", client.arguments["fields"]) def test_requests_winlogbeat_discovery_fields_before_profile_exists(self): client = _Client() GraylogStreamSource(client, "windows").fetch() self.assertIn("winlog.event_id", client.arguments["fields"]) self.assertIn("event.code", client.arguments["fields"]) self.assertIn("user.name", client.arguments["fields"]) self.assertIn("host.name", client.arguments["fields"]) self.assertIn("source.ip", client.arguments["fields"]) self.assertIn("winlog.channel", client.arguments["fields"]) self.assertIn("process.name", client.arguments["fields"]) if __name__ == "__main__": unittest.main()