From 4af316c739b3e08f8ecc6845d983e6aad7b0553c2a3c497c41a5a53b5c90684e Mon Sep 17 00:00:00 2001 From: larssand Date: Tue, 30 Jun 2026 12:16:42 +0200 Subject: [PATCH] improve agregate with status --- src/fgai/dashboard.py | 2 +- src/fgai/graylog_aggregate.py | 50 +++++++++++++++++++++++++++++++-- src/fgai/graylog_mcp.py | 7 ++++- src/fgai/monitor.py | 1 + tests/test_graylog_aggregate.py | 20 +++++++++++-- 5 files changed, 73 insertions(+), 7 deletions(-) diff --git a/src/fgai/dashboard.py b/src/fgai/dashboard.py index 9600e90..556b74b 100644 --- a/src/fgai/dashboard.py +++ b/src/fgai/dashboard.py @@ -406,7 +406,7 @@ async function refresh() { const profileNames = Object.fromEntries((data.stream_profiles || []).map(item => [item.stream_id, item.name || item.stream_id])); const profileReadiness = (data.profile_readiness || []).map(item => ({...item, profile_name: item.profile_name || profileNames[item.stream_id] || item.stream_id, stream_title: item.stream_name || item.stream_title || streamTitles[item.stream_id] || item.stream_id})); document.getElementById('diagnostics').innerHTML = - '

Stream Coverage

' + table(streamCoverage, [{label:'Stream', key:'stream_name'}, {label:'Enabled', key:'enabled', render:r => r.enabled ? 'yes' : 'no'}, {label:'Profile', render:r => esc(r.profile || 'missing')}, {label:'Entity Field', key:'entity_field'}, {label:'Tracked Fields', key:'tracked_fields'}, {label:'Ready Fields', key:'readiness'}, {label:'Raw Events', key:'events_fetched'}, {label:'Aggregate Events', key:'aggregate_events'}, {label:'Aggregate', key:'aggregate_status'}, {label:'Latest Event', key:'latest_event_time'}, {label:'Health', key:'health'}, {label:'Error', render:r => esc(r.error || '-')}], 'stream-coverage') + + '

Stream Coverage

' + table(streamCoverage, [{label:'Stream', key:'stream_name'}, {label:'Enabled', key:'enabled', render:r => r.enabled ? 'yes' : 'no'}, {label:'Profile', render:r => esc(r.profile || 'missing')}, {label:'Entity Field', key:'entity_field'}, {label:'Tracked Fields', key:'tracked_fields'}, {label:'Ready Fields', key:'readiness'}, {label:'Raw Events', key:'events_fetched'}, {label:'Aggregate Events', key:'aggregate_events'}, {label:'Aggregate', key:'aggregate_status'}, {label:'Aggregate Schema', key:'aggregate_schema_properties'}, {label:'Latest Event', key:'latest_event_time'}, {label:'Health', key:'health'}, {label:'Error', render:r => esc(r.error || '-')}], 'stream-coverage') + '

Cross-Source Correlations

' + table(correlations, [{label:'Entity', key:'entity', render:r => esc(`${r.entity || r.source_ip} (${r.entity_type || 'ip'})`)}, {label:'Streams', render:r => esc((r.streams || []).join(', '))}, {label:'Events', key:'events'}, {label:'Security Events', key:'security_events'}], 'correlations') + '

Entities

' + table(context.source_profiles || [], [{label:'Entity', key:'entity'}, {label:'Events', key:'events'}, {label:'UTM', key:'utm_events'}, {label:'Deny', key:'deny_or_threat_actions'}, {label:'Destinations', key:'distinct_destinations'}, {label:'Actions', render:r => esc((r.top_actions || []).join(', '))}], 'entities') + '

Profile Baseline Readiness

' + table(profileReadiness, [{label:'Profile', key:'profile_name'}, {label:'Stream', key:'stream_title'}, {label:'Field', key:'field'}, {label:'Buckets', key:'buckets'}, {label:'Age days', key:'age_days'}, {label:'Training days', key:'training_days'}, {label:'Ready', key:'ready', render:r => r.ready ? 'ready' : 'learning'}], 'profile-readiness') + diff --git a/src/fgai/graylog_aggregate.py b/src/fgai/graylog_aggregate.py index 6057e9c..d617eff 100644 --- a/src/fgai/graylog_aggregate.py +++ b/src/fgai/graylog_aggregate.py @@ -48,6 +48,21 @@ def _count_from_records(records: list[dict[str, object]]) -> int: return len(records) +def _schema_properties(schema: object) -> set[str]: + if not isinstance(schema, dict): + return set() + properties = schema.get("properties") + if not isinstance(properties, dict): + return set() + return {str(key) for key in properties} + + +def _filter_supported(arguments: dict[str, object], properties: set[str]) -> dict[str, object]: + if not properties: + return arguments + return {key: value for key, value in arguments.items() if key in properties} + + class GraylogAggregateSource: def __init__(self, client: GraylogMcpClient, stream: str, query: str = "*") -> None: self.client = client @@ -56,7 +71,27 @@ class GraylogAggregateSource: def fetch_count(self, *, range_seconds: int = 300) -> dict[str, object]: status = self.client.probe() - variants: list[dict[str, object]] = [ + tool_schemas = status.get("tool_schemas", {}) if isinstance(status.get("tool_schemas"), dict) else {} + aggregate_schema = tool_schemas.get("aggregate_messages", {}) if isinstance(tool_schemas, dict) else {} + properties = _schema_properties(aggregate_schema) + base = { + "query": self.query, + "streams": [self.stream] if self.stream else [], + "range_seconds": max(1, int(range_seconds)), + } + schema_variants: list[dict[str, object]] = [] + metric_keys = [key for key in ("metrics", "series") if key in properties] + group_keys = [key for key in ("group_by", "groups", "fields") if key in properties] + if properties and metric_keys: + for metric_key in metric_keys: + for metric_value in (["count()"], ["count"], [{"function": "count"}]): + candidate = {**base, metric_key: metric_value} + for group_key in group_keys: + candidate[group_key] = [] + schema_variants.append(_filter_supported(candidate, properties)) + if properties: + schema_variants.append(_filter_supported(base, properties)) + fallback_variants: list[dict[str, object]] = [ { "query": self.query, "streams": [self.stream] if self.stream else [], @@ -104,17 +139,23 @@ class GraylogAggregateSource: "range_seconds": max(1, int(range_seconds)), }, ] + variants = [*schema_variants, *fallback_variants] + seen_variants: set[str] = set() errors: list[str] = [] for arguments in variants: + variant_key = json.dumps(arguments, sort_keys=True) + if variant_key in seen_variants: + continue + seen_variants.add(variant_key) try: result = self.client.call_tool("aggregate_messages", arguments) except RuntimeError as exc: - errors.append(str(exc)) + errors.append(f"{arguments}: {exc}") continue content = result.get("result", {}).get("content", []) if isinstance(result.get("result"), dict) else [] if isinstance(result.get("result"), dict) and result["result"].get("isError"): detail = next((str(item.get("text")) for item in content if isinstance(item, dict) and item.get("type") == "text"), "Graylog aggregate failed") - errors.append(detail) + errors.append(f"{arguments}: {detail}") continue records: list[dict[str, object]] = [] for item in content if isinstance(content, list) else []: @@ -130,6 +171,7 @@ class GraylogAggregateSource: "aggregate_events": _count_from_records(records), "aggregate_records": len(records), "aggregate_arguments": arguments, + "aggregate_schema_properties": sorted(properties), } return { **status, @@ -138,4 +180,6 @@ class GraylogAggregateSource: "aggregate_events": 0, "aggregate_records": 0, "aggregate_error": "; ".join(errors[-3:]) or "aggregate_messages failed", + "aggregate_errors": errors[-6:], + "aggregate_schema_properties": sorted(properties), } diff --git a/src/fgai/graylog_mcp.py b/src/fgai/graylog_mcp.py index 8d1a402..0e5d859 100644 --- a/src/fgai/graylog_mcp.py +++ b/src/fgai/graylog_mcp.py @@ -72,8 +72,13 @@ class GraylogMcpClient: tools = self._call("tools/list") tool_list = tools.get("result", {}).get("tools", []) if isinstance(tools.get("result"), dict) else [] names = [str(item.get("name")) for item in tool_list if isinstance(item, dict) and item.get("name")] + schemas = { + str(item.get("name")): item.get("inputSchema") + for item in tool_list + if isinstance(item, dict) and item.get("name") and isinstance(item.get("inputSchema"), dict) + } version = initialized.get("result", {}).get("serverInfo", {}).get("version", "") if isinstance(initialized.get("result"), dict) else "" - return {"status": "connected", "tools": names, "server_version": version} + return {"status": "connected", "tools": names, "tool_schemas": schemas, "server_version": version} def call_tool(self, name: str, arguments: dict[str, object]) -> dict[str, object]: return self._call("tools/call", {"name": name, "arguments": arguments}) diff --git a/src/fgai/monitor.py b/src/fgai/monitor.py index 3983b36..20a660f 100644 --- a/src/fgai/monitor.py +++ b/src/fgai/monitor.py @@ -90,6 +90,7 @@ def _stream_coverage(runtime_values: dict[str, object], stream_profiles: dict[st "events_fetched": int(status.get("events_fetched", 0) or 0), "aggregate_events": int(status.get("aggregate_events", 0) or 0), "aggregate_status": str(status.get("aggregate_status", "")), + "aggregate_schema_properties": ", ".join(str(item) for item in status.get("aggregate_schema_properties", []) if item), "latest_event_time": str(status.get("latest_event_time", "")), "truncated": bool(status.get("truncated")), "partial": bool(status.get("partial")), diff --git a/tests/test_graylog_aggregate.py b/tests/test_graylog_aggregate.py index 0961586..1b15508 100644 --- a/tests/test_graylog_aggregate.py +++ b/tests/test_graylog_aggregate.py @@ -4,12 +4,16 @@ from fgai.graylog_aggregate import GraylogAggregateSource class _AggregateClient: - def __init__(self, responses): + def __init__(self, responses, schema=None): self.responses = list(responses) self.arguments = [] + self.schema = schema def probe(self): - return {"status": "connected"} + status = {"status": "connected"} + if self.schema: + status["tool_schemas"] = {"aggregate_messages": self.schema} + return status def call_tool(self, _name, arguments): self.arguments.append(arguments) @@ -41,6 +45,18 @@ class GraylogAggregateTests(unittest.TestCase): self.assertEqual(status["aggregate_events"], 42) self.assertEqual(len(client.arguments), 2) + def test_uses_tool_schema_to_avoid_unsupported_fields(self): + schema = {"properties": {"query": {}, "streams": {}, "range_seconds": {}, "series": {}}} + client = _AggregateClient([ + {"result": {"content": [{"type": "text", "text": '{"schema":[{"name":"count"}],"datarows":[[7]]}'}]}} + ], schema=schema) + + status = GraylogAggregateSource(client, "firewall").fetch_count() + + self.assertEqual(status["aggregate_events"], 7) + self.assertIn("series", client.arguments[0]) + self.assertNotIn("group_by", client.arguments[0]) + if __name__ == "__main__": unittest.main()