diff --git a/gittensor/utils/github_api_tools.py b/gittensor/utils/github_api_tools.py index 6c367adeb..85419f860 100644 --- a/gittensor/utils/github_api_tools.py +++ b/gittensor/utils/github_api_tools.py @@ -403,7 +403,11 @@ def _resolve_pr_state(raw_state: str, merged: bool = False) -> str: def _search_issue_referencing_prs_graphql( - repo: str, issue_number: int, token: str, open_only: bool = False + repo: str, + issue_number: int, + token: str, + open_only: bool = False, + session: Optional[requests.Session] = None, ) -> Optional[List[PRInfo]]: """Fetch PRs that reference an issue via GraphQL issue timeline cross-references.""" if not token: @@ -421,6 +425,7 @@ def _search_issue_referencing_prs_graphql( variables={'owner': owner, 'name': name, 'issueNumber': issue_number}, token=token, max_attempts=3, + session=session, ) if result is None: bt.logging.warning(f'GraphQL cross-reference query failed for {repo}#{issue_number}') @@ -579,6 +584,7 @@ def execute_graphql_query( token: str, max_attempts: int = 8, timeout: int = 30, + session: Optional[requests.Session] = None, ) -> Optional[Dict[str, Any]]: """ Execute a GraphQL query with retry logic and backoff. @@ -589,15 +595,17 @@ def execute_graphql_query( token: GitHub PAT for authentication max_attempts: Maximum retry attempts (default 6) timeout: Request timeout in seconds (default 30) + session: Optional pooled requests.Session to reuse connections across calls Returns: Parsed JSON response data, or None if all attempts failed """ headers = make_graphql_headers(token) + http = session if session is not None else requests for attempt in range(max_attempts): try: - response = requests.post( + response = http.post( f'{BASE_GITHUB_API_URL}/graphql', headers=headers, json={'query': query, 'variables': variables}, @@ -1011,7 +1019,10 @@ def load_miners_prs( def find_solver_from_cross_references( - repo: str, issue_number: int, token: str + repo: str, + issue_number: int, + token: str, + session: Optional[requests.Session] = None, ) -> Optional[tuple[Optional[int], Optional[int]]]: """Resolve solver from cross-referenced PRs on the issue timeline. @@ -1027,7 +1038,7 @@ def find_solver_from_cross_references( tuple ``(solver_github_id, pr_number)`` where either value may be ``None`` when no valid closing PR is found. """ - prs = _search_issue_referencing_prs_graphql(repo, issue_number, token, open_only=False) + prs = _search_issue_referencing_prs_graphql(repo, issue_number, token, open_only=False, session=session) if prs is None: return None diff --git a/tests/utils/test_github_api_tools.py b/tests/utils/test_github_api_tools.py index 7ccb62312..04da7381b 100644 --- a/tests/utils/test_github_api_tools.py +++ b/tests/utils/test_github_api_tools.py @@ -1583,5 +1583,48 @@ def test_falls_back_to_base_ref_oid_when_merge_base_fails(self, mock_merge_base, assert call_args[0][2] == 'base_branch_tip_sha', 'Should fall back to base_ref_oid' +# ============================================================================ +# Session Pooling Tests +# ============================================================================ + + +class TestExecuteGraphQLQuerySessionPooling: + """Verify execute_graphql_query reuses a pooled Session when one is supplied.""" + + @patch('gittensor.utils.github_api_tools.requests.post') + def test_uses_session_post_when_session_provided(self, mock_requests_post): + mock_session = Mock(spec=['post']) + mock_session.post.return_value = Mock(status_code=200, json=Mock(return_value={'data': {}})) + + result = execute_graphql_query('query {}', {}, 'fake_token', session=mock_session) + + assert result == {'data': {}} + mock_session.post.assert_called_once() + mock_requests_post.assert_not_called() + + @patch('gittensor.utils.github_api_tools.requests.post') + def test_falls_back_to_requests_post_when_no_session(self, mock_requests_post): + mock_requests_post.return_value = Mock(status_code=200, json=Mock(return_value={'data': {}})) + + result = execute_graphql_query('query {}', {}, 'fake_token') + + assert result == {'data': {}} + mock_requests_post.assert_called_once() + + +class TestFindSolverSessionPooling: + """Verify find_solver_from_cross_references threads the session to the GraphQL layer.""" + + @patch('gittensor.utils.github_api_tools.execute_graphql_query') + def test_session_is_forwarded_to_execute_graphql_query(self, mock_graphql): + mock_graphql.return_value = _graphql_response([]) + mock_session = Mock(spec=['post', 'get']) + + find_solver_from_cross_references('owner/repo', 12, 'fake_token', session=mock_session) + + assert mock_graphql.call_count == 1 + assert mock_graphql.call_args.kwargs.get('session') is mock_session + + if __name__ == '__main__': pytest.main([__file__, '-v'])