diff options
| -rw-r--r-- | src/srht_contrib/main.py | 12 | ||||
| -rw-r--r-- | tests/test_polling_api.py | 31 |
2 files changed, 43 insertions, 0 deletions
diff --git a/src/srht_contrib/main.py b/src/srht_contrib/main.py index 9047bcc..c6038cb 100644 --- a/src/srht_contrib/main.py +++ b/src/srht_contrib/main.py @@ -2,6 +2,7 @@ from __future__ import annotations from collections.abc import AsyncIterator from contextlib import asynccontextmanager +import logging from apscheduler.schedulers.background import BackgroundScheduler from fastapi import FastAPI @@ -21,6 +22,9 @@ from srht_contrib.services.todo import TodoIngestionService from srht_contrib.utils.identity import ActorIdentityResolver +logger = logging.getLogger(__name__) + + def build_poller(settings: Settings) -> PollerService: todo_client = SourceHutGraphQLClient(settings.todo_srht_endpoint, settings.srht_token) git_client = SourceHutGraphQLClient(settings.git_srht_endpoint, settings.srht_token) @@ -68,6 +72,7 @@ def create_app( replace_existing=True, ) scheduler.start() + _run_startup_poll(app) app.state.scheduler = scheduler try: yield @@ -96,4 +101,11 @@ def _scheduled_poll(app: FastAPI) -> None: db.close() +def _run_startup_poll(app: FastAPI) -> None: + try: + _scheduled_poll(app) + except Exception: + logger.exception("Initial scheduled poll failed during application startup") + + app = create_app() diff --git a/tests/test_polling_api.py b/tests/test_polling_api.py index 75270e0..b752524 100644 --- a/tests/test_polling_api.py +++ b/tests/test_polling_api.py @@ -19,6 +19,7 @@ class InsertingPoller: service = type("Service", (), {"client": _Closable()})() self.todo_service = service self.git_service = service + self.tracked_poll_calls: list[str] = [] def track_actor_request(self, db, actor: str, *, update_last_requested: bool = True): tracked_actor = db.scalar(select(TrackedActor).where(TrackedActor.actor == actor)) @@ -47,6 +48,12 @@ class InsertingPoller: db.commit() return 1 + def poll_tracked_actors(self, db, default_actor: str | None = None) -> dict[str, int]: + if default_actor is not None: + self.tracked_poll_calls.append(default_actor) + return {default_actor: self.poll_all(db, default_actor)} + return {} + class FailingPoller: def __init__(self) -> None: @@ -65,6 +72,9 @@ class FailingPoller: def poll_all(self, db, actor: str) -> int: raise SourceHutClientError("boom") + def poll_tracked_actors(self, db, default_actor: str | None = None) -> dict[str, int]: + raise SourceHutClientError("boom") + def test_manual_poll_uses_same_database_session(settings: Settings, db_engine, session_factory) -> None: app = create_app(settings, engine=db_engine, session_factory=session_factory, poller=InsertingPoller()) @@ -89,3 +99,24 @@ def test_manual_poll_maps_sourcehut_failures_to_502(settings: Settings, db_engin assert response.status_code == 502 assert "SourceHut polling failed" in response.json()["detail"] + + +def test_scheduler_runs_initial_poll_on_startup(settings: Settings, db_engine, session_factory) -> None: + scheduler_settings = settings.model_copy(update={"enable_scheduler": True}) + poller = InsertingPoller() + app = create_app(scheduler_settings, engine=db_engine, session_factory=session_factory, poller=poller) + + with TestClient(app): + pass + + assert poller.tracked_poll_calls == ["~ccleberg"] + + +def test_startup_poll_failure_does_not_block_app_start(settings: Settings, db_engine, session_factory) -> None: + scheduler_settings = settings.model_copy(update={"enable_scheduler": True}) + app = create_app(scheduler_settings, engine=db_engine, session_factory=session_factory, poller=FailingPoller()) + + with TestClient(app) as client: + response = client.get("/health") + + assert response.status_code == 200 |
