diff --git a/src/fgai/llm.py b/src/fgai/llm.py index 5e35338..531bcec 100644 --- a/src/fgai/llm.py +++ b/src/fgai/llm.py @@ -76,23 +76,117 @@ def ollama_summary( return str(data.get("response", "")).strip() -def ollama_dashboard_assessment(analysis: dict[str, object], model: str | None = None, timeout: int | None = None) -> str: - compact = { +def _compact_field_deviations(field_deviations: object, *, limit: int = 12) -> list[dict[str, object]]: + rows: list[dict[str, object]] = [] + if not isinstance(field_deviations, dict): + return rows + for entity, deviations in field_deviations.items(): + if not isinstance(deviations, list): + continue + for item in deviations: + if not isinstance(item, dict): + continue + rows.append({ + "entity": entity, + "stream": item.get("stream_name") or item.get("stream_title") or item.get("stream_id"), + "detector": item.get("detector"), + "field": item.get("field"), + "score": item.get("score"), + "confidence": item.get("confidence"), + "reason": item.get("reason"), + "current": item.get("current"), + "baseline": item.get("baseline"), + "baseline_age_days": item.get("baseline_age_days"), + "sample_values": item.get("sample_values", [])[:5] if isinstance(item.get("sample_values"), list) else [], + "feedback": item.get("feedback", "unreviewed"), + }) + return sorted(rows, key=lambda item: int(item.get("score") or 0), reverse=True)[:limit] + + +def _compact_event_context(event_context: object) -> dict[str, object]: + if not isinstance(event_context, dict): + return {} + return { + "source_profiles": event_context.get("source_profiles", [])[:10] if isinstance(event_context.get("source_profiles"), list) else [], + "related_activity": event_context.get("related_activity", [])[:20] if isinstance(event_context.get("related_activity"), list) else [], + "entity_labels": event_context.get("entity_labels", {}) if isinstance(event_context.get("entity_labels"), dict) else {}, + } + + +def _compact_mcp_capabilities(capabilities: object) -> dict[str, object]: + if not isinstance(capabilities, dict): + return {} + mcp = capabilities.get("graylog_mcp", {}) if isinstance(capabilities.get("graylog_mcp"), dict) else {} + profile_advisor = capabilities.get("profile_advisor", {}) if isinstance(capabilities.get("profile_advisor"), dict) else {} + threat_intel = capabilities.get("threat_intel", {}) if isinstance(capabilities.get("threat_intel"), dict) else {} + return { + "graylog_mcp": { + "status": mcp.get("status"), + "fetch_mode": mcp.get("fetch_mode"), + "coverage_status": mcp.get("coverage_status"), + "aggregate_events": mcp.get("aggregate_events"), + "raw_events_fetched": mcp.get("raw_events_fetched") or mcp.get("events_fetched"), + "partial_streams": mcp.get("partial_streams"), + "skipped_streams": mcp.get("skipped_streams"), + "coverage_warning": mcp.get("coverage_warning"), + "error": mcp.get("error"), + }, + "profile_advisor": { + "status": profile_advisor.get("status"), + "error": profile_advisor.get("error"), + "candidates": profile_advisor.get("candidates"), + }, + "threat_intel": { + "status": threat_intel.get("status"), + "provider": threat_intel.get("provider"), + }, + } + + +def _compact_profile_suggestion_for_dashboard(item: object) -> dict[str, object]: + if not isinstance(item, dict): + return {} + profile = item.get("profile", {}) if isinstance(item.get("profile"), dict) else {} + advisor = item.get("profile_advisor", {}) if isinstance(item.get("profile_advisor"), dict) else {} + return { + "stream_name": item.get("stream_name"), + "detected_log_type": item.get("detected_log_type"), + "confidence": item.get("confidence"), + "profile_exists": item.get("profile_exists"), + "advisor_status": advisor.get("status"), + "entity_fields": profile.get("entity_fields") or ([profile.get("entity_field")] if profile.get("entity_field") else []), + "timestamp_field": profile.get("timestamp_field"), + "categorical_fields": profile.get("categorical_fields", [])[:8], + "numeric_fields": profile.get("numeric_fields", [])[:5], + "relationship_fields": profile.get("relationship_fields", [])[:5], + } + + +def compact_dashboard_analysis(analysis: dict[str, object]) -> dict[str, object]: + return { "summary": analysis.get("summary", {}), "anomaly_summary": analysis.get("anomaly_summary", {}), + "baseline": analysis.get("baseline", {}), "top_anomalies": analysis.get("anomalies", [])[:5], "top_recommendations": analysis.get("recommendations", [])[:5], "block_candidates": analysis.get("block_candidates", [])[:5], "policy_findings": analysis.get("policy_findings", [])[:5], - "event_context": analysis.get("event_context", {}), + "event_context": _compact_event_context(analysis.get("event_context", {})), "diagnostics": analysis.get("diagnostics", {}), - "capabilities": analysis.get("capabilities", {}), - "cross_source_correlations": analysis.get("cross_source_correlations", [])[:20], - "incidents": analysis.get("incidents", [])[:10], - "profile_suggestions": analysis.get("profile_suggestions", [])[:10], - "field_deviations": analysis.get("field_deviations", {}), - "feedback": analysis.get("feedback", []), + "capabilities": _compact_mcp_capabilities(analysis.get("capabilities", {})), + "cross_source_correlations": analysis.get("cross_source_correlations", [])[:10], + "incidents": analysis.get("incidents", [])[:8], + "profile_suggestions": [ + _compact_profile_suggestion_for_dashboard(item) + for item in analysis.get("profile_suggestions", [])[:8] + ], + "field_deviations": _compact_field_deviations(analysis.get("field_deviations", {})), + "feedback": analysis.get("feedback", [])[:10] if isinstance(analysis.get("feedback"), list) else [], } + + +def ollama_dashboard_assessment(analysis: dict[str, object], model: str | None = None, timeout: int | None = None) -> str: + compact = compact_dashboard_analysis(analysis) return ollama_summary( [], [], diff --git a/tests/test_llm.py b/tests/test_llm.py index 0f40c0b..d2ca54c 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -1,6 +1,6 @@ import unittest -from fgai.llm import _json_object_from_text +from fgai.llm import _json_object_from_text, compact_dashboard_analysis class LlmTests(unittest.TestCase): @@ -18,6 +18,45 @@ class LlmTests(unittest.TestCase): self.assertEqual(payload["profiles"][0]["stream_id"], "firewall") + def test_dashboard_compaction_removes_heavy_evidence_payloads(self): + status = { + "summary": {"total": 1000}, + "capabilities": { + "graylog_mcp": { + "status": "connected", + "streams": [{"stream_id": f"stream-{index}", "events": index} for index in range(100)], + "aggregate_events": 1000, + "raw_events_fetched": 50, + } + }, + "event_context": { + "source_profiles": [{"entity": f"10.0.0.{index}"} for index in range(20)], + "related_activity": [{"entity": f"10.0.0.{index}"} for index in range(40)], + }, + "field_deviations": { + "alice": [ + { + "stream_name": "Windows", + "detector": "new_relationship", + "field": "relationship:username->srcip", + "score": 90, + "reason": "new srcip value", + "sample_events": [{"message": "very large raw event", "graylog_query": "username:alice"}], + } + ] + }, + } + + compact = compact_dashboard_analysis(status) + + self.assertEqual(compact["capabilities"]["graylog_mcp"]["status"], "connected") + self.assertNotIn("streams", compact["capabilities"]["graylog_mcp"]) + self.assertEqual(len(compact["event_context"]["source_profiles"]), 10) + self.assertEqual(len(compact["event_context"]["related_activity"]), 20) + self.assertEqual(compact["field_deviations"][0]["entity"], "alice") + self.assertNotIn("sample_events", compact["field_deviations"][0]) + self.assertNotIn("graylog_query", str(compact)) + if __name__ == "__main__": unittest.main()