diff --git a/src/fgai/graylog_mcp.py b/src/fgai/graylog_mcp.py index d4881b9..16fee1c 100644 --- a/src/fgai/graylog_mcp.py +++ b/src/fgai/graylog_mcp.py @@ -67,6 +67,7 @@ class GraylogMcpClient: return f"Basic {credentials}" def probe(self) -> dict[str, object]: + self.session_id = None initialized = self._call( "initialize", {"protocolVersion": "2025-06-18", "capabilities": {}, "clientInfo": {"name": "fgAI", "version": "0.1"}}, @@ -84,4 +85,14 @@ class GraylogMcpClient: 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}) + try: + return self._call("tools/call", {"name": name, "arguments": arguments}) + except RuntimeError as exc: + if not self._should_reconnect(str(exc)): + raise + self.probe() + return self._call("tools/call", {"name": name, "arguments": arguments}) + + @staticmethod + def _should_reconnect(error_text: str) -> bool: + return any(token in error_text for token in ("http_400", "http_404", "http_409", "connection_error")) diff --git a/src/fgai/monitor.py b/src/fgai/monitor.py index 004ab11..b727e47 100644 --- a/src/fgai/monitor.py +++ b/src/fgai/monitor.py @@ -310,6 +310,7 @@ def build_status( summary["aggregate_backed"] = True status = { "status_schema": 2, + "stale": False, "generated_at": int(time.time()), "log_path": log_path, "policy_path": policy_path, @@ -341,6 +342,7 @@ def build_status( "cross_source_correlations": correlations, "incidents": incidents, "data_quality": assess_data_quality(events, mcp_status), + "status_cache": {"served_from_cache": False, "reason": ""}, "anomalies": [ { "subject": finding.subject, diff --git a/tests/test_graylog_mcp.py b/tests/test_graylog_mcp.py index 47097ce..c04b60d 100644 --- a/tests/test_graylog_mcp.py +++ b/tests/test_graylog_mcp.py @@ -1,5 +1,6 @@ import json import unittest +from urllib.error import HTTPError from unittest.mock import patch from fgai.graylog_mcp import GraylogMcpClient @@ -38,6 +39,23 @@ class GraylogMcpTests(unittest.TestCase): self.assertEqual(status["status"], "connected") self.assertIn("search_messages", status["tools"]) + def test_call_tool_reinitializes_after_stale_session_error(self): + responses = iter([ + HTTPError("http://graylog/api/mcp", 404, "stale session", {}, None), + _Response({"result": {"serverInfo": {"version": "7.1.6"}}}, session="session-2"), + _Response({}, session="session-2"), + _Response({"result": {"tools": [{"name": "search_messages"}]}}, session="session-2"), + _Response({"result": {"content": []}}, session="session-2"), + ]) + client = GraylogMcpClient("http://graylog/api/mcp", "raw-token") + client.session_id = "old-session" + + with patch("fgai.graylog_mcp.request.urlopen", side_effect=responses): + result = client.call_tool("search_messages", {"query": "*"}) + + self.assertEqual(result["result"]["content"], []) + self.assertEqual(client.session_id, "session-2") + if __name__ == "__main__": unittest.main() diff --git a/tests/test_monitor.py b/tests/test_monitor.py index ca4f6ab..7cb9a97 100644 --- a/tests/test_monitor.py +++ b/tests/test_monitor.py @@ -27,6 +27,8 @@ class MonitorTests(unittest.TestCase): self.assertEqual(status["summary"]["total"], 3) self.assertEqual(status["status_schema"], 2) + self.assertFalse(status["stale"]) + self.assertFalse(status["status_cache"]["served_from_cache"]) self.assertGreaterEqual(len(status["anomalies"]), 1) def test_write_status_creates_parent_directory(self):