pull-request-fixer 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,70 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # SPDX-FileCopyrightText: 2025 The Linux Foundation
3
+
4
+ """Data models for pr-title-fixer."""
5
+
6
+ from __future__ import annotations
7
+
8
+ from dataclasses import dataclass, field
9
+ from enum import Enum
10
+ from typing import TYPE_CHECKING
11
+
12
+ if TYPE_CHECKING:
13
+ from pathlib import Path
14
+
15
+
16
+ class OutputFormat(str, Enum):
17
+ """Output format options."""
18
+
19
+ TEXT = "text"
20
+ JSON = "json"
21
+ TABLE = "table"
22
+
23
+
24
+ @dataclass
25
+ class PRInfo:
26
+ """Information about a GitHub pull request."""
27
+
28
+ number: int
29
+ title: str
30
+ repository: str
31
+ url: str
32
+ author: str
33
+ is_draft: bool
34
+ head_ref: str
35
+ head_sha: str
36
+ base_ref: str
37
+ mergeable: str
38
+ merge_state_status: str
39
+
40
+
41
+ @dataclass
42
+ class BlockedPR:
43
+ """A blocked pull request with blocking reasons."""
44
+
45
+ pr_info: PRInfo
46
+ blocking_reasons: list[str]
47
+ has_title_issues: bool = False
48
+
49
+
50
+ @dataclass
51
+ class GitHubScanResult:
52
+ """Results from scanning a GitHub organization."""
53
+
54
+ organization: str
55
+ repositories_scanned: int = 0
56
+ total_prs: int = 0
57
+ blocked_prs: list[BlockedPR] = field(default_factory=list)
58
+ prs_fixed: int = 0
59
+ errors: list[str] = field(default_factory=list)
60
+
61
+
62
+ @dataclass
63
+ class GitHubFixResult:
64
+ """Result of fixing a PR."""
65
+
66
+ pr_info: PRInfo
67
+ success: bool
68
+ message: str
69
+ files_modified: list[Path] = field(default_factory=list)
70
+ error: str | None = None
@@ -0,0 +1,37 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # SPDX-FileCopyrightText: 2025 The Linux Foundation
3
+
4
+ """PR fixer module.
5
+
6
+ Note: This module is kept for backwards compatibility but is not currently used.
7
+ The actual PR fixing logic is implemented directly in cli.py in the process_pr function.
8
+ All PR title and body fixing is handled there.
9
+
10
+ This module previously contained markdown table fixing code from the original
11
+ markdown-table-fixer tool, which has been removed as it's not relevant to this tool's purpose.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import logging
17
+ from typing import TYPE_CHECKING
18
+
19
+ if TYPE_CHECKING:
20
+ from .github_client import GitHubClient
21
+
22
+
23
+ class PRFixer:
24
+ """PR fixer class (currently unused).
25
+
26
+ The actual PR fixing logic is implemented in cli.py.
27
+ This class is kept for backwards compatibility.
28
+ """
29
+
30
+ def __init__(self, client: GitHubClient):
31
+ """Initialize PR fixer.
32
+
33
+ Args:
34
+ client: GitHub API client
35
+ """
36
+ self.client = client
37
+ self.logger = logging.getLogger("pull_request_fixer.pr_fixer")
@@ -0,0 +1,423 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # SPDX-FileCopyrightText: 2025 The Linux Foundation
3
+
4
+ """Scanner for identifying pull requests across GitHub organizations.
5
+
6
+ This module provides functionality for scanning GitHub organizations to find
7
+ pull requests that need processing. It uses GraphQL for efficient querying
8
+ and supports parallel processing with bounded concurrency.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import asyncio
14
+ import logging
15
+ from typing import TYPE_CHECKING, Any
16
+
17
+ if TYPE_CHECKING:
18
+ from collections.abc import AsyncIterator
19
+
20
+ from .github_client import GitHubClient
21
+ from .progress_tracker import ProgressTracker
22
+
23
+ from .graphql_queries import (
24
+ ORG_REPOS_ONLY,
25
+ ORG_REPOS_WITH_PRS,
26
+ REPO_OPEN_PRS_PAGE,
27
+ )
28
+
29
+ # GitHub API tuning defaults - optimized for performance and rate limit compliance
30
+ # These match dependamerge's proven values
31
+ DEFAULT_PRS_PAGE_SIZE = 30 # Pull requests per GraphQL page
32
+ DEFAULT_FILES_PAGE_SIZE = 50 # Files per pull request
33
+ DEFAULT_COMMENTS_PAGE_SIZE = 10 # Comments per pull request
34
+ DEFAULT_CONTEXTS_PAGE_SIZE = 20 # Status contexts per pull request
35
+
36
+
37
+ class PRScanner:
38
+ """Scanner for finding pull requests across GitHub organizations.
39
+
40
+ This scanner uses GraphQL to efficiently query GitHub organizations for
41
+ pull requests. It supports:
42
+ - Parallel processing with bounded concurrency
43
+ - Progress tracking for UI updates
44
+ - Filtering by various criteria (draft status, etc.)
45
+ - Pagination for large result sets
46
+ """
47
+
48
+ def __init__(
49
+ self,
50
+ client: GitHubClient,
51
+ progress_tracker: ProgressTracker | None = None,
52
+ max_repo_tasks: int = 8,
53
+ max_page_tasks: int = 16,
54
+ ):
55
+ """Initialize PR scanner.
56
+
57
+ Args:
58
+ client: GitHub API client
59
+ progress_tracker: Optional progress tracker for UI updates
60
+ max_repo_tasks: Maximum concurrent repository scans
61
+ max_page_tasks: Maximum concurrent page fetches
62
+ """
63
+ self.client = client
64
+ self.progress_tracker = progress_tracker
65
+ self.max_repo_tasks = max_repo_tasks
66
+ self.max_page_tasks = max_page_tasks
67
+ self.repo_semaphore = asyncio.Semaphore(max_repo_tasks)
68
+ self.page_semaphore = asyncio.Semaphore(max_page_tasks)
69
+ self.logger = logging.getLogger("pull_request_fixer.pr_scanner")
70
+
71
+ async def scan_organization(
72
+ self,
73
+ organization: str,
74
+ include_drafts: bool = False,
75
+ ) -> AsyncIterator[tuple[str, str, dict[str, Any]]]:
76
+ """Scan an organization for pull requests.
77
+
78
+ This method yields pull requests as they are discovered, allowing
79
+ for streaming processing without loading all results into memory.
80
+
81
+ Args:
82
+ organization: Organization name to scan
83
+ include_drafts: Whether to include draft PRs
84
+
85
+ Yields:
86
+ Tuples of (owner, repo_name, pr_data) where pr_data is a dict
87
+ containing the PR information from the GitHub API
88
+ """
89
+ self.logger.debug(f"Starting scan of organization: {organization}")
90
+
91
+ # First pass: count repositories for progress tracking
92
+ total_repos = await self._count_org_repositories(organization)
93
+ if self.progress_tracker and total_repos > 0:
94
+ self.progress_tracker.update_total_repositories(total_repos)
95
+ # Start the progress tracker now that we have the count
96
+ self.progress_tracker.start()
97
+
98
+ self.logger.debug(f"Found {total_repos} repositories in {organization}")
99
+
100
+ # If no repos found or error occurred, nothing to scan
101
+ if total_repos == 0:
102
+ self.logger.warning(
103
+ f"No repositories found in organization: {organization}"
104
+ )
105
+ return
106
+
107
+ # Second pass: scan repositories with bounded parallelism
108
+ async for owner, repo_name, pr_data in self._scan_repositories(
109
+ organization, include_drafts
110
+ ):
111
+ yield owner, repo_name, pr_data
112
+
113
+ async def _count_org_repositories(self, organization: str) -> int:
114
+ """Count total repositories in an organization.
115
+
116
+ Args:
117
+ organization: Organization name
118
+
119
+ Returns:
120
+ Total number of repositories
121
+ """
122
+ try:
123
+ result = await self.client.graphql(
124
+ ORG_REPOS_ONLY,
125
+ variables={"org": organization, "reposCursor": None},
126
+ )
127
+
128
+ # Check if organization data exists
129
+ org_data = result.get("organization")
130
+ if org_data is None:
131
+ self.logger.error(
132
+ f"Organization '{organization}' not found or token lacks access.\n"
133
+ f" Please verify:\n"
134
+ f" 1. Organization name is correct (case-sensitive)\n"
135
+ f" 2. Token has 'read:org' scope\n"
136
+ f" 3. Token user has access to the organization\n"
137
+ f" 4. Organization exists on GitHub"
138
+ )
139
+ return 0
140
+
141
+ repos = org_data.get("repositories", {})
142
+ total: int = repos.get("totalCount", 0)
143
+ self.logger.debug(
144
+ f"Organization {organization} has {total} repositories"
145
+ )
146
+ return total
147
+ except Exception as e:
148
+ self.logger.error(f"Error counting repositories: {e}")
149
+ import traceback
150
+
151
+ self.logger.debug(f"Traceback: {traceback.format_exc()}")
152
+ return 0
153
+
154
+ async def _scan_repositories(
155
+ self,
156
+ organization: str,
157
+ include_drafts: bool,
158
+ ) -> AsyncIterator[tuple[str, str, dict[str, Any]]]:
159
+ """Scan repositories in an organization for PRs.
160
+
161
+ Args:
162
+ organization: Organization name
163
+ include_drafts: Whether to include draft PRs
164
+
165
+ Yields:
166
+ Tuples of (owner, repo_name, pr_data)
167
+ """
168
+ # Create a queue for PR results
169
+ pr_queue: asyncio.Queue[tuple[str, str, dict[str, Any]] | None] = (
170
+ asyncio.Queue()
171
+ )
172
+
173
+ async def process_repository(repo_node: dict[str, Any]) -> None:
174
+ """Process a single repository and add PRs to queue."""
175
+ async with self.repo_semaphore:
176
+ repo_full_name = repo_node.get(
177
+ "nameWithOwner", "unknown/unknown"
178
+ )
179
+ if self.progress_tracker:
180
+ self.progress_tracker.start_repository(repo_full_name)
181
+
182
+ try:
183
+ owner, repo_name = self._split_owner_repo(repo_full_name)
184
+ pr_count = 0
185
+
186
+ # Fetch first page of PRs
187
+ (
188
+ first_nodes,
189
+ page_info,
190
+ ) = await self._fetch_repo_prs_first_page(owner, repo_name)
191
+
192
+ # Process first page PRs
193
+ for pr_node in first_nodes:
194
+ if self._should_include_pr(pr_node, include_drafts):
195
+ await pr_queue.put((owner, repo_name, pr_node))
196
+ pr_count += 1
197
+
198
+ # Process additional pages if present
199
+ has_next = page_info.get("hasNextPage", False)
200
+ end_cursor = page_info.get("endCursor")
201
+
202
+ if has_next and end_cursor:
203
+ async for pr_node in self._iter_repo_open_prs_pages(
204
+ owner, repo_name, end_cursor
205
+ ):
206
+ if self._should_include_pr(pr_node, include_drafts):
207
+ await pr_queue.put((owner, repo_name, pr_node))
208
+ pr_count += 1
209
+
210
+ if self.progress_tracker:
211
+ self.progress_tracker.complete_repository(pr_count)
212
+
213
+ self.logger.debug(
214
+ f"Repository {repo_full_name}: found {pr_count} PRs"
215
+ )
216
+
217
+ except Exception as e:
218
+ self.logger.error(
219
+ f"Error scanning repository {repo_full_name}: {e}"
220
+ )
221
+ if self.progress_tracker:
222
+ self.progress_tracker.add_error()
223
+
224
+ async def producer() -> None:
225
+ """Producer coroutine that scans repositories."""
226
+ tasks: list[asyncio.Task[None]] = []
227
+ try:
228
+ async for (
229
+ repo_node
230
+ ) in self._iter_org_repositories_with_open_prs(organization):
231
+ task = asyncio.create_task(process_repository(repo_node))
232
+ tasks.append(task)
233
+
234
+ # Wait for all repository processing to complete
235
+ if tasks:
236
+ await asyncio.gather(*tasks, return_exceptions=True)
237
+
238
+ finally:
239
+ # Signal completion by putting None in queue
240
+ await pr_queue.put(None)
241
+
242
+ # Start producer task
243
+ producer_task = asyncio.create_task(producer())
244
+
245
+ try:
246
+ # Yield PRs as they become available
247
+ while True:
248
+ item = await pr_queue.get()
249
+ if item is None:
250
+ break
251
+ yield item
252
+ finally:
253
+ # Ensure producer completes
254
+ await producer_task
255
+
256
+ async def _iter_org_repositories_with_open_prs(
257
+ self, organization: str
258
+ ) -> AsyncIterator[dict[str, Any]]:
259
+ """Iterate over repositories in an organization that have open PRs.
260
+
261
+ Args:
262
+ organization: Organization name
263
+
264
+ Yields:
265
+ Repository nodes with open PRs
266
+ """
267
+ has_next_page = True
268
+ end_cursor = None
269
+
270
+ while has_next_page:
271
+ variables: dict[str, Any] = {
272
+ "org": organization,
273
+ "cursor": end_cursor,
274
+ "prsPageSize": 1, # Just need to know if there are PRs
275
+ "contextsPageSize": 0, # Don't need status contexts
276
+ }
277
+
278
+ try:
279
+ result = await self.client.graphql(
280
+ ORG_REPOS_WITH_PRS, variables=variables
281
+ )
282
+ org_data = result.get("organization", {})
283
+ if not org_data:
284
+ break
285
+
286
+ repos = org_data.get("repositories", {})
287
+ nodes = repos.get("nodes", [])
288
+
289
+ # Only yield repos with open PRs
290
+ for node in nodes:
291
+ prs = node.get("pullRequests", {})
292
+ if prs.get("totalCount", 0) > 0:
293
+ yield node
294
+
295
+ page_info = repos.get("pageInfo", {})
296
+ has_next_page = page_info.get("hasNextPage", False)
297
+ end_cursor = page_info.get("endCursor")
298
+
299
+ except Exception as e:
300
+ self.logger.error(f"Error iterating repositories: {e}")
301
+ break
302
+
303
+ async def _fetch_repo_prs_first_page(
304
+ self, owner: str, repo_name: str
305
+ ) -> tuple[list[dict[str, Any]], dict[str, Any]]:
306
+ """Fetch first page of open PRs for a repository.
307
+
308
+ Args:
309
+ owner: Repository owner
310
+ repo_name: Repository name
311
+
312
+ Returns:
313
+ Tuple of (PR nodes list, page info dict)
314
+ """
315
+ variables = {
316
+ "owner": owner,
317
+ "name": repo_name,
318
+ "prsCursor": None,
319
+ "prsPageSize": DEFAULT_PRS_PAGE_SIZE,
320
+ "filesPageSize": DEFAULT_FILES_PAGE_SIZE,
321
+ "commentsPageSize": DEFAULT_COMMENTS_PAGE_SIZE,
322
+ "contextsPageSize": DEFAULT_CONTEXTS_PAGE_SIZE,
323
+ }
324
+
325
+ try:
326
+ result = await self.client.graphql(
327
+ REPO_OPEN_PRS_PAGE, variables=variables
328
+ )
329
+ repo_data = result.get("repository", {})
330
+ if not repo_data:
331
+ return [], {}
332
+
333
+ prs = repo_data.get("pullRequests", {})
334
+ pr_nodes = prs.get("nodes", [])
335
+ page_info = prs.get("pageInfo", {})
336
+ return pr_nodes, page_info
337
+
338
+ except Exception as e:
339
+ self.logger.error(
340
+ f"Error fetching PRs for {owner}/{repo_name}: {e}"
341
+ )
342
+ return [], {}
343
+
344
+ async def _iter_repo_open_prs_pages(
345
+ self, owner: str, repo_name: str, after_cursor: str
346
+ ) -> AsyncIterator[dict[str, Any]]:
347
+ """Iterate over additional pages of open PRs for a repository.
348
+
349
+ Args:
350
+ owner: Repository owner
351
+ repo_name: Repository name
352
+ after_cursor: Cursor to start from
353
+
354
+ Yields:
355
+ PR nodes
356
+ """
357
+ has_next_page = True
358
+ cursor = after_cursor
359
+
360
+ while has_next_page:
361
+ async with self.page_semaphore:
362
+ variables = {
363
+ "owner": owner,
364
+ "name": repo_name,
365
+ "prsCursor": cursor,
366
+ "prsPageSize": DEFAULT_PRS_PAGE_SIZE,
367
+ "filesPageSize": DEFAULT_FILES_PAGE_SIZE,
368
+ "commentsPageSize": DEFAULT_COMMENTS_PAGE_SIZE,
369
+ "contextsPageSize": DEFAULT_CONTEXTS_PAGE_SIZE,
370
+ }
371
+
372
+ try:
373
+ result = await self.client.graphql(
374
+ REPO_OPEN_PRS_PAGE, variables=variables
375
+ )
376
+ repo_data = result.get("repository", {})
377
+ if not repo_data:
378
+ break
379
+
380
+ prs = repo_data.get("pullRequests", {})
381
+ pr_nodes = prs.get("nodes", [])
382
+
383
+ for pr_node in pr_nodes:
384
+ yield pr_node
385
+
386
+ page_info = prs.get("pageInfo", {})
387
+ has_next_page = page_info.get("hasNextPage", False)
388
+ cursor = page_info.get("endCursor")
389
+
390
+ except Exception as e:
391
+ self.logger.error(
392
+ f"Error fetching PR page for {owner}/{repo_name}: {e}"
393
+ )
394
+ break
395
+
396
+ def _should_include_pr(
397
+ self, pr_node: dict[str, Any], include_drafts: bool
398
+ ) -> bool:
399
+ """Determine if a PR should be included in results.
400
+
401
+ Args:
402
+ pr_node: PR node from GraphQL response
403
+ include_drafts: Whether to include draft PRs
404
+
405
+ Returns:
406
+ True if PR should be included
407
+ """
408
+ return include_drafts or not pr_node.get("isDraft", False)
409
+
410
+ @staticmethod
411
+ def _split_owner_repo(full_name: str) -> tuple[str, str]:
412
+ """Split 'owner/repo' into separate components.
413
+
414
+ Args:
415
+ full_name: Repository full name (owner/repo)
416
+
417
+ Returns:
418
+ Tuple of (owner, repo_name)
419
+ """
420
+ parts = full_name.split("/", 1)
421
+ if len(parts) != 2:
422
+ return "unknown", "unknown"
423
+ return parts[0], parts[1]