diff --git a/src/github_agent_bridge/backend.py b/src/github_agent_bridge/backend.py index c0aebf0..b7960bd 100644 --- a/src/github_agent_bridge/backend.py +++ b/src/github_agent_bridge/backend.py @@ -508,13 +508,31 @@ async def _sleep_or_shutdown(shutdown_event: asyncio.Event | None, sleep_seconds return True +def _is_sqlite_contention_error(exc: sqlite3.OperationalError) -> bool: + error_code = getattr(exc, "sqlite_errorcode", None) + if isinstance(error_code, int) and error_code & 0xFF in { + sqlite3.SQLITE_BUSY, + sqlite3.SQLITE_LOCKED, + }: + return True + message = str(exc).lower() + return "database" in message and "locked" in message + + async def _session_stream_events(db: str | Path, job_id: int, *, after_id: int | None = None, sleep_seconds: float = 2.0, shutdown_event: asyncio.Event | None = None): queries = DashboardQueries(db) last_id = after_id or 0 sent_transcript_keys: set[str] = set() while shutdown_event is None or not shutdown_event.is_set(): emitted = False - events = queries.job_session_events(job_id, after_id=last_id, limit=100) + try: + events = queries.job_session_events(job_id, after_id=last_id, limit=100) + transcript = queries.job_session_transcript(job_id, limit=500) + except sqlite3.OperationalError as exc: + if not _is_sqlite_contention_error(exc): + raise + events = [] + transcript = [] for event in events: if shutdown_event is not None and shutdown_event.is_set(): return @@ -527,7 +545,6 @@ async def _session_stream_events(db: str | Path, job_id: int, *, after_id: int | if key not in sent_transcript_keys: sent_transcript_keys.add(key) yield _sse_event("transcript_entry", {"job_id": job_id, "entry": entry}) - transcript = queries.job_session_transcript(job_id, limit=500) for entry in transcript: if shutdown_event is not None and shutdown_event.is_set(): return diff --git a/tests/test_backend.py b/tests/test_backend.py index 958ef23..41603d1 100644 --- a/tests/test_backend.py +++ b/tests/test_backend.py @@ -13,7 +13,7 @@ from github_agent_bridge import __version__ from github_agent_bridge import feedback from github_agent_bridge.backend import DashboardConfig, _encode_session, _is_admin, _is_allowed, _journal_stream_events, _session_stream_events, _sign, create_app -from github_agent_bridge.dashboard_data import JOB_LIST_ORDER_SQL, get_job_detail, job_session, job_session_events, job_session_transcript, jobs_select_sql, list_all_job_actor_logins, list_job_actors, list_jobs, metrics_summary +from github_agent_bridge.dashboard_data import DashboardQueries, JOB_LIST_ORDER_SQL, get_job_detail, job_session, job_session_events, job_session_transcript, jobs_select_sql, list_all_job_actor_logins, list_job_actors, list_jobs, metrics_summary from github_agent_bridge.monitor import MonitorReport from github_agent_bridge.models import GitHubContext, Notification from github_agent_bridge.mcp import create_token @@ -1086,6 +1086,56 @@ async def first_chunks(): assert "live trajectory output" in body +def test_dashboard_sse_retries_sqlite_contention_without_closing_stream(tmp_path, monkeypatch): + db = tmp_path / "bridge.sqlite3" + q = JobQueue(db) + job, _ = q.enqueue(notif(), Policy(trusted_orgs=["gisce"])) + original_events = DashboardQueries.job_session_events + attempts = {"count": 0} + + def fail_once_then_read(*args, **kwargs): + attempts["count"] += 1 + if attempts["count"] == 1: + raise sqlite3.OperationalError("database is locked") + return original_events(*args, **kwargs) + + monkeypatch.setattr(DashboardQueries, "job_session_events", fail_once_then_read) + + async def first_chunks(): + stream = _session_stream_events(db, job.id, sleep_seconds=0) + try: + return await anext(stream), await anext(stream) + finally: + await stream.aclose() + + heartbeat, recovered = asyncio.run(first_chunks()) + + assert attempts["count"] == 2 + assert "event: session_heartbeat" in heartbeat + assert "event: session_heartbeat" in recovered + + +def test_dashboard_sse_does_not_hide_non_contention_database_errors(tmp_path, monkeypatch): + db = tmp_path / "bridge.sqlite3" + q = JobQueue(db) + job, _ = q.enqueue(notif(), Policy(trusted_orgs=["gisce"])) + + def fail_read(*args, **kwargs): + raise sqlite3.OperationalError("malformed database schema") + + monkeypatch.setattr(DashboardQueries, "job_session_events", fail_read) + + async def read_chunk(): + stream = _session_stream_events(db, job.id, sleep_seconds=0) + try: + return await anext(stream) + finally: + await stream.aclose() + + with pytest.raises(sqlite3.OperationalError, match="malformed database schema"): + asyncio.run(read_chunk()) + + def test_dashboard_sse_stream_exits_when_shutdown_is_signaled(tmp_path): db = tmp_path / "bridge.sqlite3" q = JobQueue(db)