ephraimbuddy commented on code in PR #74370:
URL: https://github.com/apache/airflow/pull/74370#discussion_r4229710641


##########
airflow-core/src/airflow/dag_processing/api_client.py:
##########
@@ -0,0 +1,440 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""Execution API client owned by a standalone Dag processor's manager loop."""
+
+from __future__ import annotations
+
+import math
+from collections.abc import Iterator
+from contextlib import contextmanager
+from contextvars import ContextVar
+from dataclasses import dataclass, field
+from pathlib import Path
+from time import monotonic
+from typing import Any
+from uuid import UUID
+
+import httpx
+import jwt
+import structlog
+from uuid6 import uuid7
+
+from airflow.api_fastapi.execution_api.datamodels.job import (
+    DagParseTokenBody,
+    DagParseTokenResponse,
+    JobCompleteBody,
+    JobHeartbeatResponse,
+    JobRegisterBody,
+    JobRegisterResponse,
+    JobState,
+    TerminalJobState,
+)
+from airflow.api_fastapi.execution_api.versions import bundle
+from airflow.sdk.api.client import BearerAuth, Client
+from airflow.sdk.execution_time.comms import GetConnection, GetVariable, 
MaskSecret
+from airflow.sdk.execution_time.request_handlers import handle_get_connection, 
handle_get_variable
+
+log = structlog.get_logger(__name__)
+
+
+@dataclass
+class DagParseContext:
+    """Credential cache owned by one parsing subprocess's supervisor."""
+
+    request: DagParseTokenBody
+    token: str | None = field(default=None, repr=False)
+    expires_at: float = 0.0
+    renew_at: float = 0.0
+
+
+# Core owns the control contracts, so it speaks the version its datamodels 
match, as the Task SDK does
+# for its own; runtime requests keep the SDK's negotiated version.
+_JOB_API_HEADERS = {"Airflow-API-Version": bundle.version_values[0]}
+
+
+class DagProcessorRegistrationRetired(RuntimeError):
+    """The manager must restart; the Job ended or no longer belongs to this 
session."""
+
+
+class DagProcessorJobAlreadyRunning(RuntimeError):
+    """The session's previous Job has not completed or stopped heartbeating 
yet."""
+
+
+def get_error_reason(error: httpx.HTTPStatusError) -> str | None:
+    try:
+        payload = error.response.json()
+    except ValueError:
+        return None
+    detail = payload.get("detail") if isinstance(payload, dict) else None
+    return detail.get("reason") if isinstance(detail, dict) else None
+
+
+class DagProcessorAPIClient(Client):
+    """
+    Register one processor process and use its Job token for subsequent API 
requests.
+
+    Create one client per process start. Token rotation and registration 
retries retain that process's
+    registration ID. Heartbeats return the server's Job state so the manager 
can stop its importers on
+    ``RESTARTING``. Closing this client closes the HTTP connection pool; call 
``complete_job`` explicitly
+    to record the outcome before closing it.
+
+    Wrap subprocess requests in ``use_parse`` and manager secret lookups in 
``use_bundle``. Early renewal is
+    best effort; ``restart_required`` tells the manager to drain and restart 
after registration retirement,
+    or once a request finds the Job completed or replaced, which raises 
``DagProcessorRegistrationRetired``.
+    The existing token remains usable until expiry, subject to the API's 
ownership checks.
+
+    After completion is attempted, only retries of that same completion are 
allowed. An expired token
+    can be renewed if the Job is still open. Retirement raises 
``DagProcessorRegistrationRetired``;
+    it does not confirm which outcome was saved or whether another process 
replaced the Job.
+    """
+
+    def __init__(
+        self,
+        *,
+        base_url: str,
+        token_file: str | Path,
+        hostname: str,
+        unixname: str | None = None,
+        bundle_names: list[str] | None = None,
+        token_reload_interval: float = 30.0,
+        **kwargs: Any,
+    ):
+        if not math.isfinite(token_reload_interval) or token_reload_interval < 
0:
+            raise ValueError("token_reload_interval must be finite and 
nonnegative")
+        self._registration = JobRegisterBody(
+            registration_id=uuid7(),
+            hostname=hostname,
+            unixname=unixname,
+            bundle_names=bundle_names,
+        )
+        self._token_file = Path(token_file)
+        self._token_reload_interval = token_reload_interval
+        self._session_token: str | None = None
+        self._reload_at = 0.0
+        self._registered_with: str | None = None
+        self._renew_at = 0.0
+        self._expires_at = 0.0
+        self._retry_renewal_at = 0.0
+        self._restart_required = False
+        self._bundle_context: ContextVar[str | None] = 
ContextVar("dag_processor_bundle", default=None)
+        self._parse_context: ContextVar[DagParseContext | None] = 
ContextVar("dag_parse", default=None)
+        self._job_id: int | None = None
+        self._completion_state: TerminalJobState | None = None
+        self._completed = False
+        super().__init__(base_url=base_url, token="", **kwargs)
+
+    @property
+    def registration_id(self) -> UUID:
+        return self._registration.registration_id
+
+    @property
+    def job_id(self) -> int | None:
+        return self._job_id
+
+    @property
+    def restart_required(self) -> bool:
+        return self._restart_required
+
+    @contextmanager
+    def use_bundle(self, bundle_name: str) -> Iterator[DagProcessorAPIClient]:
+        """Select the bundle for one subprocess request, restoring the 
previous context afterwards."""
+        if not bundle_name:
+            raise ValueError("A Dag processor request needs a nonempty bundle 
name")
+        context = self._bundle_context.set(bundle_name)
+        try:
+            yield self
+        finally:
+            self._bundle_context.reset(context)
+
+    @contextmanager
+    def use_parse(self, context: DagParseContext) -> 
Iterator[DagProcessorAPIClient]:
+        """Answer a subprocess with its own credential; restore the manager 
context afterwards."""
+        selected = self._parse_context.set(context)
+        try:
+            yield self
+        finally:
+            self._parse_context.reset(selected)
+
+    def _get_parse_token(self, context: DagParseContext, *, retry: bool) -> 
str:
+        if context.token is not None and monotonic() < context.expires_at:
+            if monotonic() < context.renew_at:
+                return context.token
+            try:
+                self._exchange_parse_token(context, retry=False, 
timeout=self._get_bounded_timeout())
+            except (httpx.HTTPError, ValueError) as error:
+                if isinstance(error, httpx.HTTPStatusError) and 
error.response.status_code < 500:
+                    raise
+                context.renew_at = monotonic() + 30
+                log.warning(
+                    "Unable to renew Dag parsing token",
+                    job_id=self._job_id,
+                    attempt_id=str(context.request.attempt_id),
+                    error_type=type(error).__name__,
+                )
+            if monotonic() < context.expires_at:
+                return context.token
+        return self._exchange_parse_token(context, retry=retry)
+
+    def _exchange_parse_token(
+        self, context: DagParseContext, *, retry: bool, timeout: httpx.Timeout 
| None = None
+    ) -> str:
+        self._ensure_job_token(retry=retry)
+        started_at = monotonic()
+        response = super().request(
+            "POST",
+            f"jobs/{self._require_job_id()}/parse-token",
+            json=context.request.model_dump(mode="json"),
+            retry=retry,
+            headers=_JOB_API_HEADERS,
+            timeout=timeout or self.timeout,
+        )
+        parsed = DagParseTokenResponse.model_validate_json(response.content)
+        try:
+            claims = jwt.decode(parsed.token, options={"verify_signature": 
False})
+            lifetime = float(claims["exp"]) - float(claims["iat"])
+            if (
+                not math.isfinite(lifetime)
+                or lifetime <= 0
+                or claims.get("scope") != "dag_parse"
+                or claims.get("sub") != str(context.request.attempt_id)
+                or claims.get("job_id") != self._job_id
+                or claims.get("dag_bundles") != [context.request.bundle_name]
+                or claims.get("relative_fileloc") != 
context.request.relative_fileloc
+            ):
+                raise ValueError
+        except (jwt.PyJWTError, KeyError, TypeError, ValueError):
+            raise ValueError("Token exchange returned an invalid Dag parsing 
token") from None
+        context.token = parsed.token
+        context.expires_at = started_at + lifetime
+        context.renew_at = started_at + lifetime * 0.8
+        return parsed.token
+
+    def _update_auth(self, response: httpx.Response) -> None:
+        # Task-token refresh headers cannot replace a provisioned or Job-bound 
credential.
+        pass
+
+    def _read_session_token(self, *, force: bool = False) -> str:
+        now = monotonic()
+        if self._session_token is None or force or now >= self._reload_at:
+            token = self._token_file.read_text().strip()
+            if not token:
+                raise ValueError(f"Dag processor token file is empty: 
{self._token_file}")
+            self._session_token = token
+            self._reload_at = now + self._token_reload_interval
+        return self._session_token
+
+    def _check_can_run(self) -> None:
+        if self._completion_state is not None:
+            raise RuntimeError("The Dag processor Job is completing; only 
completion retries are allowed")
+
+    def _require_job_id(self) -> int:
+        if self._job_id is None:
+            raise RuntimeError("Register the Dag processor Job before making 
API requests")
+        return self._job_id
+
+    def register_job(self, *, retry: bool = True) -> int:
+        """
+        Register this process, recover its registration, or renew its Job 
token.
+
+        While the session's previous Job is still alive, raise 
``DagProcessorJobAlreadyRunning``.
+        Retrying with this client retains the registration ID.
+        """
+        self._check_can_run()
+        try:
+            return self._register_job(retry=retry)
+        except httpx.HTTPStatusError as error:
+            if error.response.status_code == 409 and get_error_reason(error) 
== "job_running":
+                raise DagProcessorJobAlreadyRunning(
+                    "The previous Dag processor Job is still alive"
+                ) from error
+            raise
+
+    def _register_job(self, *, retry: bool, timeout: httpx.Timeout | None = 
None) -> int:
+        if self.restart_required:
+            raise DagProcessorRegistrationRetired(
+                "The Dag processor registration has ended; restart required"
+            )
+        session_token = self._read_session_token(force=True)
+        started_at = monotonic()
+        body = self._registration.model_dump(mode="json")
+        retried_auth = False
+        while True:
+            try:
+                response = super().request(
+                    "POST",
+                    "jobs",
+                    json=body,
+                    auth=BearerAuth(session_token),
+                    retry=retry,
+                    headers=_JOB_API_HEADERS,
+                    timeout=timeout or self.timeout,
+                )
+                break
+            except httpx.HTTPStatusError as error:
+                if error.response.status_code == 409 and 
get_error_reason(error) == "registration_retired":
+                    self._restart_required = True
+                    raise DagProcessorRegistrationRetired(
+                        "The Dag processor registration has ended; restart 
required"
+                    ) from error
+                if error.response.status_code not in (401, 403) or 
retried_auth:
+                    raise
+                rotated = self._read_session_token(force=True)
+                if rotated == session_token:
+                    raise
+                session_token = rotated
+                retried_auth = True
+
+        registered = JobRegisterResponse.model_validate_json(response.content)
+        if self._job_id is not None and self._job_id != registered.job_id:
+            raise RuntimeError("Registration returned a different Dag 
processor Job")
+        try:
+            # Unverified claims only schedule renewal; the API server remains 
the authority on validity.
+            claims = jwt.decode(registered.token, options={"verify_signature": 
False})
+            lifetime = float(claims["exp"]) - float(claims["iat"])
+            if (
+                not math.isfinite(lifetime)
+                or lifetime <= 0
+                or claims.get("scope") != "dag_processor"
+                or claims.get("job_id") != registered.job_id
+            ):
+                raise ValueError
+        except (jwt.PyJWTError, KeyError, TypeError, ValueError):
+            raise ValueError("Registration returned an invalid Dag processor 
Job token") from None
+
+        self._job_id = registered.job_id
+        self.auth = BearerAuth(registered.token)
+        self._registered_with = session_token
+        self._renew_at = started_at + lifetime * 0.8
+        self._expires_at = started_at + lifetime
+        self._retry_renewal_at = 0.0
+        return registered.job_id
+
+    def _get_bounded_timeout(self) -> httpx.Timeout:
+        return httpx.Timeout(
+            **{
+                key: min(value if value is not None else 1.0, 1.0)
+                for key, value in self.timeout.as_dict().items()
+            }
+        )
+
+    def _ensure_job_token(self, *, retry: bool) -> None:
+        self._require_job_id()
+        if monotonic() >= self._expires_at:
+            self._register_job(retry=retry, timeout=None if retry else 
self._get_bounded_timeout())
+            return
+        if self.restart_required or monotonic() < self._retry_renewal_at:
+            return
+        try:
+            if self._read_session_token() != self._registered_with or 
monotonic() >= self._renew_at:
+                # Do not spend the runtime request's retry budget on an 
optional early renewal.
+                self._register_job(retry=False, 
timeout=self._get_bounded_timeout())
+        except (httpx.HTTPError, OSError, ValueError, 
DagProcessorRegistrationRetired) as error:
+            self._retry_renewal_at = monotonic() + 30
+            log.warning(
+                "Unable to renew Dag processor Job token",
+                job_id=self._job_id,
+                error_type=type(error).__name__,
+                restart_required=self.restart_required,
+            )
+        if monotonic() >= self._expires_at:
+            self._register_job(retry=retry, timeout=None if retry else 
self._get_bounded_timeout())
+
+    def request(self, *args, retry: bool = True, **kwargs) -> httpx.Response:
+        """Use a parsing credential for subprocess requests, and the Job 
credential for manager work."""
+        self._check_can_run()
+        try:
+            headers = httpx.Headers(kwargs.get("headers"))
+            if context := self._parse_context.get():
+                kwargs["auth"] = BearerAuth(self._get_parse_token(context, 
retry=retry))
+                headers["Airflow-Dag-Bundle"] = context.request.bundle_name
+            else:
+                self._ensure_job_token(retry=retry)
+                if bundle_name := self._bundle_context.get():
+                    headers["Airflow-Dag-Bundle"] = bundle_name
+            if kwargs.get("content") is not None:
+                headers.setdefault("Content-Type", "application/json")
+            kwargs["headers"] = headers
+            return super().request(*args, retry=retry, **kwargs)
+        except httpx.HTTPStatusError as error:
+            if error.response.status_code == 403 and get_error_reason(error) 
== "job_closed":
+                self._restart_required = True
+                raise DagProcessorRegistrationRetired(
+                    "The Dag processor Job has completed or been replaced; 
restart required"
+                ) from error
+            raise
+
+    def heartbeat(self) -> JobState:
+        """Heartbeat once; the manager's next iteration retries transport 
failures."""
+        job_id = self._require_job_id()
+        response = self.request("POST", f"jobs/{job_id}/heartbeat", 
retry=False, headers=_JOB_API_HEADERS)
+        return JobHeartbeatResponse.model_validate_json(response.content).state
+
+    def complete_job(self, state: TerminalJobState) -> None:
+        """Complete this Job, retaining its identity and outcome across lost 
acknowledgments."""
+        body = JobCompleteBody(state=state)
+        job_id = self._require_job_id()
+        if self._completion_state is not None and self._completion_state != 
body.state:
+            raise ValueError("A Dag processor completion retry must keep the 
original outcome")
+        if self._completed:
+            return
+        self._completion_state = body.state
+        # A still-valid token can replay completion even after the Job closes; 
renewal cannot.
+        retried_expiry = False
+        while True:
+            if monotonic() >= self._expires_at:
+                self._register_job(retry=True)
+            try:
+                super().request(
+                    "POST",
+                    f"jobs/{job_id}/complete",
+                    json=body.model_dump(mode="json"),
+                    headers=_JOB_API_HEADERS,
+                )
+            except httpx.HTTPStatusError as error:
+                if (
+                    retried_expiry
+                    or error.response.status_code not in (401, 403)
+                    or monotonic() < self._expires_at
+                ):
+                    raise
+                retried_expiry = True
+            else:
+                self._completed = True
+                return
+
+
+class DagProcessorSecretsComms:

Review Comment:
   Fixed in #74233. API mode resets the cache and skips initialization



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to