from __future__ import annotations import logging from dataclasses import dataclass from typing import Any from vuln_scan.core.models import SecurityVulnerability from vuln_scan.github.repo_client import GitHubRepoClient logger = logging.getLogger(__name__) @dataclass class GithubRepoAlertsClient(GitHubRepoClient): def get_open_alerts( self, manifest_paths: set[str] | None = None, *, max_pages: int = 10, page_size: int = 100, ) -> list[SecurityVulnerability]: alerts: list[SecurityVulnerability] = [] cursor: str | None = None for page in range(1, max_pages + 1): logger.debug( "Fetching dependabot alerts page=%d repo=%s/%s cursor=%s", page, self.owner, self.repo, cursor, ) data = self._fetch_alerts_page(cursor, page_size) nodes, page_info = self._extract_alerts(data) logger.debug("Fetched %d alerts (page %d)", len(nodes), page) alerts.extend(self._parse_nodes(nodes, manifest_paths)) if not self._has_next_page(page_info): break cursor = page_info.get("endCursor") if not cursor: logger.warning("Missing endCursor despite hasNextPage=True") break logger.info( "Total dependabot alerts fetched for %s/%s: %d", self.owner, self.repo, len(alerts), ) return alerts def _fetch_alerts_page(self, cursor: str | None, page_size: int) -> dict[str, Any]: return self._graphql_with_retry( self._alerts_query(), { "owner": self.owner, "name": self.repo, "first": page_size, "after": cursor, }, ) def _extract_alerts( self, data: dict[str, Any], ) -> tuple[list[dict[str, Any]], dict[str, Any]]: repo_data = data.get("repository", {}) or {} alerts_data = repo_data.get("vulnerabilityAlerts", {}) or {} return ( alerts_data.get("nodes", []), alerts_data.get("pageInfo", {}) or {}, ) def _parse_nodes( self, nodes: list[dict[str, Any]], manifest_paths: set[str] | None, ) -> list[SecurityVulnerability]: results: list[SecurityVulnerability] = [] for node in nodes: if not self._should_include(node, manifest_paths): continue try: results.append(SecurityVulnerability.from_dependabot_node(node)) except Exception: logger.exception( "Failed to parse dependabot alert: %s", node.get("number"), ) return results def _should_include( self, node: dict[str, Any], manifest_paths: set[str] | None, ) -> bool: if not manifest_paths: return True return node.get("vulnerableManifestPath") in manifest_paths def _has_next_page(self, page_info: dict[str, Any]) -> bool: return bool(page_info.get("hasNextPage")) @staticmethod def _alerts_query() -> str: return """ query($owner: String!, $name: String!, $first: Int!, $after: String) { repository(owner: $owner, name: $name) { vulnerabilityAlerts(first: $first, after: $after, states:[OPEN]) { nodes { number id createdAt state vulnerableManifestPath securityVulnerability { severity vulnerableVersionRange firstPatchedVersion { identifier } package { name ecosystem } advisory { ghsaId } } } pageInfo { hasNextPage endCursor } } } } """