From a8d6afbee372bb7a59ee37535f5e29e29b3fb16c Mon Sep 17 00:00:00 2001 From: David Blain Date: Thu, 7 Nov 2024 19:06:29 +0100 Subject: [PATCH 01/18] refactor: Allow callable values in path en query parameters to be specified so those don't get evaluated during the DAG processing --- .../providers/microsoft/azure/hooks/msgraph.py | 9 ++++++++- .../providers/microsoft/azure/operators/msgraph.py | 14 ++++++++++++-- .../providers/microsoft/azure/sensors/msgraph.py | 14 ++++++++++++-- providers/tests/microsoft/azure/base.py | 1 + .../tests/microsoft/azure/hooks/test_msgraph.py | 6 ++++++ .../microsoft/azure/operators/test_msgraph.py | 4 +++- .../tests/microsoft/azure/sensors/test_msgraph.py | 3 ++- 7 files changed, 44 insertions(+), 7 deletions(-) diff --git a/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py b/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py index 4ab3aaf3ba37f..e97fe4d56c92e 100644 --- a/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py @@ -22,7 +22,7 @@ from http import HTTPStatus from io import BytesIO from json import JSONDecodeError -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Callable from urllib.parse import quote, urljoin, urlparse import httpx @@ -358,3 +358,10 @@ def error_mapping() -> dict[str, ParsableFactory | None]: "4XX": APIError, "5XX": APIError, } + + @staticmethod + def evaluate_parameters(parameters: dict[str, Any | Callable[[], Any]]): + if parameters: + for key, value in parameters.items(): + if callable(value): + parameters[key] = value() diff --git a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py index 0d187ebd5144b..c1aef2aafdfc3 100644 --- a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py @@ -36,6 +36,7 @@ if TYPE_CHECKING: from io import BytesIO + import jinja2 # Slow import. from kiota_abstractions.request_adapter import ResponseType from kiota_abstractions.request_information import QueryParams @@ -100,10 +101,10 @@ def __init__( *, url: str, response_type: ResponseType | None = None, - path_parameters: dict[str, Any] | None = None, + path_parameters: dict[str, Any | Callable[[], Any]] | None = None, url_template: str | None = None, method: str = "GET", - query_parameters: dict[str, QueryParams] | None = None, + query_parameters: dict[str, QueryParams | Callable[[], QueryParams]] | None = None, headers: dict[str, str] | None = None, data: dict[str, Any] | str | BytesIO | None = None, conn_id: str = KiotaRequestAdapterHook.default_conn_name, @@ -136,6 +137,15 @@ def __init__( self.event_handler = event_handler or default_event_handler self.serializer: ResponseSerializer = serializer() + def render_template_fields( + self, + context: Context, + jinja_env: jinja2.Environment | None = None, + ) -> None: + super().render_template_fields(context=context, jinja_env=jinja_env) + KiotaRequestAdapterHook.evaluate_parameters(self.path_parameters) + KiotaRequestAdapterHook.evaluate_parameters(self.query_parameters) + def execute(self, context: Context) -> None: self.defer( trigger=MSGraphTrigger( diff --git a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py index 6736ea59c918d..2c96df2e70c30 100644 --- a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py @@ -28,6 +28,7 @@ if TYPE_CHECKING: from datetime import timedelta from io import BytesIO + import jinja2 # Slow import. from kiota_abstractions.request_information import QueryParams from kiota_http.httpx_request_adapter import ResponseType @@ -74,10 +75,10 @@ def __init__( self, url: str, response_type: ResponseType | None = None, - path_parameters: dict[str, Any] | None = None, + path_parameters: dict[str, Any | Callable[[], Any]] | None = None, url_template: str | None = None, method: str = "GET", - query_parameters: dict[str, QueryParams] | None = None, + query_parameters: dict[str, QueryParams | Callable[[], QueryParams]] | None = None, headers: dict[str, str] | None = None, data: dict[str, Any] | str | BytesIO | None = None, conn_id: str = KiotaRequestAdapterHook.default_conn_name, @@ -105,6 +106,15 @@ def __init__( self.result_processor = result_processor self.serializer = serializer() + def render_template_fields( + self, + context: Context, + jinja_env: jinja2.Environment | None = None, + ) -> None: + super().render_template_fields(context=context, jinja_env=jinja_env) + KiotaRequestAdapterHook.evaluate_parameters(self.path_parameters) + KiotaRequestAdapterHook.evaluate_parameters(self.query_parameters) + def execute(self, context: Context): self.defer( trigger=MSGraphTrigger( diff --git a/providers/tests/microsoft/azure/base.py b/providers/tests/microsoft/azure/base.py index 98c0a59867ea4..600e4ce488e08 100644 --- a/providers/tests/microsoft/azure/base.py +++ b/providers/tests/microsoft/azure/base.py @@ -68,6 +68,7 @@ async def deferrable_operator(self, context, operator): result = None triggered_events = [] try: + operator.render_template_fields(context=context) result = operator.execute(context=context) except TaskDeferred as deferred: task = deferred diff --git a/providers/tests/microsoft/azure/hooks/test_msgraph.py b/providers/tests/microsoft/azure/hooks/test_msgraph.py index aff5d0226a1c4..8610b1738c914 100644 --- a/providers/tests/microsoft/azure/hooks/test_msgraph.py +++ b/providers/tests/microsoft/azure/hooks/test_msgraph.py @@ -218,6 +218,12 @@ async def test_throw_failed_responses_with_application_json_content_type(self): error_code = actual.get_child_node("error").get_child_node("code").get_str_value() assert error_code == "TenantThrottleThresholdExceeded" + def test_evaluate_parameters(self): + query_parameters = {"$expand": lambda: ",".join(["reports","users","datasets","dataflows","dashboards"]), "$top": 5000} + KiotaRequestAdapterHook.evaluate_parameters(query_parameters) + + assert query_parameters == {"$expand": "reports,users,datasets,dataflows,dashboards", "$top": 5000} + class TestResponseHandler: def test_default_response_handler_when_json(self): diff --git a/providers/tests/microsoft/azure/operators/test_msgraph.py b/providers/tests/microsoft/azure/operators/test_msgraph.py index fe404e48e6f0a..4a19ec852c449 100644 --- a/providers/tests/microsoft/azure/operators/test_msgraph.py +++ b/providers/tests/microsoft/azure/operators/test_msgraph.py @@ -143,11 +143,13 @@ def test_execute_when_response_is_bytes(self): task_id="drive_item_content", conn_id="msgraph_api", response_type="bytes", - url=f"/drives/{drive_id}/root/content", + url="/drives/{drive_id}/root/content", + path_parameters={"drive_id": lambda: drive_id}, ) results, events = self.execute_operator(operator) + assert operator.path_parameters == {"drive_id": drive_id} assert results == base64_encoded_content assert len(events) == 1 assert isinstance(events[0], TriggerEvent) diff --git a/providers/tests/microsoft/azure/sensors/test_msgraph.py b/providers/tests/microsoft/azure/sensors/test_msgraph.py index ba5ba35478861..8a4d5e4594439 100644 --- a/providers/tests/microsoft/azure/sensors/test_msgraph.py +++ b/providers/tests/microsoft/azure/sensors/test_msgraph.py @@ -35,13 +35,14 @@ def test_execute(self): task_id="check_workspaces_status", conn_id="powerbi", url="myorg/admin/workspaces/scanStatus/{scanId}", - path_parameters={"scanId": "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef"}, + path_parameters={"scanId": lambda: "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef"}, result_processor=lambda context, result: result["id"], timeout=350.0, ) results, events = self.execute_operator(sensor) + assert sensor.path_parameters == {"scanId": "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef"} assert isinstance(results, str) assert results == "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef" assert len(events) == 1 From eb5fdb196f4ef73e6b4bd4f768f98d977d51fefd Mon Sep 17 00:00:00 2001 From: David Blain Date: Thu, 7 Nov 2024 19:56:44 +0100 Subject: [PATCH 02/18] refactor: Re-organized imports --- .../src/airflow/providers/microsoft/azure/operators/msgraph.py | 2 +- .../src/airflow/providers/microsoft/azure/sensors/msgraph.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py index c1aef2aafdfc3..5a4243dc68d28 100644 --- a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py @@ -36,8 +36,8 @@ if TYPE_CHECKING: from io import BytesIO - import jinja2 # Slow import. + import jinja2 # Slow import. from kiota_abstractions.request_adapter import ResponseType from kiota_abstractions.request_information import QueryParams from msgraph_core import APIVersion diff --git a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py index 2c96df2e70c30..609d944d79db2 100644 --- a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py @@ -28,8 +28,8 @@ if TYPE_CHECKING: from datetime import timedelta from io import BytesIO - import jinja2 # Slow import. + import jinja2 # Slow import. from kiota_abstractions.request_information import QueryParams from kiota_http.httpx_request_adapter import ResponseType from msgraph_core import APIVersion From f48b8dae1b90619dcd7fbe9f775e2893b15c6d1e Mon Sep 17 00:00:00 2001 From: David Blain Date: Thu, 7 Nov 2024 19:57:23 +0100 Subject: [PATCH 03/18] refactor: Reformatted test_evaluate_parameters --- providers/tests/microsoft/azure/hooks/test_msgraph.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/providers/tests/microsoft/azure/hooks/test_msgraph.py b/providers/tests/microsoft/azure/hooks/test_msgraph.py index 8610b1738c914..368780b350341 100644 --- a/providers/tests/microsoft/azure/hooks/test_msgraph.py +++ b/providers/tests/microsoft/azure/hooks/test_msgraph.py @@ -219,7 +219,10 @@ async def test_throw_failed_responses_with_application_json_content_type(self): assert error_code == "TenantThrottleThresholdExceeded" def test_evaluate_parameters(self): - query_parameters = {"$expand": lambda: ",".join(["reports","users","datasets","dataflows","dashboards"]), "$top": 5000} + query_parameters = { + "$expand": lambda: ",".join(["reports","users","datasets","dataflows","dashboards"]), + "$top": 5000, + } KiotaRequestAdapterHook.evaluate_parameters(query_parameters) assert query_parameters == {"$expand": "reports,users,datasets,dataflows,dashboards", "$top": 5000} From dadb52a522aef6980827720e3789fa354422585d Mon Sep 17 00:00:00 2001 From: David Blain Date: Thu, 7 Nov 2024 19:59:02 +0100 Subject: [PATCH 04/18] refactor: parameters of evaluate_parameters is optional --- .../src/airflow/providers/microsoft/azure/hooks/msgraph.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py b/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py index e97fe4d56c92e..e3febd8e80edf 100644 --- a/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py @@ -360,7 +360,7 @@ def error_mapping() -> dict[str, ParsableFactory | None]: } @staticmethod - def evaluate_parameters(parameters: dict[str, Any | Callable[[], Any]]): + def evaluate_parameters(parameters: dict[str, Any | Callable[[], Any]] | None): if parameters: for key, value in parameters.items(): if callable(value): From c36680a9e5a8901a95d02264b1310b4a59a68631 Mon Sep 17 00:00:00 2001 From: David Blain Date: Fri, 8 Nov 2024 07:59:40 +0100 Subject: [PATCH 05/18] refactor: Reformatted test_evaluate_parameters --- providers/tests/microsoft/azure/hooks/test_msgraph.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/tests/microsoft/azure/hooks/test_msgraph.py b/providers/tests/microsoft/azure/hooks/test_msgraph.py index 368780b350341..27d3fdb852674 100644 --- a/providers/tests/microsoft/azure/hooks/test_msgraph.py +++ b/providers/tests/microsoft/azure/hooks/test_msgraph.py @@ -220,7 +220,7 @@ async def test_throw_failed_responses_with_application_json_content_type(self): def test_evaluate_parameters(self): query_parameters = { - "$expand": lambda: ",".join(["reports","users","datasets","dataflows","dashboards"]), + "$expand": lambda: ",".join(["reports", "users", "datasets", "dataflows", "dashboards"]), "$top": 5000, } KiotaRequestAdapterHook.evaluate_parameters(query_parameters) From b038c440b5845ea8769bc9607c0c4d399273ebaf Mon Sep 17 00:00:00 2001 From: David Blain Date: Fri, 8 Nov 2024 08:03:35 +0100 Subject: [PATCH 06/18] refactor: Ignore mypy error --- .../src/airflow/providers/microsoft/azure/operators/msgraph.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py index 5a4243dc68d28..60a6c9c983431 100644 --- a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py @@ -290,7 +290,7 @@ def paginate( if top and odata_count: if len(response.get("value", [])) == top and context: results = operator.pull_xcom(context=context) - skip = sum(map(lambda result: len(result["value"]), results)) + top if results else top + skip = sum(map(lambda result: len(result["value"]), results)) + top if results else top # type: ignore query_parameters["$skip"] = skip return operator.url, query_parameters return response.get("@odata.nextLink"), operator.query_parameters From 78a6dfdff32b6e0a14a58931d2ae989e4aa4ec8b Mon Sep 17 00:00:00 2001 From: David Blain Date: Fri, 8 Nov 2024 08:04:21 +0100 Subject: [PATCH 07/18] refactor: Use list comprehension to calculate total length --- .../src/airflow/providers/microsoft/azure/operators/msgraph.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py index 60a6c9c983431..1f7ee03cbf282 100644 --- a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py @@ -290,7 +290,7 @@ def paginate( if top and odata_count: if len(response.get("value", [])) == top and context: results = operator.pull_xcom(context=context) - skip = sum(map(lambda result: len(result["value"]), results)) + top if results else top # type: ignore + skip = sum([len(result["value"]) for result in results]) + top if results else top # type: ignore query_parameters["$skip"] = skip return operator.url, query_parameters return response.get("@odata.nextLink"), operator.query_parameters From 3ce3370190f546963bc2316af3f645666a303a24 Mon Sep 17 00:00:00 2001 From: David Blain Date: Tue, 12 Nov 2024 09:29:26 +0100 Subject: [PATCH 08/18] refactor: Removed evaluate_parameters method from KiotaRequestAdapterHook --- .../microsoft/azure/hooks/msgraph.py | 14 +++---------- .../microsoft/azure/operators/msgraph.py | 9 --------- .../microsoft/azure/sensors/msgraph.py | 9 --------- .../microsoft/azure/hooks/test_msgraph.py | 20 +++++-------------- .../microsoft/azure/operators/test_msgraph.py | 2 +- .../microsoft/azure/sensors/test_msgraph.py | 2 +- 6 files changed, 10 insertions(+), 46 deletions(-) diff --git a/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py b/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py index e3febd8e80edf..a7c50f974ffea 100644 --- a/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py @@ -22,10 +22,12 @@ from http import HTTPStatus from io import BytesIO from json import JSONDecodeError -from typing import TYPE_CHECKING, Any, Callable +from typing import TYPE_CHECKING, Any from urllib.parse import quote, urljoin, urlparse import httpx +from airflow.exceptions import AirflowBadRequest, AirflowException, AirflowNotFoundException +from airflow.hooks.base import BaseHook from azure.identity import ClientSecretCredential from httpx import Timeout from kiota_abstractions.api_error import APIError @@ -43,9 +45,6 @@ from msgraph_core import APIVersion, GraphClientFactory from msgraph_core._enums import NationalClouds -from airflow.exceptions import AirflowBadRequest, AirflowException, AirflowNotFoundException -from airflow.hooks.base import BaseHook - if TYPE_CHECKING: from kiota_abstractions.request_adapter import RequestAdapter from kiota_abstractions.request_information import QueryParams @@ -358,10 +357,3 @@ def error_mapping() -> dict[str, ParsableFactory | None]: "4XX": APIError, "5XX": APIError, } - - @staticmethod - def evaluate_parameters(parameters: dict[str, Any | Callable[[], Any]] | None): - if parameters: - for key, value in parameters.items(): - if callable(value): - parameters[key] = value() diff --git a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py index 1f7ee03cbf282..8022e0af7f826 100644 --- a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py @@ -137,15 +137,6 @@ def __init__( self.event_handler = event_handler or default_event_handler self.serializer: ResponseSerializer = serializer() - def render_template_fields( - self, - context: Context, - jinja_env: jinja2.Environment | None = None, - ) -> None: - super().render_template_fields(context=context, jinja_env=jinja_env) - KiotaRequestAdapterHook.evaluate_parameters(self.path_parameters) - KiotaRequestAdapterHook.evaluate_parameters(self.query_parameters) - def execute(self, context: Context) -> None: self.defer( trigger=MSGraphTrigger( diff --git a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py index 609d944d79db2..885ef0217ce44 100644 --- a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py @@ -106,15 +106,6 @@ def __init__( self.result_processor = result_processor self.serializer = serializer() - def render_template_fields( - self, - context: Context, - jinja_env: jinja2.Environment | None = None, - ) -> None: - super().render_template_fields(context=context, jinja_env=jinja_env) - KiotaRequestAdapterHook.evaluate_parameters(self.path_parameters) - KiotaRequestAdapterHook.evaluate_parameters(self.query_parameters) - def execute(self, context: Context): self.defer( trigger=MSGraphTrigger( diff --git a/providers/tests/microsoft/azure/hooks/test_msgraph.py b/providers/tests/microsoft/azure/hooks/test_msgraph.py index 27d3fdb852674..ace567187797d 100644 --- a/providers/tests/microsoft/azure/hooks/test_msgraph.py +++ b/providers/tests/microsoft/azure/hooks/test_msgraph.py @@ -22,6 +22,11 @@ from unittest.mock import Mock, patch import pytest +from airflow.exceptions import AirflowBadRequest, AirflowException, AirflowNotFoundException +from airflow.providers.microsoft.azure.hooks.msgraph import ( + DefaultResponseHandler, + KiotaRequestAdapterHook, +) from httpx import Response from kiota_http.httpx_request_adapter import HttpxRequestAdapter from kiota_serialization_json.json_parse_node import JsonParseNode @@ -29,12 +34,6 @@ from msgraph_core import APIVersion, NationalClouds from opentelemetry.trace import Span -from airflow.exceptions import AirflowBadRequest, AirflowException, AirflowNotFoundException -from airflow.providers.microsoft.azure.hooks.msgraph import ( - DefaultResponseHandler, - KiotaRequestAdapterHook, -) - from providers.tests.microsoft.conftest import ( get_airflow_connection, load_file, @@ -218,15 +217,6 @@ async def test_throw_failed_responses_with_application_json_content_type(self): error_code = actual.get_child_node("error").get_child_node("code").get_str_value() assert error_code == "TenantThrottleThresholdExceeded" - def test_evaluate_parameters(self): - query_parameters = { - "$expand": lambda: ",".join(["reports", "users", "datasets", "dataflows", "dashboards"]), - "$top": 5000, - } - KiotaRequestAdapterHook.evaluate_parameters(query_parameters) - - assert query_parameters == {"$expand": "reports,users,datasets,dataflows,dashboards", "$top": 5000} - class TestResponseHandler: def test_default_response_handler_when_json(self): diff --git a/providers/tests/microsoft/azure/operators/test_msgraph.py b/providers/tests/microsoft/azure/operators/test_msgraph.py index 4a19ec852c449..82927742b0387 100644 --- a/providers/tests/microsoft/azure/operators/test_msgraph.py +++ b/providers/tests/microsoft/azure/operators/test_msgraph.py @@ -144,7 +144,7 @@ def test_execute_when_response_is_bytes(self): conn_id="msgraph_api", response_type="bytes", url="/drives/{drive_id}/root/content", - path_parameters={"drive_id": lambda: drive_id}, + path_parameters=lambda context, jinja_env: {"drive_id": drive_id}, ) results, events = self.execute_operator(operator) diff --git a/providers/tests/microsoft/azure/sensors/test_msgraph.py b/providers/tests/microsoft/azure/sensors/test_msgraph.py index 8a4d5e4594439..a3ffa271b7fa3 100644 --- a/providers/tests/microsoft/azure/sensors/test_msgraph.py +++ b/providers/tests/microsoft/azure/sensors/test_msgraph.py @@ -35,7 +35,7 @@ def test_execute(self): task_id="check_workspaces_status", conn_id="powerbi", url="myorg/admin/workspaces/scanStatus/{scanId}", - path_parameters={"scanId": lambda: "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef"}, + path_parameters=lambda context, jinja_env: {"scanId": "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef"}, result_processor=lambda context, result: result["id"], timeout=350.0, ) From 44410eacc09cdd865ed6ab347c894f1242b2aec1 Mon Sep 17 00:00:00 2001 From: David Blain Date: Tue, 12 Nov 2024 09:31:40 +0100 Subject: [PATCH 09/18] refactor: Reorganized imports in KiotaRequestAdapterHook --- .../src/airflow/providers/microsoft/azure/hooks/msgraph.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py b/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py index a7c50f974ffea..4ab3aaf3ba37f 100644 --- a/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/hooks/msgraph.py @@ -26,8 +26,6 @@ from urllib.parse import quote, urljoin, urlparse import httpx -from airflow.exceptions import AirflowBadRequest, AirflowException, AirflowNotFoundException -from airflow.hooks.base import BaseHook from azure.identity import ClientSecretCredential from httpx import Timeout from kiota_abstractions.api_error import APIError @@ -45,6 +43,9 @@ from msgraph_core import APIVersion, GraphClientFactory from msgraph_core._enums import NationalClouds +from airflow.exceptions import AirflowBadRequest, AirflowException, AirflowNotFoundException +from airflow.hooks.base import BaseHook + if TYPE_CHECKING: from kiota_abstractions.request_adapter import RequestAdapter from kiota_abstractions.request_information import QueryParams From c21adc3e0dd27a67379a0fed185bacc0d81c2b14 Mon Sep 17 00:00:00 2001 From: David Blain Date: Tue, 12 Nov 2024 09:32:18 +0100 Subject: [PATCH 10/18] refactor: Reorganized imports in TestKiotaRequestAdapterHook --- providers/tests/microsoft/azure/hooks/test_msgraph.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/providers/tests/microsoft/azure/hooks/test_msgraph.py b/providers/tests/microsoft/azure/hooks/test_msgraph.py index ace567187797d..aff5d0226a1c4 100644 --- a/providers/tests/microsoft/azure/hooks/test_msgraph.py +++ b/providers/tests/microsoft/azure/hooks/test_msgraph.py @@ -22,11 +22,6 @@ from unittest.mock import Mock, patch import pytest -from airflow.exceptions import AirflowBadRequest, AirflowException, AirflowNotFoundException -from airflow.providers.microsoft.azure.hooks.msgraph import ( - DefaultResponseHandler, - KiotaRequestAdapterHook, -) from httpx import Response from kiota_http.httpx_request_adapter import HttpxRequestAdapter from kiota_serialization_json.json_parse_node import JsonParseNode @@ -34,6 +29,12 @@ from msgraph_core import APIVersion, NationalClouds from opentelemetry.trace import Span +from airflow.exceptions import AirflowBadRequest, AirflowException, AirflowNotFoundException +from airflow.providers.microsoft.azure.hooks.msgraph import ( + DefaultResponseHandler, + KiotaRequestAdapterHook, +) + from providers.tests.microsoft.conftest import ( get_airflow_connection, load_file, From 3d422dd0a79705fb03bed55dfee201678a6ff457 Mon Sep 17 00:00:00 2001 From: David Blain Date: Tue, 12 Nov 2024 09:33:48 +0100 Subject: [PATCH 11/18] refactor: Removed obsolete jinja2 import from MSGraphAsyncOperator and MSGraphSensor --- .../src/airflow/providers/microsoft/azure/operators/msgraph.py | 1 - .../src/airflow/providers/microsoft/azure/sensors/msgraph.py | 1 - 2 files changed, 2 deletions(-) diff --git a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py index 8022e0af7f826..27ea47a654c7e 100644 --- a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py @@ -37,7 +37,6 @@ if TYPE_CHECKING: from io import BytesIO - import jinja2 # Slow import. from kiota_abstractions.request_adapter import ResponseType from kiota_abstractions.request_information import QueryParams from msgraph_core import APIVersion diff --git a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py index 885ef0217ce44..d086b3b44e796 100644 --- a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py @@ -29,7 +29,6 @@ from datetime import timedelta from io import BytesIO - import jinja2 # Slow import. from kiota_abstractions.request_information import QueryParams from kiota_http.httpx_request_adapter import ResponseType from msgraph_core import APIVersion From e5f84116573a686095611220c530574892344606 Mon Sep 17 00:00:00 2001 From: David Blain Date: Tue, 12 Nov 2024 09:39:09 +0100 Subject: [PATCH 12/18] refactor: Removed Callable types from path_parameters and query_parameters --- .../airflow/providers/microsoft/azure/operators/msgraph.py | 4 ++-- .../src/airflow/providers/microsoft/azure/sensors/msgraph.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py index 27ea47a654c7e..57ebf6504c0de 100644 --- a/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/operators/msgraph.py @@ -100,10 +100,10 @@ def __init__( *, url: str, response_type: ResponseType | None = None, - path_parameters: dict[str, Any | Callable[[], Any]] | None = None, + path_parameters: dict[str, Any] | None = None, url_template: str | None = None, method: str = "GET", - query_parameters: dict[str, QueryParams | Callable[[], QueryParams]] | None = None, + query_parameters: dict[str, QueryParams] | None = None, headers: dict[str, str] | None = None, data: dict[str, Any] | str | BytesIO | None = None, conn_id: str = KiotaRequestAdapterHook.default_conn_name, diff --git a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py index d086b3b44e796..6736ea59c918d 100644 --- a/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py +++ b/providers/src/airflow/providers/microsoft/azure/sensors/msgraph.py @@ -74,10 +74,10 @@ def __init__( self, url: str, response_type: ResponseType | None = None, - path_parameters: dict[str, Any | Callable[[], Any]] | None = None, + path_parameters: dict[str, Any] | None = None, url_template: str | None = None, method: str = "GET", - query_parameters: dict[str, QueryParams | Callable[[], QueryParams]] | None = None, + query_parameters: dict[str, QueryParams] | None = None, headers: dict[str, str] | None = None, data: dict[str, Any] | str | BytesIO | None = None, conn_id: str = KiotaRequestAdapterHook.default_conn_name, From e70c8528509880de0a61e0544b095f76b5485114 Mon Sep 17 00:00:00 2001 From: David Blain Date: Tue, 12 Nov 2024 11:51:40 +0100 Subject: [PATCH 13/18] refactor: Changed type for xcom_pull method of MockedTaskInstance --- providers/tests/microsoft/conftest.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/providers/tests/microsoft/conftest.py b/providers/tests/microsoft/conftest.py index de25d24fb05e1..bf6a291ee9dbc 100644 --- a/providers/tests/microsoft/conftest.py +++ b/providers/tests/microsoft/conftest.py @@ -132,14 +132,15 @@ def __init__( def xcom_pull( self, - task_ids: Iterable[str] | str | None = None, + task_ids: str | Iterable[str] | None = None, dag_id: str | None = None, key: str = XCOM_RETURN_KEY, include_prior_dates: bool = False, session: Session = NEW_SESSION, *, - map_indexes: Iterable[int] | int | None = None, - default: Any | None = None, + map_indexes: int | Iterable[int] | None = None, + default: Any = None, + run_id: str | None = None, ) -> Any: if map_indexes: return values.get(f"{task_ids or self.task_id}_{dag_id or self.dag_id}_{key}_{map_indexes}") From 86ac4279b69f5ed40d6b74a55bf5c173191c02d3 Mon Sep 17 00:00:00 2001 From: David Blain Date: Tue, 12 Nov 2024 14:30:10 +0100 Subject: [PATCH 14/18] refactor: Removed duplicate run_id in xcom_pull method of MockedTaskInstance --- providers/tests/microsoft/conftest.py | 1 - 1 file changed, 1 deletion(-) diff --git a/providers/tests/microsoft/conftest.py b/providers/tests/microsoft/conftest.py index 40383cdd9e10a..bf6a291ee9dbc 100644 --- a/providers/tests/microsoft/conftest.py +++ b/providers/tests/microsoft/conftest.py @@ -137,7 +137,6 @@ def xcom_pull( key: str = XCOM_RETURN_KEY, include_prior_dates: bool = False, session: Session = NEW_SESSION, - run_id: str | None = None, *, map_indexes: int | Iterable[int] | None = None, default: Any = None, From 483d5826b497b542bc9360ae50147b7190a4db37 Mon Sep 17 00:00:00 2001 From: David Blain Date: Wed, 13 Nov 2024 07:54:26 +0100 Subject: [PATCH 15/18] refactor: Only test operator and sensor with lambda parameter when Airflow is 2.10.0 or higher --- .../microsoft/azure/operators/test_msgraph.py | 29 +++++++++++++++++++ .../microsoft/azure/sensors/test_msgraph.py | 29 +++++++++++++++++++ 2 files changed, 58 insertions(+) diff --git a/providers/tests/microsoft/azure/operators/test_msgraph.py b/providers/tests/microsoft/azure/operators/test_msgraph.py index 82927742b0387..fd6543935c685 100644 --- a/providers/tests/microsoft/azure/operators/test_msgraph.py +++ b/providers/tests/microsoft/azure/operators/test_msgraph.py @@ -36,6 +36,8 @@ mock_response, ) +from tests_common.test_utils.compat import AIRFLOW_V_2_10_PLUS + if TYPE_CHECKING: from airflow.utils.context import Context @@ -138,6 +140,33 @@ def test_execute_when_response_is_bytes(self): drive_id = "82f9d24d-6891-4790-8b6d-f1b2a1d0ca22" response = mock_response(200, content) + with self.patch_hook_and_request_adapter(response): + operator = MSGraphAsyncOperator( + task_id="drive_item_content", + conn_id="msgraph_api", + response_type="bytes", + url="/drives/{drive_id}/root/content", + path_parameters={"drive_id": drive_id}, + ) + + results, events = self.execute_operator(operator) + + assert operator.path_parameters == {"drive_id": drive_id} + assert results == base64_encoded_content + assert len(events) == 1 + assert isinstance(events[0], TriggerEvent) + assert events[0].payload["status"] == "success" + assert events[0].payload["type"] == "builtins.bytes" + assert events[0].payload["response"] == base64_encoded_content + + @pytest.mark.db_test + @pytest.mark.skipif(not AIRFLOW_V_2_10_PLUS, reason="Lambda parameters works in Airflow >= 2.10.0") + def test_execute_with_lambda_parameter_when_response_is_bytes(self): + content = load_file("resources", "dummy.pdf", mode="rb", encoding=None) + base64_encoded_content = b64encode(content).decode(locale.getpreferredencoding()) + drive_id = "82f9d24d-6891-4790-8b6d-f1b2a1d0ca22" + response = mock_response(200, content) + with self.patch_hook_and_request_adapter(response): operator = MSGraphAsyncOperator( task_id="drive_item_content", diff --git a/providers/tests/microsoft/azure/sensors/test_msgraph.py b/providers/tests/microsoft/azure/sensors/test_msgraph.py index a3ffa271b7fa3..f0bfb8d77f387 100644 --- a/providers/tests/microsoft/azure/sensors/test_msgraph.py +++ b/providers/tests/microsoft/azure/sensors/test_msgraph.py @@ -18,18 +18,47 @@ import json +import pytest from airflow.providers.microsoft.azure.sensors.msgraph import MSGraphSensor from airflow.triggers.base import TriggerEvent from providers.tests.microsoft.azure.base import Base from providers.tests.microsoft.conftest import load_json, mock_json_response +from tests_common.test_utils.compat import AIRFLOW_V_2_10_PLUS + class TestMSGraphSensor(Base): def test_execute(self): status = load_json("resources", "status.json") response = mock_json_response(200, status) + with self.patch_hook_and_request_adapter(response): + sensor = MSGraphSensor( + task_id="check_workspaces_status", + conn_id="powerbi", + url="myorg/admin/workspaces/scanStatus/{scanId}", + path_parameters={"scanId": "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef"}, + result_processor=lambda context, result: result["id"], + timeout=350.0, + ) + + results, events = self.execute_operator(sensor) + + assert sensor.path_parameters == {"scanId": "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef"} + assert isinstance(results, str) + assert results == "0a1b1bf3-37de-48f7-9863-ed4cda97a9ef" + assert len(events) == 1 + assert isinstance(events[0], TriggerEvent) + assert events[0].payload["status"] == "success" + assert events[0].payload["type"] == "builtins.dict" + assert events[0].payload["response"] == json.dumps(status) + + @pytest.mark.skipif(not AIRFLOW_V_2_10_PLUS, reason="Lambda parameters works in Airflow >= 2.10.0") + def test_execute_with_lambda_parameter(self): + status = load_json("resources", "status.json") + response = mock_json_response(200, status) + with self.patch_hook_and_request_adapter(response): sensor = MSGraphSensor( task_id="check_workspaces_status", From beda59a191dfc3a6c6bea18c6c7a7eeff8e522ba Mon Sep 17 00:00:00 2001 From: David Blain Date: Wed, 13 Nov 2024 09:02:06 +0100 Subject: [PATCH 16/18] refactor: Reorganized imports TestMSGraphSensor --- providers/tests/microsoft/azure/sensors/test_msgraph.py | 1 - 1 file changed, 1 deletion(-) diff --git a/providers/tests/microsoft/azure/sensors/test_msgraph.py b/providers/tests/microsoft/azure/sensors/test_msgraph.py index f0bfb8d77f387..074b0ffb36747 100644 --- a/providers/tests/microsoft/azure/sensors/test_msgraph.py +++ b/providers/tests/microsoft/azure/sensors/test_msgraph.py @@ -24,7 +24,6 @@ from providers.tests.microsoft.azure.base import Base from providers.tests.microsoft.conftest import load_json, mock_json_response - from tests_common.test_utils.compat import AIRFLOW_V_2_10_PLUS From 75442fd6adc8099a729a6dee6a14177f8904706e Mon Sep 17 00:00:00 2001 From: David Blain Date: Wed, 13 Nov 2024 09:03:10 +0100 Subject: [PATCH 17/18] refactor: Reorganized imports TestMSGraphSensor --- providers/tests/microsoft/azure/sensors/test_msgraph.py | 1 + 1 file changed, 1 insertion(+) diff --git a/providers/tests/microsoft/azure/sensors/test_msgraph.py b/providers/tests/microsoft/azure/sensors/test_msgraph.py index 074b0ffb36747..8b8ec793d65ce 100644 --- a/providers/tests/microsoft/azure/sensors/test_msgraph.py +++ b/providers/tests/microsoft/azure/sensors/test_msgraph.py @@ -19,6 +19,7 @@ import json import pytest + from airflow.providers.microsoft.azure.sensors.msgraph import MSGraphSensor from airflow.triggers.base import TriggerEvent From 1d9e24ff5ccef02b3bb6e5970011a02f6822808a Mon Sep 17 00:00:00 2001 From: David Blain Date: Wed, 13 Nov 2024 09:37:27 +0100 Subject: [PATCH 18/18] refactor: Reorganized imports TestMSGraphAsyncOperator --- providers/tests/microsoft/azure/operators/test_msgraph.py | 1 - 1 file changed, 1 deletion(-) diff --git a/providers/tests/microsoft/azure/operators/test_msgraph.py b/providers/tests/microsoft/azure/operators/test_msgraph.py index fd6543935c685..2c9c8129d5d08 100644 --- a/providers/tests/microsoft/azure/operators/test_msgraph.py +++ b/providers/tests/microsoft/azure/operators/test_msgraph.py @@ -35,7 +35,6 @@ mock_json_response, mock_response, ) - from tests_common.test_utils.compat import AIRFLOW_V_2_10_PLUS if TYPE_CHECKING: