87 lines
3.4 KiB
Python
87 lines
3.4 KiB
Python
import unittest
|
|
|
|
from fgai.graylog_source import GraylogStreamSource
|
|
|
|
|
|
class _Client:
|
|
def __init__(self):
|
|
self.arguments = None
|
|
|
|
def probe(self):
|
|
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_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_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_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"])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|