diff --git a/AGENTS.md b/AGENTS.md index a2b23351f..7f73d76bf 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -7,7 +7,7 @@ This guide is designed for AI agents working on the Ymir AI workflows project. I **Consult these first:** - **[README-agents.md](README-agents.md)** — Full setup, running agents, environment variables, Jira mocking - **[README.md](README.md)** — Project overview, development environment setup -- **[CONTRIBUTING.md](CONTRIBUTING.md)** — Code merge policy +- **[CONTRIBUTING.md](CONTRIBUTING.md)** — Code merge policy, Mocking in tests ## Agent Architecture @@ -121,7 +121,7 @@ If you introduce a new service as a dependency to our agents, make sure to read ## Code Changes Checklist -- [ ] Write tests first (especially for tools/git operations) +- [ ] Write tests first (especially for tools/git operations) — make sure they use `flexmock` for mocking. - [ ] Run `make check-in-container` — all tests pass - [ ] Test with `DRY_RUN=true` — don't touch real Jira/git - [ ] Use rebase merge (see [CONTRIBUTING.md](CONTRIBUTING.md)) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 5996f7898..b42a2b30b 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -42,4 +42,6 @@ Prefer rebase-merging over creating a merge commit, unless preserving the branch ## Mocking -`flexmock` is the preferred framework for mocking in tests. +`flexmock` is the preferred mock framework in tests ahead of `pytest-mock` and `unittest.mock`. +Do not use `AsyncMock`, `MagicMock` and `patch` constructs, use `flexmock` instead. +Since `pytest` is used as the general testing framework, `monkeypatch` may be used in some mock cases (environment variables, etc). diff --git a/ymir/agents/tests/unit/test_rebase_consolidation.py b/ymir/agents/tests/unit/test_rebase_consolidation.py index 608fcc04e..adbaae270 100644 --- a/ymir/agents/tests/unit/test_rebase_consolidation.py +++ b/ymir/agents/tests/unit/test_rebase_consolidation.py @@ -1,5 +1,7 @@ import pytest +from flexmock import flexmock +from ymir.agents import rebase_consolidation from ymir.agents.rebase_agent import _consolidated_issue_keys from ymir.agents.rebase_consolidation import ( add_jira_tickets_to_latest_changelog_entry, @@ -9,7 +11,8 @@ has_new_latest_changelog_entry, uses_autochangelog, ) -from ymir.common.models import ConsolidatedIssue +from ymir.common.constants import JiraLabels +from ymir.common.models import ConsolidatedIssue, RebaseData from ymir.common.utils import extract_text_from_adf @@ -391,8 +394,6 @@ class TestTerminalLabels: def test_jql_excludes_all_triage_decision_labels(self): """JQL must exclude all triage decision labels to avoid re-queueing decided siblings.""" - from ymir.common.constants import JiraLabels - jql = build_rebase_siblings_jql("RHEL-100", "postgresql", "rhel-9.8") # Verify each triage decision label appears in the JQL exclusion @@ -412,8 +413,6 @@ def test_jql_excludes_all_completion_labels(self): Regression test for RHEL-248139 where ymir_backported was not excluded. """ - from ymir.common.constants import JiraLabels - jql = build_rebase_siblings_jql("RHEL-100", "postgresql", "rhel-9.8") # These were the missing labels that caused RHEL-248139 @@ -430,8 +429,6 @@ def test_jql_excludes_errored_labels(self): Per jira_label_workflow_routing.md: ERRORED labels (triage/backport/rebase_errored) block retry and need human attention, so they're terminal for sibling queueing. """ - from ymir.common.constants import JiraLabels - jql = build_rebase_siblings_jql("RHEL-100", "postgresql", "rhel-9.8") # ERRORED labels block retry → must exclude @@ -452,8 +449,6 @@ def test_jql_includes_failed_labels(self): "May auto-retry", so excluding them breaks the retry mechanism where a new sibling triggers re-queueing of failed issues. """ - from ymir.common.constants import JiraLabels - jql = build_rebase_siblings_jql("RHEL-100", "postgresql", "rhel-9.8") # FAILED labels may auto-retry → must NOT exclude @@ -475,8 +470,6 @@ def test_jql_does_not_exclude_sibling_marker(self): queue_siblings_for_triage() handles the re-queueing check in its defensive filter. """ - from ymir.common.constants import JiraLabels - jql = build_rebase_siblings_jql("RHEL-100", "postgresql", "rhel-9.8") assert f'"{JiraLabels.REBASE_SIBLING.value}"' not in jql, ( @@ -492,8 +485,6 @@ def test_jql_exclusion_applies_before_50_result_limit(self): This test verifies the exclusion is in the JQL string (server-side filtering). """ - from ymir.common.constants import JiraLabels - jql = build_rebase_siblings_jql("RHEL-100", "postgresql", "rhel-9.8") # Critical: the exclusion MUST be in the JQL query string itself @@ -552,10 +543,6 @@ async def test_find_triaged_rebase_siblings_no_unbound_error_on_jira_failure(): The code then checks ``if downstream_component is None and primary_details:``. Without ``primary_details = None`` before try, this raises UnboundLocalError. """ - from unittest.mock import AsyncMock, patch - - from ymir.common.models import RebaseData - rebase_data = RebaseData( package="postgis", version="3.5.2", @@ -564,17 +551,14 @@ async def test_find_triaged_rebase_siblings_no_unbound_error_on_jira_failure(): ) # Simulate get_jira_details failure (MCP gateway down, network error, etc.) - with patch( - "ymir.agents.rebase_consolidation.run_tool", - new_callable=AsyncMock, - side_effect=Exception("MCP gateway unreachable"), - ): - # Should NOT raise UnboundLocalError — should return empty results - result, summary = await find_triaged_rebase_siblings( - jira_issue="RHEL-250764", - rebase_data=rebase_data, - available_tools=[], - downstream_component=None, - ) + flexmock(rebase_consolidation).should_receive("run_tool").and_raise(Exception("MCP Gateway unreachable")) + + # Should NOT raise UnboundLocalError — should return empty results + result, summary = await find_triaged_rebase_siblings( + jira_issue="RHEL-250764", + rebase_data=rebase_data, + available_tools=[], + downstream_component=None, + ) assert result == [] assert summary == "" diff --git a/ymir/agents/tests/unit/test_triage_agent.py b/ymir/agents/tests/unit/test_triage_agent.py index 8fb9b502c..cc187da55 100644 --- a/ymir/agents/tests/unit/test_triage_agent.py +++ b/ymir/agents/tests/unit/test_triage_agent.py @@ -1,5 +1,4 @@ from contextlib import asynccontextmanager -from unittest.mock import AsyncMock, patch import pytest from flexmock import flexmock @@ -33,6 +32,7 @@ Resolution, Task, TriageEligibility, + TriageInputSchema, TriageOutputSchema, ) from ymir.common.version_utils import extract_downstream_package, is_modular, parse_module_stream @@ -406,30 +406,33 @@ def _cve_eligibility(*, needs_internal_fix: bool) -> CVEEligibilityResult: ) +async def _older_zstream_true(*_args, **_kwargs): + return True + + +async def _older_zstream_false(*_args, **_kwargs): + return False + + @pytest.mark.asyncio async def test_determine_target_branch_modular_internal_fix_no_ystream_uses_cs(): """RHEL 8 has no Y-stream, so even CVEs needing internal fix go to centos-stream.""" - with ( - patch( - "ymir.agents.triage_agent.is_older_zstream", - new_callable=AsyncMock, - return_value=False, - ), - patch( - "ymir.agents.triage_agent.load_rhel_config", - new_callable=AsyncMock, - return_value={ - "current_y_streams": {"9": "rhel-9.9", "10": "rhel-10.3"}, - "current_z_streams": {"8": "rhel-8.10.z", "9": "rhel-9.8.z", "10": "rhel-10.2.z"}, - }, - ), - ): - branch, namespace = await determine_target_branch( - _cve_eligibility(needs_internal_fix=True), - _modular_backport_data(), - jira_summary=_MODULAR_SUMMARY, - downstream_component="squid", - ) + + async def _mock_load_rhel_config(*_args, **_kwargs): + return { + "current_y_streams": {"9": "rhel-9.9", "10": "rhel-10.3"}, + "current_z_streams": {"8": "rhel-8.10.z", "9": "rhel-9.8.z", "10": "rhel-10.2.z"}, + } + + flexmock(t_agent).should_receive("is_older_zstream").replace_with(_older_zstream_false) + flexmock(t_agent).should_receive("load_rhel_config").replace_with(_mock_load_rhel_config) + + branch, namespace = await determine_target_branch( + _cve_eligibility(needs_internal_fix=True), + _modular_backport_data(), + jira_summary=_MODULAR_SUMMARY, + downstream_component="squid", + ) assert branch == "stream-squid-4-rhel-8.10.0" assert namespace == "centos-stream" @@ -437,34 +440,27 @@ async def test_determine_target_branch_modular_internal_fix_no_ystream_uses_cs() @pytest.mark.asyncio async def test_determine_target_branch_modular_internal_fix_with_ystream_uses_rhel(): """RHEL 9 has a Y-stream, so CVEs needing internal fix go to rhel.""" + + async def _mock_load_rhel_config(*_args, **_kwargs): + return {"current_y_streams": {"9": "rhel-9.9", "10": "rhel-10.3"}} + summary = "CVE-2026-32748 squid:4/squid: Squid: Denial of Service [rhel-9.8.z]" - with ( - patch( - "ymir.agents.triage_agent.is_older_zstream", - new_callable=AsyncMock, - return_value=False, - ), - patch( - "ymir.agents.triage_agent.load_rhel_config", - new_callable=AsyncMock, - return_value={"current_y_streams": {"9": "rhel-9.9", "10": "rhel-10.3"}}, - ), - ): - branch, namespace = await determine_target_branch( - _cve_eligibility(needs_internal_fix=True), - _modular_backport_data(fix_version="rhel-9.8.z"), - jira_summary=summary, - downstream_component="squid", - ) + + flexmock(t_agent).should_receive("is_older_zstream").replace_with(_older_zstream_false) + flexmock(t_agent).should_receive("load_rhel_config").replace_with(_mock_load_rhel_config) + + branch, namespace = await determine_target_branch( + _cve_eligibility(needs_internal_fix=True), + _modular_backport_data(fix_version="rhel-9.8.z"), + jira_summary=summary, + downstream_component="squid", + ) assert branch == "stream-squid-4-rhel-9.8.0" assert namespace == "rhel" @pytest.mark.asyncio async def test_determine_target_branch_modular_cs_eligible_uses_centos_stream(): - async def _older_zstream_false(*_args, **_kwargs): - return False - flexmock(t_agent).should_receive("is_older_zstream").replace_with(_older_zstream_false) branch, namespace = await determine_target_branch( @@ -479,9 +475,6 @@ async def _older_zstream_false(*_args, **_kwargs): @pytest.mark.asyncio async def test_determine_target_branch_modular_older_zstream_uses_rhel(): - async def _older_zstream_true(*_args, **_kwargs): - return True - flexmock(t_agent).should_receive("is_older_zstream").replace_with(_older_zstream_true) branch, namespace = await determine_target_branch( @@ -500,32 +493,26 @@ async def test_render_prompt_modular_rhel8_no_internal_fix(): for modular issues even when CVE eligibility says needs_internal_fix=True. Otherwise the prompt tells the LLM to clone from rhel namespace instead of centos-stream.""" - from ymir.common.models import TriageInputSchema as InputSchema - input_data = InputSchema(issue="RHEL-999") + async def _mock_load_rhel_config(*_args, **_kwargs): + return { + "current_y_streams": {"9": "rhel-9.9", "10": "rhel-10.3"}, + "current_z_streams": {"8": "rhel-8.10.z", "9": "rhel-9.8.z", "10": "rhel-10.2.z"}, + } + + input_data = TriageInputSchema(issue="RHEL-999") summary = "CVE-2026-32748 squid:4/squid: Denial of Service [rhel-8.10.z]" - with ( - patch( - "ymir.agents.triage_agent.is_older_zstream", - new_callable=AsyncMock, - return_value=False, - ), - patch( - "ymir.agents.triage_agent.load_rhel_config", - new_callable=AsyncMock, - return_value={ - "current_y_streams": {"9": "rhel-9.9", "10": "rhel-10.3"}, - "current_z_streams": {"8": "rhel-8.10.z", "9": "rhel-9.8.z", "10": "rhel-10.2.z"}, - }, - ), - ): - prompt = await render_prompt( - input_data, - fix_version="rhel-8.10.z", - cve_eligibility_result=_cve_eligibility(needs_internal_fix=True), - jira_summary=summary, - downstream_component="squid", - ) + + flexmock(t_agent).should_receive("is_older_zstream").replace_with(_older_zstream_false) + flexmock(t_agent).should_receive("load_rhel_config").replace_with(_mock_load_rhel_config) + + prompt = await render_prompt( + input_data, + fix_version="rhel-8.10.z", + cve_eligibility_result=_cve_eligibility(needs_internal_fix=True), + jira_summary=summary, + downstream_component="squid", + ) assert "stream-squid-4-rhel-8.10.0" not in prompt assert "redhat/rhel/rpms" not in prompt @@ -534,29 +521,23 @@ async def test_render_prompt_modular_rhel8_no_internal_fix(): async def test_render_prompt_modular_rhel9_has_internal_fix(): """RHEL 9 has a Y-stream, so render_prompt SHOULD set needs_internal_fix for modular issues when CVE eligibility says needs_internal_fix=True.""" - from ymir.common.models import TriageInputSchema as InputSchema - input_data = InputSchema(issue="RHEL-999") + async def _mock_load_rhel_config(*_args, **_kwargs): + return {"current_y_streams": {"9": "rhel-9.9", "10": "rhel-10.3"}} + + input_data = TriageInputSchema(issue="RHEL-999") summary = "CVE-2026-32748 squid:4/squid: Denial of Service [rhel-9.8.z]" - with ( - patch( - "ymir.agents.triage_agent.is_older_zstream", - new_callable=AsyncMock, - return_value=False, - ), - patch( - "ymir.agents.triage_agent.load_rhel_config", - new_callable=AsyncMock, - return_value={"current_y_streams": {"9": "rhel-9.9", "10": "rhel-10.3"}}, - ), - ): - prompt = await render_prompt( - input_data, - fix_version="rhel-9.8.z", - cve_eligibility_result=_cve_eligibility(needs_internal_fix=True), - jira_summary=summary, - downstream_component="squid", - ) + + flexmock(t_agent).should_receive("is_older_zstream").replace_with(_older_zstream_false) + flexmock(t_agent).should_receive("load_rhel_config").replace_with(_mock_load_rhel_config) + + prompt = await render_prompt( + input_data, + fix_version="rhel-9.8.z", + cve_eligibility_result=_cve_eligibility(needs_internal_fix=True), + jira_summary=summary, + downstream_component="squid", + ) assert "stream-squid-4-rhel-9.8.0" in prompt diff --git a/ymir/api/tests/test_auth.py b/ymir/api/tests/test_auth.py index 38e5ac2c4..f742e6c84 100644 --- a/ymir/api/tests/test_auth.py +++ b/ymir/api/tests/test_auth.py @@ -2,7 +2,6 @@ import json import time -from unittest.mock import patch import pytest import pytest_asyncio @@ -10,6 +9,8 @@ from aiohttp.test_utils import TestClient, TestServer from jwcrypto import jwk, jwt +from ymir.api import auth + # --------------------------------------------------------------------------- # Helpers: generate a key pair + sign a test token # --------------------------------------------------------------------------- @@ -65,39 +66,36 @@ def valid_token(keyset): @pytest_asyncio.fixture -async def app_and_client(keyset): +async def app_and_client(keyset, monkeypatch): """Create a minimal aiohttp app with the OIDC middleware and a test route.""" - from ymir.api import auth # Patch env vars and JWKS cache for the middleware. - with ( - patch.object(auth, "OIDC_PROVIDER_URL", "http://localhost:8084/realms/ymir"), - patch.object(auth, "OIDC_ISSUER", "http://localhost:8084/realms/ymir"), - patch.object(auth, "OIDC_CLIENT_ID", "ymir-trace-ui"), - patch.object(auth, "OIDC_CORS_ALLOWED_ORIGIN", "http://localhost:8082"), - patch.object(auth, "OIDC_CORS_ALLOWED_ORIGIN_ALT", "DISABLED"), - patch.object(auth, "_jwks_cache", keyset), - patch.object(auth, "_jwks_fetched_at", time.monotonic()), - ): - - async def protected_handler(request: web.Request) -> web.Response: - return web.json_response( - { - "user": request.get("remote_user", "unknown"), - } - ) - - async def public_handler(request: web.Request) -> web.Response: - return web.json_response({"status": "ok"}) - - app = web.Application(middlewares=[auth.oidc_middleware]) - app.router.add_get("/healthz", public_handler) - app.router.add_get("/readyz", public_handler) - app.router.add_post("/api/jira/webhook", public_handler) - app.router.add_post("/api/consolidation", protected_handler) - - async with TestClient(TestServer(app)) as client: - yield app, client + monkeypatch.setattr(auth, "OIDC_PROVIDER_URL", "http://localhost:8084/realms/ymir") + monkeypatch.setattr(auth, "OIDC_ISSUER", "http://localhost:8084/realms/ymir") + monkeypatch.setattr(auth, "OIDC_CLIENT_ID", "ymir-trace-ui") + monkeypatch.setattr(auth, "OIDC_CORS_ALLOWED_ORIGIN", "http://localhost:8082") + monkeypatch.setattr(auth, "OIDC_CORS_ALLOWED_ORIGIN_ALT", "DISABLED") + monkeypatch.setattr(auth, "_jwks_cache", keyset) + monkeypatch.setattr(auth, "_jwks_fetched_at", time.monotonic()) + + async def protected_handler(request: web.Request) -> web.Response: + return web.json_response( + { + "user": request.get("remote_user", "unknown"), + } + ) + + async def public_handler(request: web.Request) -> web.Response: + return web.json_response({"status": "ok"}) + + app = web.Application(middlewares=[auth.oidc_middleware]) + app.router.add_get("/healthz", public_handler) + app.router.add_get("/readyz", public_handler) + app.router.add_post("/api/jira/webhook", public_handler) + app.router.add_post("/api/consolidation", protected_handler) + + async with TestClient(TestServer(app)) as client: + yield app, client # --------------------------------------------------------------------------- @@ -339,22 +337,19 @@ async def test_user_identity_falls_back_to_email(app_and_client, keyset): @pytest.mark.asyncio -async def test_oidc_disabled_passes_through(): +async def test_oidc_disabled_passes_through(monkeypatch): """When OIDC_PROVIDER_URL is empty, all requests pass through.""" - from ymir.api import auth - with ( - patch.object(auth, "OIDC_PROVIDER_URL", ""), - patch.object(auth, "OIDC_CORS_ALLOWED_ORIGIN", ""), - patch.object(auth, "OIDC_CORS_ALLOWED_ORIGIN_ALT", ""), - ): + monkeypatch.setattr(auth, "OIDC_PROVIDER_URL", "") + monkeypatch.setattr(auth, "OIDC_CORS_ALLOWED_ORIGIN", "") + monkeypatch.setattr(auth, "OIDC_CORS_ALLOWED_ORIGIN_ALT", "") - async def handler(request: web.Request) -> web.Response: - return web.json_response({"ok": True}) + async def handler(request: web.Request) -> web.Response: + return web.json_response({"ok": True}) - app = web.Application(middlewares=[auth.oidc_middleware]) - app.router.add_post("/api/consolidation", handler) + app = web.Application(middlewares=[auth.oidc_middleware]) + app.router.add_post("/api/consolidation", handler) - async with TestClient(TestServer(app)) as client: - resp = await client.post("/api/consolidation", json={}) - assert resp.status == 200 + async with TestClient(TestServer(app)) as client: + resp = await client.post("/api/consolidation", json={}) + assert resp.status == 200 diff --git a/ymir/api/tests/test_command_parser.py b/ymir/api/tests/test_command_parser.py index 80e1cbcc7..ef7eafca5 100644 --- a/ymir/api/tests/test_command_parser.py +++ b/ymir/api/tests/test_command_parser.py @@ -3,10 +3,10 @@ from __future__ import annotations import json -from unittest.mock import AsyncMock import pytest from aiohttp import web +from flexmock import flexmock from ymir.api import command_parser @@ -27,13 +27,16 @@ def _parse_response_body(resp: web.Response) -> dict: @pytest.mark.asyncio async def test_dispatch_calls_registered_handler(): - handler = AsyncMock(return_value=web.json_response({"ok": True}, status=201)) - command_parser.register("hello", handler) + async def _mock_json_response(*_args, **_kwargs): + return web.json_response({"ok": True}, status=201) request = object() + handler = flexmock() + handler.should_receive("handle").with_args(["world"], request).replace_with(_mock_json_response).once() + + command_parser.register("hello", handler.handle) resp = await command_parser.dispatch("hello world", request) - handler.assert_awaited_once_with(["world"], request) assert resp.status == 201 @@ -55,22 +58,30 @@ async def test_dispatch_empty_command(): @pytest.mark.asyncio async def test_dispatch_case_insensitive(): - handler = AsyncMock(return_value=web.json_response({"ok": True})) - command_parser.register("greet", handler) + async def _mock_json_response(*_args, **_kwargs): + return web.json_response({"ok": True}) request = object() + handler = flexmock() + handler.should_receive("handle").with_args(["Alice"], request).replace_with(_mock_json_response).once() + + command_parser.register("greet", handler.handle) await command_parser.dispatch("GREET Alice", request) - handler.assert_awaited_once_with(["Alice"], request) @pytest.mark.asyncio async def test_dispatch_handles_quoted_args(): - handler = AsyncMock(return_value=web.json_response({"ok": True})) - command_parser.register("echo", handler) + async def _mock_json_response(*_args, **_kwargs): + return web.json_response({"ok": True}) request = object() + handler = flexmock() + handler.should_receive("handle").with_args(["hello world", "foo"], request).replace_with( + _mock_json_response + ).once() + + command_parser.register("echo", handler.handle) await command_parser.dispatch('echo "hello world" foo', request) - handler.assert_awaited_once_with(["hello world", "foo"], request) @pytest.mark.asyncio diff --git a/ymir/api/tests/test_jira_reply.py b/ymir/api/tests/test_jira_reply.py index 4408aabf2..10230315e 100644 --- a/ymir/api/tests/test_jira_reply.py +++ b/ymir/api/tests/test_jira_reply.py @@ -1,32 +1,33 @@ """Unit tests for the Jira reply helper.""" -from unittest.mock import AsyncMock, MagicMock, patch - import pytest +from flexmock import flexmock +from ymir.api import jira_reply from ymir.api.jira_reply import post_comment +class _AsyncContextManager: + def __init__(self, value): + self.value = value + + async def __aenter__(self): + return self.value + + async def __aexit__(self, *_args): + return False + + @pytest.mark.asyncio async def test_post_comment_success(monkeypatch): monkeypatch.setenv("JIRA_URL", "https://issues.example.com") monkeypatch.setenv("JIRA_EMAIL", "bot@example.com") monkeypatch.setenv("JIRA_TOKEN", "fake-token") # pragma: allowlist secret - mock_response = AsyncMock() - mock_response.status = 201 - mock_response.__aenter__ = AsyncMock(return_value=mock_response) - mock_response.__aexit__ = AsyncMock(return_value=False) + mock_response = flexmock(status=201) - mock_session = AsyncMock() - mock_session.post = MagicMock(return_value=mock_response) - mock_session.__aenter__ = AsyncMock(return_value=mock_session) - mock_session.__aexit__ = AsyncMock(return_value=False) - - with patch("ymir.api.jira_reply.aiohttp.ClientSession", return_value=mock_session): - await post_comment("RHEL-12345", "Command failed: bad args") - - mock_session.post.assert_called_once_with( + mock_session = flexmock() + mock_session.should_receive("post").with_args( "https://issues.example.com/rest/api/2/issue/RHEL-12345/comment", json={"body": "Command failed: bad args"}, headers={ @@ -34,17 +35,22 @@ async def test_post_comment_success(monkeypatch): "Content-Type": "application/json", "Accept": "application/json", }, + ).and_return(_AsyncContextManager(mock_response)).once() + + flexmock(jira_reply.aiohttp).should_receive("ClientSession").replace_with( + lambda: _AsyncContextManager(mock_session) ) + await post_comment("RHEL-12345", "Command failed: bad args") + @pytest.mark.asyncio async def test_post_comment_no_jira_url(monkeypatch): monkeypatch.delenv("JIRA_URL", raising=False) - with patch("ymir.api.jira_reply.aiohttp.ClientSession") as mock_cls: - await post_comment("RHEL-12345", "should not call Jira") + flexmock(jira_reply.aiohttp).should_receive("ClientSession").never() - mock_cls.assert_not_called() + await post_comment("RHEL-12345", "should not call Jira") @pytest.mark.asyncio @@ -54,19 +60,20 @@ async def test_post_comment_jira_error(monkeypatch): monkeypatch.setenv("JIRA_EMAIL", "bot@example.com") monkeypatch.setenv("JIRA_TOKEN", "fake-token") # pragma: allowlist secret - mock_response = AsyncMock() - mock_response.status = 403 - mock_response.text = AsyncMock(return_value="Forbidden") - mock_response.__aenter__ = AsyncMock(return_value=mock_response) - mock_response.__aexit__ = AsyncMock(return_value=False) + async def _mock_response_text(*_args, **_kwargs): + return "Forbidden" - mock_session = AsyncMock() - mock_session.post = MagicMock(return_value=mock_response) - mock_session.__aenter__ = AsyncMock(return_value=mock_session) - mock_session.__aexit__ = AsyncMock(return_value=False) + mock_response = flexmock(status=403) + mock_response.should_receive("text").replace_with(_mock_response_text) - with patch("ymir.api.jira_reply.aiohttp.ClientSession", return_value=mock_session): - await post_comment("RHEL-12345", "some error") + mock_session = flexmock() + mock_session.should_receive("post").and_return(_AsyncContextManager(mock_response)) + + flexmock(jira_reply.aiohttp).should_receive("ClientSession").replace_with( + lambda: _AsyncContextManager(mock_session) + ) + + await post_comment("RHEL-12345", "some error") @pytest.mark.asyncio @@ -76,13 +83,14 @@ async def test_post_comment_network_exception(monkeypatch): monkeypatch.setenv("JIRA_EMAIL", "bot@example.com") monkeypatch.setenv("JIRA_TOKEN", "fake-token") # pragma: allowlist secret - mock_session = AsyncMock() - mock_session.post = MagicMock(side_effect=OSError("connection refused")) - mock_session.__aenter__ = AsyncMock(return_value=mock_session) - mock_session.__aexit__ = AsyncMock(return_value=False) + mock_session = flexmock() + mock_session.should_receive("post").and_raise(OSError("connection refused")) + + flexmock(jira_reply.aiohttp).should_receive("ClientSession").replace_with( + lambda: _AsyncContextManager(mock_session) + ) - with patch("ymir.api.jira_reply.aiohttp.ClientSession", return_value=mock_session): - await post_comment("RHEL-12345", "some error") + await post_comment("RHEL-12345", "some error") @pytest.mark.asyncio @@ -92,18 +100,22 @@ async def test_post_comment_trailing_slash(monkeypatch): monkeypatch.setenv("JIRA_EMAIL", "bot@example.com") monkeypatch.setenv("JIRA_TOKEN", "fake-token") # pragma: allowlist secret - mock_response = AsyncMock() - mock_response.status = 201 - mock_response.__aenter__ = AsyncMock(return_value=mock_response) - mock_response.__aexit__ = AsyncMock(return_value=False) + mock_response = flexmock(status=201) - mock_session = AsyncMock() - mock_session.post = MagicMock(return_value=mock_response) - mock_session.__aenter__ = AsyncMock(return_value=mock_session) - mock_session.__aexit__ = AsyncMock(return_value=False) + call_url = "http://should/be//overwritten//by/mock_post" + + def _mock_post(*_args, **_kwargs): + nonlocal call_url + call_url = _args[0] + return _AsyncContextManager(mock_response) + + mock_session = flexmock() + mock_session.should_receive("post").replace_with(_mock_post) + + flexmock(jira_reply.aiohttp).should_receive("ClientSession").replace_with( + lambda: _AsyncContextManager(mock_session) + ) - with patch("ymir.api.jira_reply.aiohttp.ClientSession", return_value=mock_session): - await post_comment("RHEL-12345", "msg") + await post_comment("RHEL-12345", "msg") - call_url = mock_session.post.call_args[0][0] assert "//" not in call_url.split("://")[1] diff --git a/ymir/api/tests/test_jira_webhook.py b/ymir/api/tests/test_jira_webhook.py index 606d54bf9..3739e32ad 100644 --- a/ymir/api/tests/test_jira_webhook.py +++ b/ymir/api/tests/test_jira_webhook.py @@ -4,12 +4,13 @@ import hashlib import hmac as hmac_mod import json -from unittest.mock import AsyncMock, patch import pytest import pytest_asyncio from aiohttp.test_utils import TestClient, TestServer +from flexmock import flexmock +from ymir.api import jira_reply, jira_webhook from ymir.api.jira_webhook import ( _extract_command, _extract_command_from_adf, @@ -80,12 +81,11 @@ def _mock_rh_employee(): Individual tests override this when they need to exercise the rejection path. """ - with patch( - "ymir.api.jira_webhook._is_rh_employee", - new_callable=AsyncMock, - return_value=True, - ): - yield + + async def _mock_rh_employee_true(*_args, **_kwargs): + return True + + flexmock(jira_webhook).should_receive("_is_rh_employee").replace_with(_mock_rh_employee_true) @pytest_asyncio.fixture @@ -396,14 +396,19 @@ async def test_fail_closed_when_secret_not_configured(client, monkeypatch): @pytest.mark.asyncio async def test_error_triggers_jira_comment(client): """A 400 response from dispatch should schedule a Jira comment.""" + call_args = [] + + async def _mock_post_comment(*_args, **_kwargs): + call_args.append(_args) + payload = _comment_payload(_adf_mention_body("do-something-unknown arg1")) - with patch("ymir.api.jira_webhook.jira_reply.post_comment", new_callable=AsyncMock) as mock_post: - resp = await _signed_post(client, payload) - assert resp.status == 400 - await asyncio.sleep(0) # let fire-and-forget create_task drain + flexmock(jira_reply).should_receive("post_comment").replace_with(_mock_post_comment).once() + + resp = await _signed_post(client, payload) + + assert resp.status == 400 + await asyncio.sleep(0) # let fire-and-forget create_task drain - mock_post.assert_awaited_once() - call_args = mock_post.call_args assert call_args[0][0] == "RHEL-99999" assert "unknown command" in call_args[0][1] @@ -412,11 +417,11 @@ async def test_error_triggers_jira_comment(client): async def test_success_does_not_trigger_jira_comment(client): """A 201 success response should NOT schedule a Jira comment.""" payload = _comment_payload(_adf_mention_body("consolidate expat rhel-9.8.0")) - with patch("ymir.api.jira_webhook.jira_reply.post_comment", new_callable=AsyncMock) as mock_post: - resp = await _signed_post(client, payload) - assert resp.status == 201 + flexmock(jira_reply).should_receive("post_comment").never() - mock_post.assert_not_awaited() + resp = await _signed_post(client, payload) + + assert resp.status == 201 @pytest.mark.asyncio @@ -429,30 +434,36 @@ async def test_error_comment_with_missing_issue_key(client): "author": {"accountId": "user-123", "displayName": "Test User"}, }, } - with patch("ymir.api.jira_webhook.jira_reply.post_comment", new_callable=AsyncMock) as mock_post: - resp = await client.post( - "/api/jira/webhook", - json=payload, - headers=_sign(payload), - ) - assert resp.status == 400 - await asyncio.sleep(0) + flexmock(jira_reply).should_receive("post_comment").never() + + resp = await client.post( + "/api/jira/webhook", + json=payload, + headers=_sign(payload), + ) - mock_post.assert_not_awaited() + assert resp.status == 400 + await asyncio.sleep(0) @pytest.mark.asyncio async def test_malformed_consolidate_triggers_comment(client): """Malformed consolidate args should trigger an error comment.""" + call_args = [] + + async def _mock_post_comment(*_args, **_kwargs): + call_args.append(_args) + payload = _comment_payload(_adf_mention_body("consolidate")) - with patch("ymir.api.jira_webhook.jira_reply.post_comment", new_callable=AsyncMock) as mock_post: - resp = await _signed_post(client, payload) - assert resp.status == 400 - await asyncio.sleep(0) # let fire-and-forget create_task drain + flexmock(jira_reply).should_receive("post_comment").replace_with(_mock_post_comment).once() - mock_post.assert_awaited_once() - assert mock_post.call_args[0][0] == "RHEL-99999" - assert "invalid consolidate arguments" in mock_post.call_args[0][1] + resp = await _signed_post(client, payload) + + assert resp.status == 400 + await asyncio.sleep(0) # let fire-and-forget create_task drain + + assert call_args[0][0] == "RHEL-99999" + assert "invalid consolidate arguments" in call_args[0][1] # -- Red Hat Employee group check --------------------------------------------- @@ -471,13 +482,15 @@ async def test_rh_employee_command_accepted(client): @pytest.mark.asyncio async def test_non_rh_employee_rejected(client): """A non-employee comment author should receive a 403.""" - with patch( - "ymir.api.jira_webhook._is_rh_employee", - new_callable=AsyncMock, - return_value=False, - ): - payload = _comment_payload(_adf_mention_body("consolidate expat rhel-9.8.0")) - resp = await _signed_post(client, payload) + + async def _mock_rh_employee_false(*_args, **_kwargs): + return False + + flexmock(jira_webhook).should_receive("_is_rh_employee").replace_with(_mock_rh_employee_false) + payload = _comment_payload(_adf_mention_body("consolidate expat rhel-9.8.0")) + + resp = await _signed_post(client, payload) + assert resp.status == 403 body = await resp.json() assert "not a Red Hat employee" in body["error"] @@ -507,13 +520,12 @@ async def test_missing_author_account_id_rejected(client): @pytest.mark.asyncio async def test_jira_api_failure_fails_closed(client): """If the Jira user lookup raises, the command must be rejected (fail closed).""" - with patch( - "ymir.api.jira_webhook._is_rh_employee", - new_callable=AsyncMock, - side_effect=Exception("connection refused"), - ): - payload = _comment_payload(_adf_mention_body("consolidate expat rhel-9.8.0")) - resp = await _signed_post(client, payload) + + flexmock(jira_webhook).should_receive("_is_rh_employee").and_raise(Exception("connection refused")) + payload = _comment_payload(_adf_mention_body("consolidate expat rhel-9.8.0")) + + resp = await _signed_post(client, payload) + assert resp.status == 403 body = await resp.json() assert "not a Red Hat employee" in body["error"] @@ -522,13 +534,11 @@ async def test_jira_api_failure_fails_closed(client): @pytest.mark.asyncio async def test_non_command_comment_skips_employee_check(client): """Comments without a bot mention must be ignored without calling the employee check.""" - with patch( - "ymir.api.jira_webhook._is_rh_employee", - new_callable=AsyncMock, - ) as mock_check: - payload = _comment_payload(_adf_plain_body("just a regular comment")) - resp = await _signed_post(client, payload) + + flexmock(jira_webhook).should_receive("_is_rh_employee").never() + payload = _comment_payload(_adf_plain_body("just a regular comment")) + + resp = await _signed_post(client, payload) assert resp.status == 200 body = await resp.json() assert body["ignored"] is True - mock_check.assert_not_awaited()