Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 19 additions & 2 deletions src/github_agent_bridge/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
52 changes: 51 additions & 1 deletion tests/test_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading