diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/srht_contrib/jobs/poller.py | 12 | ||||
| -rw-r--r-- | src/srht_contrib/main.py | 4 | ||||
| -rw-r--r-- | src/srht_contrib/models.py | 10 | ||||
| -rw-r--r-- | src/srht_contrib/services/git.py | 63 | ||||
| -rw-r--r-- | src/srht_contrib/services/srht_client.py | 18 |
5 files changed, 83 insertions, 24 deletions
diff --git a/src/srht_contrib/jobs/poller.py b/src/srht_contrib/jobs/poller.py index 94a7d96..b2fe70f 100644 --- a/src/srht_contrib/jobs/poller.py +++ b/src/srht_contrib/jobs/poller.py @@ -91,7 +91,6 @@ class PollerService: inserted = 0 inserted += self._poll_service(db, actor, self.todo_service.service_name, self.todo_service.fetch_recent_events) self._sync_tracked_repositories(db, actor) - git_repositories = self._tracked_repositories_for_actor(db, actor) inserted += self._poll_service( db, actor, @@ -99,7 +98,7 @@ class PollerService: lambda actor, since: self.git_service.fetch_recent_events( actor=actor, since=since, - repositories=git_repositories, + db=db, ), ) return inserted @@ -172,15 +171,6 @@ class PollerService: ) db.flush() - def _tracked_repositories_for_actor(self, db: Session, actor: str) -> list[str]: - rows = db.scalars( - select(TrackedRepository.repo_name) - .where(TrackedRepository.service == self.git_service.service_name) - .where(TrackedRepository.actor == actor) - .order_by(TrackedRepository.repo_name) - ).all() - return list(rows) - def _update_tracked_actor_poll_state(self, db: Session, actor: str, status: str, error: str | None) -> None: tracked_actor = self.track_actor_request(db, actor, update_last_requested=False) tracked_actor.last_poll_status = status diff --git a/src/srht_contrib/main.py b/src/srht_contrib/main.py index 3737569..d8284d6 100644 --- a/src/srht_contrib/main.py +++ b/src/srht_contrib/main.py @@ -26,8 +26,8 @@ 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) + todo_client = SourceHutGraphQLClient(settings.todo_srht_endpoint, settings.srht_token, request_delay=settings.srht_request_delay_seconds) + git_client = SourceHutGraphQLClient(settings.git_srht_endpoint, settings.srht_token, request_delay=settings.srht_request_delay_seconds) todo_service = TodoIngestionService(todo_client, settings) git_service = GitIngestionService(git_client, settings) return PollerService(todo_service=todo_service, git_service=git_service, settings=settings) diff --git a/src/srht_contrib/models.py b/src/srht_contrib/models.py index 475bd22..e0725ff 100644 --- a/src/srht_contrib/models.py +++ b/src/srht_contrib/models.py @@ -81,6 +81,16 @@ class TrackedActor(Base): last_backfill_error: Mapped[str | None] = mapped_column(Text, nullable=True) +class DiscoveredRepository(Base): + __tablename__ = "discovered_repositories" + __table_args__ = (UniqueConstraint("actor", "name", name="uq_discovered_repository_actor_name"),) + + id: Mapped[int] = mapped_column(Integer, primary_key=True) + actor: Mapped[str] = mapped_column(String(255), nullable=False) + name: Mapped[str] = mapped_column(String(255), nullable=False) + discovered_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + + class ServiceBackfillState(Base): __tablename__ = "service_backfill_states" __table_args__ = (UniqueConstraint("actor", "service", "scope", name="uq_service_backfill_state_actor_service_scope"),) diff --git a/src/srht_contrib/services/git.py b/src/srht_contrib/services/git.py index 7d5360b..26b2957 100644 --- a/src/srht_contrib/services/git.py +++ b/src/srht_contrib/services/git.py @@ -6,7 +6,11 @@ from datetime import UTC, datetime, timedelta import logging from typing import Any +from sqlalchemy import delete, select +from sqlalchemy.orm import Session + from srht_contrib.config import Settings +from srht_contrib.models import DiscoveredRepository, TrackedRepository from srht_contrib.schemas import NormalizedEvent from srht_contrib.services.srht_client import SourceHutClientError, SourceHutGraphQLClient from srht_contrib.services.types import BackfillBatchResult @@ -88,9 +92,10 @@ class GitIngestionService: actor: str, since: datetime | None = None, repositories: list[str] | None = None, + db: Session | None = None, ) -> GitPollResult: since_dt = ensure_utc(since or (datetime.now(tz=UTC) - timedelta(days=30))) - discovered_repositories = repositories or self._repositories_for_actor(actor) + discovered_repositories = repositories or self._repositories_for_actor(actor, db) if not discovered_repositories: logger.info("git poll skipped for actor=%s because no repositories were discovered", actor) return GitPollResult(events=[], cursor=datetime.now(tz=UTC).isoformat()) @@ -108,18 +113,44 @@ class GitIngestionService: logger.info("git poll complete for actor=%s normalized_events=%s", actor, len(events)) return GitPollResult(events=events, cursor=datetime.now(tz=UTC).isoformat()) - def _repositories_for_actor(self, actor: str) -> list[str]: + def _repositories_for_actor(self, actor: str, db: Session | None = None) -> list[str]: configured = { self._canonical_repository_name(actor, repository) for repository in self.settings.git_tracked_repositories } - discovered = set(self._discover_owned_repositories(actor)) - repositories = sorted(configured | discovered) + tracked = set() + discovered = set() + now = datetime.now(tz=UTC) + + if db is not None: + tracked = set( + db.scalars( + select(TrackedRepository.repo_name) + .where(TrackedRepository.service == self.service_name) + .where(TrackedRepository.actor == actor) + ).all() + ) + cutoff = now - timedelta(seconds=self.settings.git_repo_discovery_ttl_seconds) + discovered = set( + db.scalars( + select(DiscoveredRepository.name) + .where(DiscoveredRepository.actor == actor) + .where(DiscoveredRepository.discovered_at > cutoff) + ).all() + ) + + if not discovered: + discovered = set(self._discover_owned_repositories(actor)) + if db is not None: + self._refresh_discovered_repositories(db, actor, discovered, discovered_at=now) + + repositories = sorted(configured | tracked | discovered) logger.info( - "git repositories selected for actor=%s count=%s configured=%s discovered=%s", + "git repositories selected for actor=%s count=%s configured=%s tracked=%s discovered=%s", actor, len(repositories), len(configured), + len(tracked), len(discovered), ) return repositories @@ -268,7 +299,8 @@ class GitIngestionService: repository_owner = ((repository.get("owner") or {}).get("canonicalName") or actor).strip() if not name or not repository_owner: continue - repositories.append(f"{repository_owner}/{name}") + repo_name = f"{repository_owner}/{name}" + repositories.append(repo_name) if not cursor: break @@ -276,6 +308,25 @@ class GitIngestionService: return repositories @staticmethod + def _refresh_discovered_repositories( + db: Session, + actor: str, + repositories: set[str], + *, + discovered_at: datetime, + ) -> None: + db.execute(delete(DiscoveredRepository).where(DiscoveredRepository.actor == actor)) + for repository in sorted(repositories): + db.add( + DiscoveredRepository( + actor=actor, + name=repository, + discovered_at=discovered_at, + ) + ) + db.flush() + + @staticmethod def _canonical_repository_name(default_actor: str, repository: str) -> str: owner, repo_name = GitIngestionService._split_repository(default_actor, repository) return f"~{owner}/{repo_name}" diff --git a/src/srht_contrib/services/srht_client.py b/src/srht_contrib/services/srht_client.py index ee09bff..1118f0d 100644 --- a/src/srht_contrib/services/srht_client.py +++ b/src/srht_contrib/services/srht_client.py @@ -1,6 +1,7 @@ from __future__ import annotations import logging +import time from typing import Any import httpx @@ -27,11 +28,13 @@ class SourceHutGraphQLClient: *, timeout: float = 15.0, max_retries: int = 2, + request_delay: float = 0.5, transport: httpx.BaseTransport | None = None, ) -> None: self.endpoint = endpoint self.timeout = timeout self.max_retries = max_retries + self.request_delay = request_delay headers = { "Authorization": f"Bearer {token}", "Content-Type": "application/json", @@ -42,7 +45,12 @@ class SourceHutGraphQLClient: payload = {"query": query, "variables": variables or {}} attempts = self.max_retries + 1 - for attempt in range(1, attempts + 1): + for attempt in range(attempts): + if attempt > 0: + time.sleep(2 ** (attempt - 1)) + elif self.request_delay > 0: + time.sleep(self.request_delay) + try: response = self._client.post(self.endpoint, json=payload) response.raise_for_status() @@ -51,21 +59,21 @@ class SourceHutGraphQLClient: logger.warning( "SourceHut HTTP failure from %s on attempt %s/%s: status=%s", self.endpoint, - attempt, + attempt + 1, attempts, exc.response.status_code, ) - if exc.response.status_code >= 500 and attempt < attempts: + if exc.response.status_code >= 500 and attempt < attempts - 1: continue raise SourceHutClientError(f"HTTP error from SourceHut: {exc.response.status_code}") from exc except httpx.HTTPError as exc: logger.warning( "SourceHut network failure from %s on attempt %s/%s", self.endpoint, - attempt, + attempt + 1, attempts, ) - if attempt < attempts: + if attempt < attempts - 1: continue raise SourceHutClientError("Network error while contacting SourceHut") from exc |
