from __future__ import annotations import logging from concurrent.futures import ThreadPoolExecutor, as_completed from dataclasses import dataclass from typing import Any from vuln_scan.core.models import SecurityVulnerability from vuln_scan.github.base_client import GitHubBaseClient logger = logging.getLogger(__name__) _DEFAULT_BATCH_SIZE = 10 @dataclass class GitHubAdvisoryClient(GitHubBaseClient): max_workers: int = 4 def get_advisories( self, ecosystem: str, package_name: str, *, max_pages: int = 10, page_size: int = 100, ) -> list[SecurityVulnerability]: """ Fetch all advisories for a package using GraphQL with pagination. """ query = """ query($ecosystem: SecurityAdvisoryEcosystem!, $package: String!, $cursor: String, $page_size: Int!) { securityVulnerabilities( first: $page_size, after: $cursor, ecosystem: $ecosystem, package: $package ) { nodes { severity vulnerableVersionRange package { name ecosystem } firstPatchedVersion { identifier } advisory { ghsaId publishedAt } } pageInfo { hasNextPage endCursor } } } """ advisories: list[SecurityVulnerability] = [] cursor: str | None = None page = 0 while True: page += 1 if page > max_pages: logger.warning( "Reached max_pages=%d for %s/%s, stopping pagination", max_pages, ecosystem, package_name, ) break logger.debug( "Fetching advisories page=%d ecosystem=%s package=%s cursor=%s", page, ecosystem, package_name, cursor, ) data = self._graphql_with_retry( query, { "ecosystem": ecosystem, "package": package_name, "cursor": cursor, "page_size": page_size, }, ) nodes, page_info = self._extract_security_vulnerabilities_page(data) logger.debug( "Fetched %d advisories (page %d)", len(nodes), page, ) advisories.extend(self._parse_advisory_nodes(nodes, ecosystem, package_name)) if not page_info.get("hasNextPage"): logger.debug( "No more pages for %s/%s", ecosystem, package_name, ) break cursor = page_info.get("endCursor") if not cursor: logger.warning("Missing endCursor despite hasNextPage=True, stopping to avoid loop") break logger.debug( "Total advisories fetched for %s/%s: %d", ecosystem, package_name, len(advisories), ) return advisories def get_advisories_batch( self, packages: list[tuple[str, str]], *, batch_size: int = _DEFAULT_BATCH_SIZE, page_size: int = 100, ) -> dict[tuple[str, str], list[SecurityVulnerability]]: """ Fetch advisories for many packages using GraphQL aliasing. Packages whose first page indicates additional pages are re-fetched individually via `get_advisories()` to ensure complete results. """ if not packages: return {} results: dict[tuple[str, str], list[SecurityVulnerability]] = {pkg: [] for pkg in packages} chunks = [packages[i : i + batch_size] for i in range(0, len(packages), batch_size)] max_workers = max(1, min(self.max_workers, len(chunks))) logger.info( "Fetching advisories: %d packages in %d chunks with %d workers", len(packages), len(chunks), max_workers, ) total_cost = 0 last_rate_limit: dict[str, Any] | None = None with ThreadPoolExecutor(max_workers=max_workers) as executor: future_to_chunk = { executor.submit(self._fetch_chunk, chunk, page_size=page_size): chunk for chunk in chunks } for future in as_completed(future_to_chunk): chunk = future_to_chunk[future] try: chunk_results, rate_limit = future.result() results.update(chunk_results) if rate_limit: total_cost += int(rate_limit.get("cost", 0) or 0) last_rate_limit = rate_limit except Exception: logger.exception("Batch chunk failed for packages: %s", chunk) if last_rate_limit: logger.info( "GitHub GraphQL batch rate limit summary: total_cost=%d remaining=%s limit=%s resetAt=%s", total_cost, last_rate_limit.get("remaining"), last_rate_limit.get("limit"), last_rate_limit.get("resetAt"), ) return results def _fetch_chunk( self, chunk: list[tuple[str, str]], *, page_size: int, ) -> tuple[dict[tuple[str, str], list[SecurityVulnerability]], dict[str, Any] | None]: """ Fetch one chunk of packages in a single GraphQL request. Returns: (chunk_results, rate_limit_info) """ query, variables = self._build_batch_query(chunk, page_size=page_size) logger.debug("Batch-fetching %d packages in one GraphQL request", len(chunk)) data = self._graphql_with_retry(query, variables) rate_limit = data.get("rateLimit") if rate_limit: logger.debug( "GitHub GraphQL batch rate limit: cost=%s remaining=%s limit=%s resetAt=%s", rate_limit.get("cost"), rate_limit.get("remaining"), rate_limit.get("limit"), rate_limit.get("resetAt"), ) chunk_results: dict[tuple[str, str], list[SecurityVulnerability]] = {} needs_pagination: list[tuple[str, str]] = [] for i, (ecosystem, package_name) in enumerate(chunk): key = (ecosystem, package_name) section = data.get(f"result_{i}", {}) or {} nodes = section.get("nodes", []) or [] page_info = section.get("pageInfo", {}) or {} chunk_results[key] = self._parse_advisory_nodes( nodes, ecosystem, package_name, ) if page_info.get("hasNextPage"): logger.debug( "%s/%s has more pages; fetching complete advisory history individually", ecosystem, package_name, ) needs_pagination.append(key) for ecosystem, package_name in needs_pagination: logger.info( "Paginating advisories for %s/%s (exceeded first page)", ecosystem, package_name, ) chunk_results[(ecosystem, package_name)] = self.get_advisories( ecosystem, package_name, page_size=page_size, ) return chunk_results, rate_limit def _build_batch_query( self, chunk: list[tuple[str, str]], *, page_size: int, ) -> tuple[str, dict[str, object]]: fragment = """ nodes { severity vulnerableVersionRange package { name ecosystem } firstPatchedVersion { identifier } advisory { ghsaId publishedAt } } pageInfo { hasNextPage endCursor } """ variables: dict[str, object] = {"page_size": page_size} alias_lines: list[str] = [] var_decls = ["$page_size: Int!"] for i, (ecosystem, package_name) in enumerate(chunk): eco_var = f"eco_{i}" pkg_var = f"pkg_{i}" variables[eco_var] = ecosystem variables[pkg_var] = package_name var_decls.append(f"$eco_{i}: SecurityAdvisoryEcosystem!") var_decls.append(f"$pkg_{i}: String!") alias_lines.append( f" result_{i}: securityVulnerabilities(" f"first: $page_size, " f"ecosystem: ${eco_var}, " f"package: ${pkg_var}" f") {{ {fragment} }}" ) query = ( "query(" + ", ".join(var_decls) + ") {\n" + "\n".join(alias_lines) + "\n rateLimit { cost remaining limit resetAt }\n" + "}" ) return query, variables def _extract_security_vulnerabilities_page( self, data: dict[str, Any], ) -> tuple[list[dict[str, Any]], dict[str, Any]]: advisories_data = data.get("securityVulnerabilities", {}) or {} nodes = advisories_data.get("nodes", []) or [] page_info = advisories_data.get("pageInfo", {}) or {} return nodes, page_info def _parse_advisory_nodes( self, nodes: list[dict[str, Any]], ecosystem: str, package_name: str, ) -> list[SecurityVulnerability]: advisories: list[SecurityVulnerability] = [] for node in nodes: try: advisories.append(SecurityVulnerability.from_advisory_node(node)) except Exception: logger.exception( "Failed to parse advisory for %s/%s", ecosystem, package_name, ) return advisories