From 584b46d2c64daeab1e71844c00eaa2474717fa56 Mon Sep 17 00:00:00 2001 From: Hemkumar Chheda Date: Thu, 14 May 2026 13:04:12 +0530 Subject: [PATCH 1/3] Avoid false trigger Dag run conflicts after ambiguous retry closes: #66905 --- task-sdk/src/airflow/sdk/api/client.py | 42 +++++++++++++--- task-sdk/tests/task_sdk/api/test_client.py | 57 ++++++++++++++++++++++ 2 files changed, 91 insertions(+), 8 deletions(-) diff --git a/task-sdk/src/airflow/sdk/api/client.py b/task-sdk/src/airflow/sdk/api/client.py index 269978ac9dd1d..0bffa6b136137 100644 --- a/task-sdk/src/airflow/sdk/api/client.py +++ b/task-sdk/src/airflow/sdk/api/client.py @@ -870,9 +870,18 @@ def trigger( ) try: - self.client.post( - f"dag-runs/{dag_id}/{run_id}", content=body.model_dump_json(exclude_defaults=True) + self.client._request_without_retry( + "POST", f"dag-runs/{dag_id}/{run_id}", content=body.model_dump_json(exclude_defaults=True) ) + except httpx.RequestError: + if not reset_dag_run and self._dag_run_exists(dag_id=dag_id, run_id=run_id): + log.info( + "Dag Run exists after ambiguous trigger response; treating trigger as successful.", + dag_id=dag_id, + run_id=run_id, + ) + return OKResponse(ok=True) + raise except ServerResponseError as e: if e.response.status_code == HTTPStatus.CONFLICT: if reset_dag_run: @@ -885,6 +894,15 @@ def trigger( return OKResponse(ok=True) + def _dag_run_exists(self, dag_id: str, run_id: str) -> bool: + try: + self.client.get(f"dag-runs/{dag_id}/{run_id}") + except ServerResponseError as e: + if e.response.status_code == HTTPStatus.NOT_FOUND: + return False + raise + return True + def clear(self, dag_id: str, run_id: str) -> OKResponse: """Clear a Dag run via the API server.""" self.client.post(f"dag-runs/{dag_id}/{run_id}/clear") @@ -1125,6 +1143,19 @@ def _update_auth(self, response: httpx.Response): log.debug("Execution API issued us a refreshed Task token") self.auth = BearerAuth(new_token) + @staticmethod + def _ensure_json_content_type(kwargs: dict[str, Any]) -> None: + # Set content type as convenience if not already set + if kwargs.get("content", None) is not None and "content-type" not in ( + kwargs.get("headers", {}) or {} + ): + kwargs["headers"] = {"content-type": "application/json"} + + def _request_without_retry(self, *args, **kwargs): + """Implement a convenience for httpx.Client.request without retrying.""" + self._ensure_json_content_type(kwargs) + return super().request(*args, **kwargs) + @retry( retry=retry_if_exception(_should_retry_api_request), stop=stop_after_attempt(API_RETRIES), @@ -1134,12 +1165,7 @@ def _update_auth(self, response: httpx.Response): ) def request(self, *args, **kwargs): """Implement a convenience for httpx.Client.request with a retry layer.""" - # Set content type as convenience if not already set - if kwargs.get("content", None) is not None and "content-type" not in ( - kwargs.get("headers", {}) or {} - ): - kwargs["headers"] = {"content-type": "application/json"} - + self._ensure_json_content_type(kwargs) return super().request(*args, **kwargs) # We "group" or "namespace" operations by what they operate on, rather than a flat namespace with all diff --git a/task-sdk/tests/task_sdk/api/test_client.py b/task-sdk/tests/task_sdk/api/test_client.py index a179ff08436b2..94f4fc278ef0c 100644 --- a/task-sdk/tests/task_sdk/api/test_client.py +++ b/task-sdk/tests/task_sdk/api/test_client.py @@ -1280,6 +1280,63 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) + def test_trigger_treats_ambiguous_request_error_as_success_when_dag_run_exists(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + raise httpx.ReadError("Trigger response was lost", request=request) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=200, json={"detail": "exists"}) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == OKResponse(ok=True) + assert requests == [ + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_reraises_ambiguous_request_error_when_dag_run_is_missing(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + raise httpx.ReadError("Trigger response was lost", request=request) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + + with pytest.raises(httpx.ReadError, match="Trigger response was lost"): + client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert requests == [ + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_reraises_ambiguous_request_error_when_resetting_dag_run(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + raise httpx.ReadError("Trigger response was lost", request=request) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + + with pytest.raises(httpx.ReadError, match="Trigger response was lost"): + client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id", reset_dag_run=True) + + assert requests == [("POST", "/dag-runs/test_trigger/test_run_id")] + def test_trigger_conflict(self): """Test that if the dag run already exists, the client returns an error when default reset_dag_run=False""" From e30cdc50d16de817763a3b04b3c73368bf8745ef Mon Sep 17 00:00:00 2001 From: Hemkumar Chheda Date: Mon, 18 May 2026 22:14:27 +0530 Subject: [PATCH 2/3] Address trigger Dag run precheck review feedback --- task-sdk/src/airflow/sdk/api/client.py | 79 +++- task-sdk/tests/task_sdk/api/test_client.py | 512 ++++++++++++++++++++- 2 files changed, 567 insertions(+), 24 deletions(-) diff --git a/task-sdk/src/airflow/sdk/api/client.py b/task-sdk/src/airflow/sdk/api/client.py index 0bffa6b136137..108718c976724 100644 --- a/task-sdk/src/airflow/sdk/api/client.py +++ b/task-sdk/src/airflow/sdk/api/client.py @@ -869,12 +869,23 @@ def trigger( run_after=run_after, ) + dag_run_exists_before_trigger = self._dag_run_exists(dag_id=dag_id, run_id=run_id) + if dag_run_exists_before_trigger is True: + if reset_dag_run: + log.info("Dag Run already exists; Resetting Dag Run.", dag_id=dag_id, run_id=run_id) + return self.clear(run_id=run_id, dag_id=dag_id) + log.info("Dag Run already exists!", dag_id=dag_id, run_id=run_id) + return ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) + try: self.client._request_without_retry( "POST", f"dag-runs/{dag_id}/{run_id}", content=body.model_dump_json(exclude_defaults=True) ) - except httpx.RequestError: - if not reset_dag_run and self._dag_run_exists(dag_id=dag_id, run_id=run_id): + except (httpx.ReadError, httpx.ReadTimeout, httpx.RemoteProtocolError): + if ( + dag_run_exists_before_trigger is False + and self._dag_run_exists(dag_id=dag_id, run_id=run_id, retry=True) is True + ): log.info( "Dag Run exists after ambiguous trigger response; treating trigger as successful.", dag_id=dag_id, @@ -885,24 +896,77 @@ def trigger( except ServerResponseError as e: if e.response.status_code == HTTPStatus.CONFLICT: if reset_dag_run: - log.info("Dag Run already exists; Resetting Dag Run.", dag_id=dag_id, run_id=run_id) + log.info( + "Dag Run already exists after trigger attempt; Resetting Dag Run.", + detail=e.detail, + dag_id=dag_id, + run_id=run_id, + ) return self.clear(run_id=run_id, dag_id=dag_id) - - log.info("Dag Run already exists!", detail=e.detail, dag_id=dag_id, run_id=run_id) + log.info( + "Dag Run already exists after trigger attempt.", + detail=e.detail, + dag_id=dag_id, + run_id=run_id, + ) return ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) raise return OKResponse(ok=True) - def _dag_run_exists(self, dag_id: str, run_id: str) -> bool: + def _dag_run_exists(self, dag_id: str, run_id: str, *, retry: bool = False) -> bool | None: + """Return whether the Dag run exists, or None when the detail endpoint is unavailable.""" + if self.client._dry_run: + return None + try: - self.client.get(f"dag-runs/{dag_id}/{run_id}") + if retry: + self.client.get(f"dag-runs/{dag_id}/{run_id}") + else: + self.client._request_without_retry("GET", f"dag-runs/{dag_id}/{run_id}") + except httpx.RequestError: + return None except ServerResponseError as e: if e.response.status_code == HTTPStatus.NOT_FOUND: return False + if e.response.status_code == HTTPStatus.METHOD_NOT_ALLOWED: + # Older execution API servers may not support the Dag run detail endpoint yet. + return None + if ( + e.response.status_code == HTTPStatus.UNPROCESSABLE_ENTITY + and run_id == "previous" + and self._is_legacy_previous_dag_run_route_response(e.response) + ): + return None + if e.response.status_code >= HTTPStatus.INTERNAL_SERVER_ERROR: + return None + raise + except httpx.HTTPStatusError as e: + if e.response.status_code == HTTPStatus.NOT_FOUND: + return False + if e.response.status_code == HTTPStatus.METHOD_NOT_ALLOWED: + return None + if e.response.status_code >= HTTPStatus.INTERNAL_SERVER_ERROR: + return None raise return True + @staticmethod + def _is_legacy_previous_dag_run_route_response(response: httpx.Response) -> bool: + """Return whether a 422 came from the legacy ``/{dag_id}/previous`` endpoint.""" + try: + detail = response.json()["detail"] + except (KeyError, ValueError, TypeError): + return False + if not isinstance(detail, list): + return False + return any( + isinstance(error, dict) + and error.get("type") == "missing" + and error.get("loc") == ["query", "logical_date"] + for error in detail + ) + def clear(self, dag_id: str, run_id: str) -> OKResponse: """Clear a Dag run via the API server.""" self.client.post(f"dag-runs/{dag_id}/{run_id}/clear") @@ -1104,6 +1168,7 @@ def __init__(self, *, base_url: str | None, dry_run: bool = False, token: str, * if (not base_url) ^ dry_run: raise ValueError(f"Can only specify one of {base_url=} or {dry_run=}") auth = BearerAuth(token) + self._dry_run: bool = dry_run if dry_run: # If dry run is requested, install a no op handler so that simple tasks can "heartbeat" using a diff --git a/task-sdk/tests/task_sdk/api/test_client.py b/task-sdk/tests/task_sdk/api/test_client.py index 94f4fc278ef0c..49c33c0ba88e0 100644 --- a/task-sdk/tests/task_sdk/api/test_client.py +++ b/task-sdk/tests/task_sdk/api/test_client.py @@ -1258,8 +1258,13 @@ def handle_request(request: httpx.Request) -> httpx.Response: class TestDagRunOperations: def test_trigger(self): - # Simulate a successful response from the server when triggering a dag run + # Simulate a successful response from the server when triggering a Dag run + requests: list[tuple[str, str]] = [] + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) if request.url.path == "/dag-runs/test_trigger/test_run_id": actual_body = json.loads(request.read()) assert actual_body["logical_date"] == "2025-01-01T00:00:00Z" @@ -1279,16 +1284,67 @@ def handle_request(request: httpx.Request) -> httpx.Response: ) assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_dry_run_skips_precheck_conflict(self): + client = make_client_w_dry_run() + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") - def test_trigger_treats_ambiguous_request_error_as_success_when_dag_run_exists(self): + assert result == OKResponse(ok=True) + + def test_trigger_pre_existing_dag_run_returns_conflict_without_posting(self): requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=200, json={"detail": "exists"}) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": - raise httpx.ReadError("Trigger response was lost", request=request) + return httpx.Response(status_code=500, json={"detail": "POST should not happen"}) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_clears_pre_existing_dag_run_without_posting_when_resetting(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": return httpx.Response(status_code=200, json={"detail": "exists"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=500, json={"detail": "POST should not happen"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id/clear": + return httpx.Response(status_code=204) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id", reset_dag_run=True) + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id/clear"), + ] + + def test_trigger_posts_when_precheck_endpoint_is_not_supported(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=405, json={"detail": "Method Not Allowed"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=204) return httpx.Response(status_code=422) client = make_client(transport=httpx.MockTransport(handle_request)) @@ -1296,19 +1352,72 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), ("POST", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_returns_conflict_after_unsupported_precheck(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=405, json={"detail": "Method Not Allowed"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response( + status_code=409, + json={ + "detail": { + "reason": "already_exists", + "message": "A Dag Run already exists for Dag test_trigger with run id test_run_id", + } + }, + ) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) + assert requests == [ ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), ] - def test_trigger_reraises_ambiguous_request_error_when_dag_run_is_missing(self): + def test_trigger_posts_when_precheck_request_error_leaves_existence_unknown(self): requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + raise httpx.ReadError("Could not check Dag run existence", request=request) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": - raise httpx.ReadError("Trigger response was lost", request=request) + return httpx.Response(status_code=204) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_reraises_ambiguous_post_when_precheck_request_error_leaves_existence_unknown(self): + requests: list[tuple[str, str]] = [] + get_attempts = 0 + + def handle_request(request: httpx.Request) -> httpx.Response: + nonlocal get_attempts + requests.append((request.method, request.url.path)) if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + get_attempts += 1 + if get_attempts == 1: + raise httpx.ReadError("Could not check Dag run existence", request=request) + return httpx.Response(status_code=200, json={"detail": "exists"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + raise httpx.ReadError("Trigger response was lost", request=request) return httpx.Response(status_code=422) client = make_client(transport=httpx.MockTransport(handle_request)) @@ -1317,15 +1426,306 @@ def handle_request(request: httpx.Request) -> httpx.Response: client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_posts_when_non_json_precheck_server_error_leaves_existence_unknown(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=503, text="Service unavailable") + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=204) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_posts_when_previous_run_id_hits_legacy_previous_endpoint(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/previous": + return httpx.Response( + status_code=422, + json={ + "detail": [ + { + "type": "missing", + "loc": ["query", "logical_date"], + "msg": "Field required", + "input": None, + } + ] + }, + ) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/previous": + return httpx.Response(status_code=204) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="previous") + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/previous"), + ("POST", "/dag-runs/test_trigger/previous"), + ] + + def test_trigger_raises_unexpected_precheck_validation_error(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/previous": + return httpx.Response( + status_code=422, + json={ + "detail": [ + { + "type": "value_error", + "loc": ["query", "state"], + "msg": "Invalid state", + "input": "not-a-state", + } + ] + }, + ) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + + with pytest.raises(ServerResponseError) as err: + client.dag_runs.trigger(dag_id="test_trigger", run_id="previous") + + assert err.value.response.status_code == 422 + assert requests == [ + ("GET", "/dag-runs/test_trigger/previous"), + ] + + def test_trigger_treats_read_error_as_success_when_dag_run_appears_after_missing_precheck(self): + requests: list[tuple[str, str]] = [] + dag_run_exists = False + + def handle_request(request: httpx.Request) -> httpx.Response: + nonlocal dag_run_exists + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + if dag_run_exists: + return httpx.Response(status_code=200, json={"detail": "exists"}) + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + dag_run_exists = True + raise httpx.ReadError("Trigger response was lost", request=request) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_treats_read_timeout_as_success_when_dag_run_appears_after_missing_precheck(self): + requests: list[tuple[str, str]] = [] + dag_run_exists = False + + def handle_request(request: httpx.Request) -> httpx.Response: + nonlocal dag_run_exists + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + if dag_run_exists: + return httpx.Response(status_code=200, json={"detail": "exists"}) + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + dag_run_exists = True + raise httpx.ReadTimeout("Trigger response timed out", request=request) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_retries_followup_probe_after_ambiguous_response(self): + requests: list[tuple[str, str]] = [] + dag_run_exists = False + followup_attempts = 0 + + with time_machine.travel("2023-01-01T00:00:00Z", tick=False): + + def handle_request(request: httpx.Request) -> httpx.Response: + nonlocal dag_run_exists, followup_attempts + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + if not dag_run_exists: + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + followup_attempts += 1 + if followup_attempts == 1: + return httpx.Response(status_code=500, json={"detail": "Internal Server Error"}) + return httpx.Response(status_code=200, json={"detail": "exists"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + dag_run_exists = True + raise httpx.ReadError("Trigger response was lost", request=request) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_treats_read_error_as_success_when_resetting_missing_dag_run_appears(self): + requests: list[tuple[str, str]] = [] + dag_run_exists = False + + def handle_request(request: httpx.Request) -> httpx.Response: + nonlocal dag_run_exists + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + if dag_run_exists: + return httpx.Response(status_code=200, json={"detail": "exists"}) + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + dag_run_exists = True + raise httpx.ReadError("Trigger response was lost", request=request) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger( + dag_id="test_trigger", + run_id="test_run_id", + reset_dag_run=True, + ) + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_treats_remote_protocol_error_as_success_when_dag_run_appears_after_missing_precheck( + self, + ): + requests: list[tuple[str, str]] = [] + dag_run_exists = False + + def handle_request(request: httpx.Request) -> httpx.Response: + nonlocal dag_run_exists + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + if dag_run_exists: + return httpx.Response(status_code=200, json={"detail": "exists"}) + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + dag_run_exists = True + raise httpx.RemoteProtocolError("Trigger response was lost", request=request) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_returns_conflict_after_missing_precheck(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response( + status_code=409, + json={ + "detail": { + "reason": "already_exists", + "message": "A Dag Run already exists for Dag test_trigger with run id test_run_id", + } + }, + ) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert result == ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), ("POST", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_conflict_after_unsupported_precheck_reset_dag_run(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=405, json={"detail": "Method Not Allowed"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response( + status_code=409, + json={ + "detail": { + "reason": "already_exists", + "message": "A Dag Run already exists for Dag test_trigger with run id test_run_id", + } + }, + ) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id/clear": + return httpx.Response(status_code=204) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + result = client.dag_runs.trigger( + dag_id="test_trigger", + run_id="test_run_id", + reset_dag_run=True, + ) + + assert result == OKResponse(ok=True) + assert requests == [ ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id/clear"), ] - def test_trigger_reraises_ambiguous_request_error_when_resetting_dag_run(self): + def test_trigger_reraises_read_error_when_dag_run_is_missing_after_missing_precheck(self): requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": raise httpx.ReadError("Trigger response was lost", request=request) return httpx.Response(status_code=422) @@ -1333,21 +1733,75 @@ def handle_request(request: httpx.Request) -> httpx.Response: client = make_client(transport=httpx.MockTransport(handle_request)) with pytest.raises(httpx.ReadError, match="Trigger response was lost"): - client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id", reset_dag_run=True) + client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") - assert requests == [("POST", "/dag-runs/test_trigger/test_run_id")] + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_reraises_connect_error_even_if_dag_run_exists_after_missing_precheck(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + raise httpx.ConnectError("Could not connect", request=request) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + + with pytest.raises(httpx.ConnectError, match="Could not connect"): + client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ] + + def test_trigger_reraises_pool_timeout_even_if_dag_run_exists_after_missing_precheck(self): + requests: list[tuple[str, str]] = [] + + def handle_request(request: httpx.Request) -> httpx.Response: + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": + raise httpx.PoolTimeout("Could not get a connection from the pool", request=request) + return httpx.Response(status_code=422) + + client = make_client(transport=httpx.MockTransport(handle_request)) + + with pytest.raises(httpx.PoolTimeout, match="Could not get a connection from the pool"): + client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") + + assert requests == [ + ("GET", "/dag-runs/test_trigger/test_run_id"), + ("POST", "/dag-runs/test_trigger/test_run_id"), + ] def test_trigger_conflict(self): - """Test that if the dag run already exists, the client returns an error when default reset_dag_run=False""" + """Test that if the Dag run already exists, the client returns an error when default reset_dag_run=False""" + + requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: - if request.url.path == "/dag-runs/test_trigger_conflict/test_run_id": + requests.append((request.method, request.url.path)) + if request.method == "GET" and request.url.path == "/dag-runs/test_trigger_conflict/test_run_id": + return httpx.Response(status_code=200, json={"detail": "exists"}) + if request.method == "POST" and request.url.path == "/dag-runs/test_trigger_conflict/test_run_id": return httpx.Response( status_code=409, json={ "detail": { "reason": "already_exists", - "message": "A Dag Run already exists for Dag test_trigger_conflict with run id test_run_id", + "message": ( + "A Dag Run already exists for Dag test_trigger_conflict " + "with run id test_run_id" + ), } }, ) @@ -1357,22 +1811,42 @@ def handle_request(request: httpx.Request) -> httpx.Response: result = client.dag_runs.trigger(dag_id="test_trigger_conflict", run_id="test_run_id") assert result == ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) + assert requests == [ + ("GET", "/dag-runs/test_trigger_conflict/test_run_id"), + ] def test_trigger_conflict_reset_dag_run(self): - """Test that if dag run already exists and reset_dag_run=True, the client clears the dag run""" + """Test that if the Dag run already exists and reset_dag_run=True, the client clears the Dag run""" + + requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: - if request.url.path == "/dag-runs/test_trigger_conflict_reset/test_run_id": + requests.append((request.method, request.url.path)) + if ( + request.method == "GET" + and request.url.path == "/dag-runs/test_trigger_conflict_reset/test_run_id" + ): + return httpx.Response(status_code=200, json={"detail": "exists"}) + if ( + request.method == "POST" + and request.url.path == "/dag-runs/test_trigger_conflict_reset/test_run_id" + ): return httpx.Response( status_code=409, json={ "detail": { "reason": "already_exists", - "message": "A Dag Run already exists for Dag test_trigger_conflict with run id test_run_id", + "message": ( + "A Dag Run already exists for Dag test_trigger_conflict_reset " + "with run id test_run_id" + ), } }, ) - if request.url.path == "/dag-runs/test_trigger_conflict_reset/test_run_id/clear": + if ( + request.method == "POST" + and request.url.path == "/dag-runs/test_trigger_conflict_reset/test_run_id/clear" + ): return httpx.Response(status_code=204) return httpx.Response(status_code=422) @@ -1384,9 +1858,13 @@ def handle_request(request: httpx.Request) -> httpx.Response: ) assert result == OKResponse(ok=True) + assert requests == [ + ("GET", "/dag-runs/test_trigger_conflict_reset/test_run_id"), + ("POST", "/dag-runs/test_trigger_conflict_reset/test_run_id/clear"), + ] def test_clear(self): - """Test that the client can clear a dag run""" + """Test that the client can clear a Dag run""" def handle_request(request: httpx.Request) -> httpx.Response: if request.url.path == "/dag-runs/test_clear/test_run_id/clear": From 8a9259d5c0fd71c04a79805db7a17b03bb5069db Mon Sep 17 00:00:00 2001 From: Hemkumar Chheda Date: Tue, 19 May 2026 19:04:38 +0530 Subject: [PATCH 3/3] Use get_count for TriggerDagRunOperator pre-check instead of _dag_run_exists Replace the custom _dag_run_exists helper (GET dag-runs/{dag_id}/{run_id} with compatibility fallback for 405/422/5xx) with direct self.get_count(run_ids=[run_id]) calls. The dag-runs/count endpoint has been available since Airflow 3.0.0, so the 40-line compatibility layer is unnecessary. Also removes _is_legacy_previous_dag_run_route_response, updates all trigger tests to mock GET /dag-runs/count, and adds a run_ids query param assertion. --- task-sdk/src/airflow/sdk/api/client.py | 68 +---- task-sdk/tests/task_sdk/api/test_client.py | 306 ++++----------------- 2 files changed, 69 insertions(+), 305 deletions(-) diff --git a/task-sdk/src/airflow/sdk/api/client.py b/task-sdk/src/airflow/sdk/api/client.py index 108718c976724..e784111ef99d1 100644 --- a/task-sdk/src/airflow/sdk/api/client.py +++ b/task-sdk/src/airflow/sdk/api/client.py @@ -869,10 +869,16 @@ def trigger( run_after=run_after, ) - dag_run_exists_before_trigger = self._dag_run_exists(dag_id=dag_id, run_id=run_id) - if dag_run_exists_before_trigger is True: + # GET dag-runs/count is available since Airflow 3.0.0. Failures propagate rather + # than falling through to POST — if the count endpoint is unreachable, the POST + # would fail too. None signals dry-run (skip pre-check, no-op POST). + dag_run_count_before = ( + None if self.client._dry_run else self.get_count(dag_id=dag_id, run_ids=[run_id]).count + ) + if dag_run_count_before is not None and dag_run_count_before > 0: if reset_dag_run: log.info("Dag Run already exists; Resetting Dag Run.", dag_id=dag_id, run_id=run_id) + # TODO: Make clear() idempotent as a follow-up. return self.clear(run_id=run_id, dag_id=dag_id) log.info("Dag Run already exists!", dag_id=dag_id, run_id=run_id) return ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) @@ -882,10 +888,7 @@ def trigger( "POST", f"dag-runs/{dag_id}/{run_id}", content=body.model_dump_json(exclude_defaults=True) ) except (httpx.ReadError, httpx.ReadTimeout, httpx.RemoteProtocolError): - if ( - dag_run_exists_before_trigger is False - and self._dag_run_exists(dag_id=dag_id, run_id=run_id, retry=True) is True - ): + if dag_run_count_before == 0 and self.get_count(dag_id=dag_id, run_ids=[run_id]).count > 0: log.info( "Dag Run exists after ambiguous trigger response; treating trigger as successful.", dag_id=dag_id, @@ -914,59 +917,6 @@ def trigger( return OKResponse(ok=True) - def _dag_run_exists(self, dag_id: str, run_id: str, *, retry: bool = False) -> bool | None: - """Return whether the Dag run exists, or None when the detail endpoint is unavailable.""" - if self.client._dry_run: - return None - - try: - if retry: - self.client.get(f"dag-runs/{dag_id}/{run_id}") - else: - self.client._request_without_retry("GET", f"dag-runs/{dag_id}/{run_id}") - except httpx.RequestError: - return None - except ServerResponseError as e: - if e.response.status_code == HTTPStatus.NOT_FOUND: - return False - if e.response.status_code == HTTPStatus.METHOD_NOT_ALLOWED: - # Older execution API servers may not support the Dag run detail endpoint yet. - return None - if ( - e.response.status_code == HTTPStatus.UNPROCESSABLE_ENTITY - and run_id == "previous" - and self._is_legacy_previous_dag_run_route_response(e.response) - ): - return None - if e.response.status_code >= HTTPStatus.INTERNAL_SERVER_ERROR: - return None - raise - except httpx.HTTPStatusError as e: - if e.response.status_code == HTTPStatus.NOT_FOUND: - return False - if e.response.status_code == HTTPStatus.METHOD_NOT_ALLOWED: - return None - if e.response.status_code >= HTTPStatus.INTERNAL_SERVER_ERROR: - return None - raise - return True - - @staticmethod - def _is_legacy_previous_dag_run_route_response(response: httpx.Response) -> bool: - """Return whether a 422 came from the legacy ``/{dag_id}/previous`` endpoint.""" - try: - detail = response.json()["detail"] - except (KeyError, ValueError, TypeError): - return False - if not isinstance(detail, list): - return False - return any( - isinstance(error, dict) - and error.get("type") == "missing" - and error.get("loc") == ["query", "logical_date"] - for error in detail - ) - def clear(self, dag_id: str, run_id: str) -> OKResponse: """Clear a Dag run via the API server.""" self.client.post(f"dag-runs/{dag_id}/{run_id}/clear") diff --git a/task-sdk/tests/task_sdk/api/test_client.py b/task-sdk/tests/task_sdk/api/test_client.py index 49c33c0ba88e0..329b97d359ed3 100644 --- a/task-sdk/tests/task_sdk/api/test_client.py +++ b/task-sdk/tests/task_sdk/api/test_client.py @@ -1263,8 +1263,10 @@ def test_trigger(self): def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + assert request.url.params["dag_id"] == "test_trigger" + assert request.url.params["run_ids"] == "test_run_id" + return httpx.Response(status_code=200, json=0) if request.url.path == "/dag-runs/test_trigger/test_run_id": actual_body = json.loads(request.read()) assert actual_body["logical_date"] == "2025-01-01T00:00:00Z" @@ -1285,7 +1287,7 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), ] @@ -1300,8 +1302,8 @@ def test_trigger_pre_existing_dag_run_returns_conflict_without_posting(self): def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=200, json={"detail": "exists"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=1) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": return httpx.Response(status_code=500, json={"detail": "POST should not happen"}) return httpx.Response(status_code=422) @@ -1311,7 +1313,7 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ] def test_trigger_clears_pre_existing_dag_run_without_posting_when_resetting(self): @@ -1319,8 +1321,8 @@ def test_trigger_clears_pre_existing_dag_run_without_posting_when_resetting(self def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=200, json={"detail": "exists"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=1) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": return httpx.Response(status_code=500, json={"detail": "POST should not happen"}) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id/clear": @@ -1332,37 +1334,17 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id/clear"), ] - def test_trigger_posts_when_precheck_endpoint_is_not_supported(self): + def test_trigger_returns_conflict_from_post_when_run_was_missing_before(self): requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=405, json={"detail": "Method Not Allowed"}) - if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=204) - return httpx.Response(status_code=422) - - client = make_client(transport=httpx.MockTransport(handle_request)) - result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") - - assert result == OKResponse(ok=True) - assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), - ("POST", "/dag-runs/test_trigger/test_run_id"), - ] - - def test_trigger_returns_conflict_after_unsupported_precheck(self): - requests: list[tuple[str, str]] = [] - - def handle_request(request: httpx.Request) -> httpx.Response: - requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=405, json={"detail": "Method Not Allowed"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": return httpx.Response( status_code=409, @@ -1380,139 +1362,10 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), - ("POST", "/dag-runs/test_trigger/test_run_id"), - ] - - def test_trigger_posts_when_precheck_request_error_leaves_existence_unknown(self): - requests: list[tuple[str, str]] = [] - - def handle_request(request: httpx.Request) -> httpx.Response: - requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - raise httpx.ReadError("Could not check Dag run existence", request=request) - if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=204) - return httpx.Response(status_code=422) - - client = make_client(transport=httpx.MockTransport(handle_request)) - result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") - - assert result == OKResponse(ok=True) - assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), - ("POST", "/dag-runs/test_trigger/test_run_id"), - ] - - def test_trigger_reraises_ambiguous_post_when_precheck_request_error_leaves_existence_unknown(self): - requests: list[tuple[str, str]] = [] - get_attempts = 0 - - def handle_request(request: httpx.Request) -> httpx.Response: - nonlocal get_attempts - requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - get_attempts += 1 - if get_attempts == 1: - raise httpx.ReadError("Could not check Dag run existence", request=request) - return httpx.Response(status_code=200, json={"detail": "exists"}) - if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": - raise httpx.ReadError("Trigger response was lost", request=request) - return httpx.Response(status_code=422) - - client = make_client(transport=httpx.MockTransport(handle_request)) - - with pytest.raises(httpx.ReadError, match="Trigger response was lost"): - client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") - - assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), - ("POST", "/dag-runs/test_trigger/test_run_id"), - ] - - def test_trigger_posts_when_non_json_precheck_server_error_leaves_existence_unknown(self): - requests: list[tuple[str, str]] = [] - - def handle_request(request: httpx.Request) -> httpx.Response: - requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=503, text="Service unavailable") - if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=204) - return httpx.Response(status_code=422) - - client = make_client(transport=httpx.MockTransport(handle_request)) - result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") - - assert result == OKResponse(ok=True) - assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), ] - def test_trigger_posts_when_previous_run_id_hits_legacy_previous_endpoint(self): - requests: list[tuple[str, str]] = [] - - def handle_request(request: httpx.Request) -> httpx.Response: - requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/previous": - return httpx.Response( - status_code=422, - json={ - "detail": [ - { - "type": "missing", - "loc": ["query", "logical_date"], - "msg": "Field required", - "input": None, - } - ] - }, - ) - if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/previous": - return httpx.Response(status_code=204) - return httpx.Response(status_code=422) - - client = make_client(transport=httpx.MockTransport(handle_request)) - result = client.dag_runs.trigger(dag_id="test_trigger", run_id="previous") - - assert result == OKResponse(ok=True) - assert requests == [ - ("GET", "/dag-runs/test_trigger/previous"), - ("POST", "/dag-runs/test_trigger/previous"), - ] - - def test_trigger_raises_unexpected_precheck_validation_error(self): - requests: list[tuple[str, str]] = [] - - def handle_request(request: httpx.Request) -> httpx.Response: - requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/previous": - return httpx.Response( - status_code=422, - json={ - "detail": [ - { - "type": "value_error", - "loc": ["query", "state"], - "msg": "Invalid state", - "input": "not-a-state", - } - ] - }, - ) - return httpx.Response(status_code=422) - - client = make_client(transport=httpx.MockTransport(handle_request)) - - with pytest.raises(ServerResponseError) as err: - client.dag_runs.trigger(dag_id="test_trigger", run_id="previous") - - assert err.value.response.status_code == 422 - assert requests == [ - ("GET", "/dag-runs/test_trigger/previous"), - ] - def test_trigger_treats_read_error_as_success_when_dag_run_appears_after_missing_precheck(self): requests: list[tuple[str, str]] = [] dag_run_exists = False @@ -1520,10 +1373,8 @@ def test_trigger_treats_read_error_as_success_when_dag_run_appears_after_missing def handle_request(request: httpx.Request) -> httpx.Response: nonlocal dag_run_exists requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - if dag_run_exists: - return httpx.Response(status_code=200, json={"detail": "exists"}) - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=1 if dag_run_exists else 0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": dag_run_exists = True raise httpx.ReadError("Trigger response was lost", request=request) @@ -1534,9 +1385,9 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ] def test_trigger_treats_read_timeout_as_success_when_dag_run_appears_after_missing_precheck(self): @@ -1546,10 +1397,8 @@ def test_trigger_treats_read_timeout_as_success_when_dag_run_appears_after_missi def handle_request(request: httpx.Request) -> httpx.Response: nonlocal dag_run_exists requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - if dag_run_exists: - return httpx.Response(status_code=200, json={"detail": "exists"}) - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=1 if dag_run_exists else 0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": dag_run_exists = True raise httpx.ReadTimeout("Trigger response timed out", request=request) @@ -1560,9 +1409,9 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ] def test_trigger_retries_followup_probe_after_ambiguous_response(self): @@ -1575,13 +1424,13 @@ def test_trigger_retries_followup_probe_after_ambiguous_response(self): def handle_request(request: httpx.Request) -> httpx.Response: nonlocal dag_run_exists, followup_attempts requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": + if request.method == "GET" and request.url.path == "/dag-runs/count": if not dag_run_exists: - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + return httpx.Response(status_code=200, json=0) followup_attempts += 1 if followup_attempts == 1: return httpx.Response(status_code=500, json={"detail": "Internal Server Error"}) - return httpx.Response(status_code=200, json={"detail": "exists"}) + return httpx.Response(status_code=200, json=1) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": dag_run_exists = True raise httpx.ReadError("Trigger response was lost", request=request) @@ -1592,10 +1441,10 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), - ("GET", "/dag-runs/test_trigger/test_run_id"), - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), + ("GET", "/dag-runs/count"), ] def test_trigger_treats_read_error_as_success_when_resetting_missing_dag_run_appears(self): @@ -1605,10 +1454,8 @@ def test_trigger_treats_read_error_as_success_when_resetting_missing_dag_run_app def handle_request(request: httpx.Request) -> httpx.Response: nonlocal dag_run_exists requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - if dag_run_exists: - return httpx.Response(status_code=200, json={"detail": "exists"}) - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=1 if dag_run_exists else 0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": dag_run_exists = True raise httpx.ReadError("Trigger response was lost", request=request) @@ -1623,9 +1470,9 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ] def test_trigger_treats_remote_protocol_error_as_success_when_dag_run_appears_after_missing_precheck( @@ -1637,10 +1484,8 @@ def test_trigger_treats_remote_protocol_error_as_success_when_dag_run_appears_af def handle_request(request: httpx.Request) -> httpx.Response: nonlocal dag_run_exists requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - if dag_run_exists: - return httpx.Response(status_code=200, json={"detail": "exists"}) - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=1 if dag_run_exists else 0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": dag_run_exists = True raise httpx.RemoteProtocolError("Trigger response was lost", request=request) @@ -1651,46 +1496,18 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), - ("POST", "/dag-runs/test_trigger/test_run_id"), - ("GET", "/dag-runs/test_trigger/test_run_id"), - ] - - def test_trigger_returns_conflict_after_missing_precheck(self): - requests: list[tuple[str, str]] = [] - - def handle_request(request: httpx.Request) -> httpx.Response: - requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) - if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response( - status_code=409, - json={ - "detail": { - "reason": "already_exists", - "message": "A Dag Run already exists for Dag test_trigger with run id test_run_id", - } - }, - ) - return httpx.Response(status_code=422) - - client = make_client(transport=httpx.MockTransport(handle_request)) - result = client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") - - assert result == ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) - assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ] - def test_trigger_conflict_after_unsupported_precheck_reset_dag_run(self): + def test_trigger_conflict_from_post_clears_run_when_resetting(self): requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=405, json={"detail": "Method Not Allowed"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": return httpx.Response( status_code=409, @@ -1714,18 +1531,18 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), ("POST", "/dag-runs/test_trigger/test_run_id/clear"), ] - def test_trigger_reraises_read_error_when_dag_run_is_missing_after_missing_precheck(self): + def test_trigger_reraises_read_error_when_dag_run_is_missing_after_precheck(self): requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": raise httpx.ReadError("Trigger response was lost", request=request) return httpx.Response(status_code=422) @@ -1736,18 +1553,18 @@ def handle_request(request: httpx.Request) -> httpx.Response: client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ] - def test_trigger_reraises_connect_error_even_if_dag_run_exists_after_missing_precheck(self): + def test_trigger_reraises_connect_error_even_if_dag_run_exists_after_precheck(self): requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": raise httpx.ConnectError("Could not connect", request=request) return httpx.Response(status_code=422) @@ -1758,17 +1575,17 @@ def handle_request(request: httpx.Request) -> httpx.Response: client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), ] - def test_trigger_reraises_pool_timeout_even_if_dag_run_exists_after_missing_precheck(self): + def test_trigger_reraises_pool_timeout_even_if_dag_run_exists_after_precheck(self): requests: list[tuple[str, str]] = [] def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger/test_run_id": - return httpx.Response(status_code=404, json={"detail": "Dag run not found"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=0) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger/test_run_id": raise httpx.PoolTimeout("Could not get a connection from the pool", request=request) return httpx.Response(status_code=422) @@ -1779,7 +1596,7 @@ def handle_request(request: httpx.Request) -> httpx.Response: client.dag_runs.trigger(dag_id="test_trigger", run_id="test_run_id") assert requests == [ - ("GET", "/dag-runs/test_trigger/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger/test_run_id"), ] @@ -1790,8 +1607,8 @@ def test_trigger_conflict(self): def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if request.method == "GET" and request.url.path == "/dag-runs/test_trigger_conflict/test_run_id": - return httpx.Response(status_code=200, json={"detail": "exists"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=1) if request.method == "POST" and request.url.path == "/dag-runs/test_trigger_conflict/test_run_id": return httpx.Response( status_code=409, @@ -1812,7 +1629,7 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == ErrorResponse(error=ErrorType.DAGRUN_ALREADY_EXISTS) assert requests == [ - ("GET", "/dag-runs/test_trigger_conflict/test_run_id"), + ("GET", "/dag-runs/count"), ] def test_trigger_conflict_reset_dag_run(self): @@ -1822,11 +1639,8 @@ def test_trigger_conflict_reset_dag_run(self): def handle_request(request: httpx.Request) -> httpx.Response: requests.append((request.method, request.url.path)) - if ( - request.method == "GET" - and request.url.path == "/dag-runs/test_trigger_conflict_reset/test_run_id" - ): - return httpx.Response(status_code=200, json={"detail": "exists"}) + if request.method == "GET" and request.url.path == "/dag-runs/count": + return httpx.Response(status_code=200, json=1) if ( request.method == "POST" and request.url.path == "/dag-runs/test_trigger_conflict_reset/test_run_id" @@ -1859,7 +1673,7 @@ def handle_request(request: httpx.Request) -> httpx.Response: assert result == OKResponse(ok=True) assert requests == [ - ("GET", "/dag-runs/test_trigger_conflict_reset/test_run_id"), + ("GET", "/dag-runs/count"), ("POST", "/dag-runs/test_trigger_conflict_reset/test_run_id/clear"), ]