kaxil commented on code in PR #74370: URL: https://github.com/apache/airflow/pull/74370#discussion_r4218428028
########## airflow-core/src/airflow/cli/commands/dag_processor_token_command.py: ########## @@ -0,0 +1,76 @@ +# 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. +"""Dag processor token command.""" + +from __future__ import annotations + +import contextlib +import logging +import os +import tempfile +import time +from pathlib import Path + +import uuid6 + +from airflow.api_fastapi.execution_api.app import create_jwt_generator +from airflow.api_fastapi.execution_api.dag_processor_tokens import generate_dag_processor_session_token +from airflow.configuration import conf +from airflow.dag_processing.bundles.manager import _load_bundle_config_snapshot +from airflow.utils.providers_configuration_loader import providers_configuration_loaded + +log = logging.getLogger(__name__) + + +def write_token_file(path: Path, token: str) -> None: + """Replace the token file in one step, so a processor rereading it never sees a partial token.""" + fd, tmp_path = tempfile.mkstemp(dir=path.parent, prefix=f".{path.name}.") + try: + with os.fdopen(fd, "w") as tmp: + tmp.write(token) + tmp.flush() + os.fsync(tmp.fileno()) + os.replace(tmp_path, path) + except BaseException: + with contextlib.suppress(FileNotFoundError): + os.unlink(tmp_path) + raise + + +# Not wrapped in ``action_cli``: its audit logging writes to the metadata database, which a provisioning +# component that only holds the signing key cannot reach. +@providers_configuration_loaded +def dag_processor_token(args) -> None: + """Write a Dag processor session token to a file, and with ``--rotate`` keep replacing it.""" + # Names only: provisioning runs where the signing key is, which may not have the bundle classes installed. + configured = _load_bundle_config_snapshot().names + bundle_names = set(args.bundle_name or configured) + if unknown := bundle_names - configured: + raise SystemExit(f"Bundles not found: {', '.join(sorted(unknown))}") + + valid_for = args.valid_for or conf.getint("execution_api", "jwt_expiration_time") + token_file = Path(args.token_file) + session_id = uuid6.uuid7() Review Comment: A restart of this process ends up restarting every processor it feeds. Nothing in the `while True` loop below catches a failed `generate_dag_processor_session_token` or `write_token_file` (ENOSPC, an unreadable key), so one bad write exits it. When a supervisor brings it back, `uuid7()` picks a new session, the processor re-registers with its existing registration_id, `_resume_registration` sees `job.session_id != token.id` and answers `registration_retired`, and the client sets restart_required. Could the loop retry with backoff, and the session survive a restart (reuse the `sub` of a still-valid token file, or accept `--session-id`)? ########## airflow-core/src/airflow/cli/commands/dag_processor_token_command.py: ########## @@ -0,0 +1,76 @@ +# 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. +"""Dag processor token command.""" + +from __future__ import annotations + +import contextlib +import logging +import os +import tempfile +import time +from pathlib import Path + +import uuid6 + +from airflow.api_fastapi.execution_api.app import create_jwt_generator +from airflow.api_fastapi.execution_api.dag_processor_tokens import generate_dag_processor_session_token +from airflow.configuration import conf +from airflow.dag_processing.bundles.manager import _load_bundle_config_snapshot +from airflow.utils.providers_configuration_loader import providers_configuration_loaded + +log = logging.getLogger(__name__) + + +def write_token_file(path: Path, token: str) -> None: + """Replace the token file in one step, so a processor rereading it never sees a partial token.""" + fd, tmp_path = tempfile.mkstemp(dir=path.parent, prefix=f".{path.name}.") Review Comment: `mkstemp` creates the file 0600 and owned by the provisioner, and `os.replace` swaps it in on every rotation. Any chgrp/chmod/ACL an operator applied to "give the processor read access" (as the command description asks) is gone after the first rotation. If the processor runs as a different uid, `_read_session_token` then fails with PermissionError, which `_ensure_job_token` swallows as a renewal failure until the Job token expires. Carry the existing file's mode and group onto the temp fd before the replace, or document that both must run as the same user? ########## airflow-core/tests/unit/api_fastapi/execution_api/versions/head/test_jobs.py: ########## @@ -0,0 +1,418 @@ +# 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. +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from uuid import UUID + +import jwt +import pytest +from fastapi import Request +from sqlalchemy import event, insert, select, update + +from airflow.api_fastapi.execution_api.datamodels.token import ( + DagProcessorSessionToken, + DagProcessorToken, + ExecutionToken, +) +from airflow.api_fastapi.execution_api.security import require_auth +from airflow.jobs.job import Job, JobState +from airflow.models.dagbundle import DagBundleModel +from airflow.models.team import Team + +from tests_common.test_utils.config import conf_vars +from tests_common.test_utils.db import clear_db_dag_bundles, clear_db_jobs, clear_db_teams + +pytestmark = pytest.mark.db_test + +SESSION_ID = UUID("00000000-0000-0000-0000-0000000000aa") +OTHER_SESSION_ID = UUID("00000000-0000-0000-0000-0000000000bb") +REGISTRATION_ID = UUID("00000000-0000-0000-0000-00000000000a") +OTHER_REGISTRATION_ID = UUID("00000000-0000-0000-0000-00000000000b") +NOW = datetime(2026, 10, 5, 12, 0, tzinfo=timezone.utc) +SESSION_EXPIRY = NOW + timedelta(hours=1) + + [email protected](autouse=True) +def clean_db(): + clear_db_jobs() + clear_db_dag_bundles() + clear_db_teams() + yield + clear_db_jobs() + clear_db_dag_bundles() + clear_db_teams() + + [email protected](autouse=True) +def frozen_time(time_machine): + time_machine.move_to(NOW, tick=False) + + [email protected] +def authenticate(exec_app): + def _authenticate(**claims) -> None: + claims.setdefault("exp", SESSION_EXPIRY.timestamp()) + token_type = ( + DagProcessorSessionToken if claims["scope"] == "dag_processor_session" else DagProcessorToken + ) + + async def _auth(request: Request) -> ExecutionToken: + return token_type.model_validate({"id": SESSION_ID, "claims": claims}) + + exec_app.dependency_overrides[require_auth] = _auth + + return _authenticate + + [email protected] +def as_session(authenticate): + authenticate( + scope="dag_processor_session", + dag_bundles=frozenset({"bundle_a", "bundle_b"}), + exp=SESSION_EXPIRY.timestamp(), + ) + + +def _create_job( + session, + *, + session_id: UUID | None = SESSION_ID, + registration_id: UUID = REGISTRATION_ID, + state: JobState = JobState.RUNNING, + latest_heartbeat: datetime = NOW, + end_date: datetime | None = None, +) -> Job: + job = Job(job_type="DagProcessorJob", state=state) + job.session_id = session_id + job.registration_id = registration_id + job.hostname = "processor-1" + job.unixname = None + job.bundle_names = ["bundle_a", "bundle_b"] + job.latest_heartbeat = latest_heartbeat + job.end_date = end_date + session.add(job) + session.commit() + return job + + +def _register(client, registration_id: UUID = REGISTRATION_ID, **body): + return client.post( + "/execution/jobs", json={"registration_id": str(registration_id), "hostname": "processor-1", **body} + ) + + +def _decode(token: str) -> dict: + return jwt.decode(token, options={"verify_signature": False}) + + [email protected]("as_session") +class TestRegisterJob: + def test_registers_a_running_job_for_every_granted_bundle(self, client, session): + response = _register(client, unixname="airflow") + + assert response.status_code == 201, response.json() + job = session.get(Job, response.json()["job_id"]) + assert job.job_type == "DagProcessorJob" + assert job.state == JobState.RUNNING + assert (job.start_date, job.latest_heartbeat) == (NOW, NOW) + assert (job.hostname, job.unixname) == ("processor-1", "airflow") + assert job.bundle_names == ["bundle_a", "bundle_b"] + assert (job.session_id, job.registration_id) == (SESSION_ID, REGISTRATION_ID) + + def test_returns_a_token_for_the_job_that_never_outlives_the_session_token(self, client): + response = _register(client, bundle_names=["bundle_b"]) + + assert response.status_code == 201, response.json() + claims = _decode(response.json()["token"]) + assert claims["scope"] == "dag_processor" + assert claims["sub"] == str(SESSION_ID) + assert claims["job_id"] == response.json()["job_id"] + assert claims["dag_bundles"] == ["bundle_b"] + assert claims["exp"] <= SESSION_EXPIRY.timestamp() Review Comment: This passes without the cap in `generate_dag_processor_token`. The session expiry here is NOW + 1h while the generator lifetime is `[execution_api] jwt_expiration_time` (600s by default), so `exp` is already well under it. Swapping `valid_for=min(generator.valid_for, remaining)` for `valid_for=generator.valid_for` keeps this green, and nothing else covers the Job-token cap. Setting the session `exp` to NOW + 60 and asserting `exp == NOW + 60` would pin it. ########## airflow-core/src/airflow/api_fastapi/execution_api/routes/jobs.py: ########## @@ -0,0 +1,295 @@ +# 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. + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import svcs +from cadwyn import VersionedAPIRouter +from fastapi import HTTPException, Security, status +from sqlalchemy import insert, select, update +from sqlalchemy.exc import IntegrityError + +from airflow._shared.timezones import timezone +from airflow.api_fastapi.auth.tokens import JWTGenerator +from airflow.api_fastapi.common.db.common import SessionDep +from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc +from airflow.api_fastapi.execution_api.dag_processor_tokens import ( + ExpiredDagProcessorToken, + generate_dag_parse_token, + generate_dag_processor_token, +) +from airflow.api_fastapi.execution_api.datamodels.job import ( + DagParseTokenBody, + DagParseTokenResponse, + JobCompleteBody, + JobHeartbeatResponse, + JobRegisterBody, + JobRegisterResponse, +) +from airflow.api_fastapi.execution_api.datamodels.token import DagProcessorSessionToken, DagProcessorToken +from airflow.api_fastapi.execution_api.deps import DepContainer +from airflow.api_fastapi.execution_api.security import ( + JOB_UNCHECKED_SCOPE, + CurrentDagProcessorSessionToken, + CurrentDagProcessorToken, + ExecutionAPIRoute, + require_auth, +) +from airflow.configuration import conf +from airflow.jobs.dag_processor_job_runner import DagProcessorJobRunner +from airflow.jobs.job import Job, JobState +from airflow.models.dagbundle import DagBundleModel +from airflow.models.team import JobTeam + +if TYPE_CHECKING: + from uuid import UUID + + from sqlalchemy.orm import Session + +router = VersionedAPIRouter(route_class=ExecutionAPIRoute) + +_JOB_NOT_FOUND = create_openapi_http_exception_doc( + [(status.HTTP_404_NOT_FOUND, "Job not found for this token")] +) + + +def _check_token_job(job_id: int, token: DagProcessorToken) -> None: + if job_id != token.claims.job_id: + raise HTTPException( + status.HTTP_404_NOT_FOUND, + detail={"reason": "not_found", "message": f"Job {job_id} not found for this token"}, + ) + + +def _get_registered_job(registration_id: UUID, *, session: Session) -> Job | None: + return session.scalar(select(Job).where(Job.registration_id == registration_id).with_for_update()) + + +def _build_registration_response( + job_id: int, bundle_names: list[str], token: DagProcessorSessionToken, services: svcs.Container +) -> JobRegisterResponse: + try: + credential = generate_dag_processor_token( + services.get(JWTGenerator), + session_id=token.id, + job_id=job_id, + bundle_names=bundle_names, + session_expiry=token.claims.exp, + ) + except ExpiredDagProcessorToken as error: + raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(error)) from error + return JobRegisterResponse(job_id=job_id, token=credential) + + +def _resume_registration( + job: Job, + body: JobRegisterBody, + bundle_names: list[str], + token: DagProcessorSessionToken, + services: svcs.Container, +) -> JobRegisterResponse: + """Return the Job an earlier registration created, with a fresh token, if that Job is still the session's.""" + if job.session_id != token.id or job.end_date is not None: + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_retired", + "message": f"Registration {body.registration_id} has ended; a restarted processor uses a new one", + }, + ) + if (job.hostname, job.unixname, job.bundle_names) != (body.hostname, body.unixname, bundle_names): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_conflict", + "message": f"Registration {body.registration_id} was made with different details", + }, + ) + return _build_registration_response(job.id, bundle_names, token, services) + + [email protected]( + "", Review Comment: This brings back the shape the comment in routes/__init__.py warns about: an empty path under an include-time prefix. I checked a minimal router chain (child `""` route, included with `prefix="/jobs"`, then included without a prefix): it works on FastAPI 0.136.1 and raises "Prefix and path cannot be both empty" on 0.137.0. It's latent behind the `<0.137` pin, but declaring full paths on this router the way health.py does would avoid it. ########## airflow-core/src/airflow/api_fastapi/execution_api/routes/jobs.py: ########## @@ -0,0 +1,295 @@ +# 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. + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import svcs +from cadwyn import VersionedAPIRouter +from fastapi import HTTPException, Security, status +from sqlalchemy import insert, select, update +from sqlalchemy.exc import IntegrityError + +from airflow._shared.timezones import timezone +from airflow.api_fastapi.auth.tokens import JWTGenerator +from airflow.api_fastapi.common.db.common import SessionDep +from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc +from airflow.api_fastapi.execution_api.dag_processor_tokens import ( + ExpiredDagProcessorToken, + generate_dag_parse_token, + generate_dag_processor_token, +) +from airflow.api_fastapi.execution_api.datamodels.job import ( + DagParseTokenBody, + DagParseTokenResponse, + JobCompleteBody, + JobHeartbeatResponse, + JobRegisterBody, + JobRegisterResponse, +) +from airflow.api_fastapi.execution_api.datamodels.token import DagProcessorSessionToken, DagProcessorToken +from airflow.api_fastapi.execution_api.deps import DepContainer +from airflow.api_fastapi.execution_api.security import ( + JOB_UNCHECKED_SCOPE, + CurrentDagProcessorSessionToken, + CurrentDagProcessorToken, + ExecutionAPIRoute, + require_auth, +) +from airflow.configuration import conf +from airflow.jobs.dag_processor_job_runner import DagProcessorJobRunner +from airflow.jobs.job import Job, JobState +from airflow.models.dagbundle import DagBundleModel +from airflow.models.team import JobTeam + +if TYPE_CHECKING: + from uuid import UUID + + from sqlalchemy.orm import Session + +router = VersionedAPIRouter(route_class=ExecutionAPIRoute) + +_JOB_NOT_FOUND = create_openapi_http_exception_doc( + [(status.HTTP_404_NOT_FOUND, "Job not found for this token")] +) + + +def _check_token_job(job_id: int, token: DagProcessorToken) -> None: + if job_id != token.claims.job_id: + raise HTTPException( + status.HTTP_404_NOT_FOUND, + detail={"reason": "not_found", "message": f"Job {job_id} not found for this token"}, + ) + + +def _get_registered_job(registration_id: UUID, *, session: Session) -> Job | None: + return session.scalar(select(Job).where(Job.registration_id == registration_id).with_for_update()) + + +def _build_registration_response( + job_id: int, bundle_names: list[str], token: DagProcessorSessionToken, services: svcs.Container +) -> JobRegisterResponse: + try: + credential = generate_dag_processor_token( + services.get(JWTGenerator), + session_id=token.id, + job_id=job_id, + bundle_names=bundle_names, + session_expiry=token.claims.exp, + ) + except ExpiredDagProcessorToken as error: + raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(error)) from error + return JobRegisterResponse(job_id=job_id, token=credential) + + +def _resume_registration( + job: Job, + body: JobRegisterBody, + bundle_names: list[str], + token: DagProcessorSessionToken, + services: svcs.Container, +) -> JobRegisterResponse: + """Return the Job an earlier registration created, with a fresh token, if that Job is still the session's.""" + if job.session_id != token.id or job.end_date is not None: + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_retired", + "message": f"Registration {body.registration_id} has ended; a restarted processor uses a new one", + }, + ) + if (job.hostname, job.unixname, job.bundle_names) != (body.hostname, body.unixname, bundle_names): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_conflict", + "message": f"Registration {body.registration_id} was made with different details", + }, + ) + return _build_registration_response(job.id, bundle_names, token, services) + + [email protected]( + "", + status_code=status.HTTP_201_CREATED, + dependencies=[Security(require_auth, scopes=["token:dag_processor_session"])], + responses=create_openapi_http_exception_doc( + [ + ( + status.HTTP_403_FORBIDDEN, + "A requested bundle is not granted to the session, or its credential has expired", + ), + ( + status.HTTP_409_CONFLICT, + "The registration has ended, conflicts with an earlier one, or another process's Job is running", + ), + ] + ), +) +def register_job( + body: JobRegisterBody, + session: SessionDep, + token: DagProcessorSessionToken = CurrentDagProcessorSessionToken, + services: svcs.Container = DepContainer, +) -> JobRegisterResponse: + """ + Register the Job of a Dag processor session in exchange for its management credential. + + A registration creates one Job. Repeating it while that Job is open returns the Job with a fresh token; + once the Job completes or is replaced, the registration is refused for good. A new registration is + refused while the session's Job is alive, and otherwise ends and replaces it, which ends every token + issued for the replaced Job. + """ + requested = token.claims.dag_bundles if body.bundle_names is None else set(body.bundle_names) + if ungranted := requested - token.claims.dag_bundles: + raise HTTPException( + status.HTTP_403_FORBIDDEN, + detail={"reason": "bundle_not_granted", "message": f"Bundles not granted: {sorted(ungranted)}"}, + ) + bundle_names = sorted(requested) + + if registered := _get_registered_job(body.registration_id, session=session): + return _resume_registration(registered, body, bundle_names, token, services) + + now = timezone.utcnow() + previous = session.scalar(select(Job).where(Job.session_id == token.id).with_for_update()) + if previous is not None: + if previous.end_date is None and previous.is_alive(): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={"reason": "job_running", "message": f"Session already has running Job {previous.id}"}, + ) + if previous.end_date is None: + previous.state = JobState.FAILED + previous.end_date = now + previous.session_id = None + session.flush() + + try: + with session.begin_nested(): + # A core insert: Job.__init__ would stamp the API server's host and user and fire listeners. + session.execute( + insert(Job).values( + job_type=DagProcessorJobRunner.job_type, + state=JobState.RUNNING, + start_date=now, + latest_heartbeat=now, + hostname=body.hostname, + unixname=body.unixname, + bundle_names=bundle_names, + session_id=token.id, + registration_id=body.registration_id, + ) + ) + job_id = session.scalars(select(Job.id).where(Job.session_id == token.id)).one() + # Issue the credential before releasing the savepoint so an expiry failure + # also rolls back the new Job under SQLite's legacy transaction control. + response = _build_registration_response(job_id, bundle_names, token, services) + except IntegrityError: + # A retry sent before the original request finished can lose the race to it; resume the winner. + if registered := _get_registered_job(body.registration_id, session=session): + return _resume_registration(registered, body, bundle_names, token, services) + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={"reason": "job_running", "message": "Session registered another Job concurrently"}, + ) + if conf.getboolean("core", "multi_team"): + team_names = DagBundleModel.get_team_names(bundle_names, session=session) Review Comment: `dag_processor_command._get_team_names` resolves teams from bundle config on purpose. Its docstring says the job row is written before `sync_bundles()`, so a DB lookup sees no rows on a fresh deployment and stale rows right after a bundle changes team. Registration happens at the same point in startup, so this gets the same empty or stale result, and `_resume_registration` never refreshes it. Only the JobTeam rows behind the public Jobs team filter are affected, not authorization, but could this use the config-based lookup too? ########## airflow-core/src/airflow/api_fastapi/execution_api/security.py: ########## @@ -279,7 +330,122 @@ async def _require_live_attempt(token: TIToken, *, allow_callback: bool) -> None ) -CurrentTIToken: TIToken = Depends(require_auth) +async def _require_open_dag_processor_job( + request: Request, token_id: UUID, claims: DagProcessorClaims | DagParseClaims +) -> None: + """Refuse processor or parsing access after the Job ends or is replaced.""" + if request.scope.get(_REQUEST_SCOPE_JOB_KEY): + return + + from airflow.jobs.job import Job + + async with create_session_async() as session: + session_id = claims.session_id if isinstance(claims, DagParseClaims) else token_id + job_id = await session.scalar( + select(Job.id).where( + Job.id == claims.job_id, Job.session_id == session_id, Job.end_date.is_(None) + ) + ) + if job_id is None: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "reason": "job_closed", + "message": "The Job this token was issued for has completed or been replaced", + }, + ) + request.scope[_REQUEST_SCOPE_JOB_KEY] = job_id + + +CurrentExecutionToken: ExecutionToken = Depends(require_auth) + + +def require_task_token(token: ExecutionToken = CurrentExecutionToken) -> TIToken: + if not isinstance(token, TIToken): + raise HTTPException(status.HTTP_403_FORBIDDEN, detail="Task token required") + return token + + +def require_dag_processor_session_token( + token: ExecutionToken = CurrentExecutionToken, +) -> DagProcessorSessionToken: + if not isinstance(token, DagProcessorSessionToken): + raise HTTPException(status.HTTP_403_FORBIDDEN, detail="Dag processor session token required") + return token + + +def require_dag_processor_token(token: ExecutionToken = CurrentExecutionToken) -> DagProcessorToken: + if not isinstance(token, DagProcessorToken): + raise HTTPException(status.HTTP_403_FORBIDDEN, detail="Dag processor token required") + return token + + +CurrentTIToken: TIToken = Depends(require_task_token) +CurrentDagProcessorSessionToken: DagProcessorSessionToken = Depends(require_dag_processor_session_token) +CurrentDagProcessorToken: DagProcessorToken = Depends(require_dag_processor_token) + +DAG_BUNDLE_HEADER = "Airflow-Dag-Bundle" + +ExecutionOrDagParseToken = Security(require_auth, scopes=["token:execution", "token:dag_parse"]) +ExecutionOrProcessorSecretsToken = Security( + require_auth, scopes=["token:execution", "token:dag_processor", "token:dag_parse"] +) +"""Bundle preparation needs Connection and Variable reads before a file can be discovered.""" + + +async def get_selected_dag_bundle(request: Request, token=CurrentExecutionToken) -> str | None: + """Select a granted management bundle or use the immutable bundle of a parsing attempt.""" + if token.claims.scope not in ("dag_processor", "dag_parse"): + return None + bundle_name = request.headers.get(DAG_BUNDLE_HEADER) Review Comment: Reading `Airflow-Dag-Bundle` from `request.headers` keeps it out of the OpenAPI spec, so the spec doesn't say the connection and variable routes take it or can return 400 when it's missing (their 403 descriptions also still say "Task does not have access"). Declaring it as `Annotated[str | None, Header(alias=...)]` on this dependency would surface it, and those routes could list the 400. ########## airflow-core/docs/security/jwt_token_authentication.rst: ########## @@ -486,6 +494,54 @@ See :doc:`/security/security_model` for the full security implications, deployme guidance, and the planned strategic and tactical improvements. +Dag processor HTTP client credentials +------------------------------------- + +The core ``DagProcessorAPIClient`` lets processor integrations authenticate to the Execution +API without holding its signing key. The standard ``airflow dag-processor`` command does +not yet use this client. Provisioning a token file alone does not switch that command to +HTTP or remove its need for database access. + +On a trusted host with the API signing key, provision a session token for the bundles +the client may access:: + + airflow dag-processor-token --token-file /run/airflow/processor.jwt \ + --bundle-name dags-folder --rotate + +The command writes the token atomically with owner-only permissions. Keep the provisioner +running and mount its directory rather than a single file so token replacement stays +visible. Each processor process needs its own session. + +The three credentials have different issuers and permissions: + +- ``dag_processor_session`` is issued by trusted provisioning. It can only register a + Job or renew that Job's credential through ``POST /jobs``. +- ``dag_processor`` is returned by Job registration. It can heartbeat and complete that + Job, read Connections and Variables for a granted bundle, and exchange a parsing + credential through ``POST /jobs/{job_id}/parse-token``. It cannot outlive the session + token used to obtain it. +- ``dag_parse`` is returned by the parsing-token exchange for one attempt, bundle, and + relative file location. It permits parse-time requests, including callback-context reads, Review Comment: `dag_parse` can also PUT and DELETE Variables, which matches what a parse subprocess can do today but isn't obvious from "permits parse-time requests". On L533, "uses the bundle's current team" holds for reads, but `Variable.set` upserts on `key` alone, so a write can overwrite another team's Variable with the same key (same as task tokens on main, but worth stating here). The unchanged "All other endpoints require scope=execution" line above the token-validation section is also no longer true. ########## airflow-core/src/airflow/api_fastapi/execution_api/AGENTS.md: ########## @@ -77,4 +77,4 @@ Adding a new Execution API feature touches multiple packages. All of these must ## Token Scope Infrastructure -Token types (`"execution"`, `"workload"`), route-level enforcement via `ExecutionAPIRoute` + `require_auth`, and the `ti:self` path-parameter validation are documented in the module docstring of `security.py`. +Token types (`"execution"`, `"workload"`, `"callback"`, `"dag_processor_session"`, `"dag_processor"`), route-level enforcement via `ExecutionAPIRoute` + `require_auth`, and the `ti:self` path-parameter validation are documented in the module docstring of `security.py`. A route that admits `dag_processor` tokens must also bind the request to a granted Dag bundle; `test_token_scope_boundaries.py` lists those routes and the parse-time messages that use them. Review Comment: The token-type list leaves out `dag_parse`, and the bundle rule doesn't match the code: the `/jobs` heartbeat and complete routes admit `dag_processor` with no bundle binding, while the routes that need binding are the Connection, Variable, XCom and Dag run routes that admit `dag_processor` or `dag_parse`. ########## airflow-core/src/airflow/jobs/job.py: ########## @@ -103,6 +104,10 @@ class Job(Base, LoggingMixin): hostname: Mapped[str | None] = mapped_column(String(500)) unixname: Mapped[str | None] = mapped_column(String(1000)) bundle_names: Mapped[list[str] | None] = mapped_column(ExtendedJSON, nullable=True) + session_id: Mapped[UUID | None] = mapped_column(Uuid(), nullable=True, unique=True) + """Component session that registered this Job over the Execution API; ``None`` for Jobs written directly.""" Review Comment: `register_job` also sets this to None on a replaced API Job (routes/jobs.py L180), so None doesn't only mean "written directly". ########## airflow-core/tests/unit/api_fastapi/execution_api/test_dag_processor_client.py: ########## @@ -0,0 +1,352 @@ +# 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. +from __future__ import annotations + +import socket +import threading +from datetime import datetime, timezone +from unittest import mock + +import httpx +import jwt +import pytest +import uvicorn +from sqlalchemy import select, update +from tenacity import wait_none +from uuid6 import uuid7 + +from airflow.api_fastapi.app import create_app +from airflow.api_fastapi.auth.tokens import JWTGenerator, JWTValidator +from airflow.api_fastapi.execution_api.app import lifespan +from airflow.api_fastapi.execution_api.datamodels.job import DagParseTokenBody, JobState, TerminalJobState +from airflow.dag_processing.api_client import ( + DagParseContext, + DagProcessorAPIClient, + DagProcessorRegistrationRetired, +) +from airflow.jobs.job import Job +from airflow.models.dagbundle import DagBundleModel +from airflow.models.variable import Variable +from airflow.sdk.api.client import Client + +from tests_common.test_utils.config import conf_vars +from tests_common.test_utils.db import ( + clear_db_jobs, + clear_db_variables, +) + +pytestmark = pytest.mark.db_test + +SECRET = "processor-client-test-secret-" * 3 +AUDIENCE = "urn:airflow.apache.org:task" +SESSION_ID = "00000000-0000-0000-0000-0000000000aa" +HEARTBEAT_EXPIRED = datetime(2026, 10, 5, 11, 0, tzinfo=timezone.utc) + + [email protected](autouse=True) +def clean_db(): + clear_db_jobs() + clear_db_variables() + yield + clear_db_jobs() + clear_db_variables() + + [email protected](autouse=True) +def freeze_time(time_machine): + time_machine.move_to("2026-10-05T12:00:00Z", tick=False) + + [email protected] +def api_requests(): Review Comment: `api_requests` is filled by the middleware below but no test asserts on it. Drop it, or assert which token each request carried? ########## 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: Heads-up for when #74233 installs this: the SDK's `_get_connection` / `_get_variable` check `SecretCache` before they reach SUPERVISOR_COMMS, and they look up without a team or bundle. The manager calls `SecretCache.init()` before forking parse processes, so with `[secrets] use_cache = True` a connection fetched under `use_bundle("a")` (or by a bundle-a parse) is served from the shared cache to bundle b, and the server-side bundle/team check never runs. The cache probably needs keying by bundle, or bypassing, on this path. ########## 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: Review Comment: The SDK's `raise_on_4xx_5xx` hook only reads the body when the response is `application/json`. For a non-JSON 403/409 (an ingress or proxy error page) the body is never read, so `error.response.json()` raises `httpx.ResponseNotRead`, a RuntimeError rather than ValueError. It escapes here, and since it isn't an `httpx.HTTPError` it also escapes the best-effort renewal handlers at L178 and L345. A local server returning a text/html 403 makes `client.variables.get()` raise `ResponseNotRead`. Catching `httpx.StreamError` here too would keep those paths best-effort. ########## airflow-core/src/airflow/api_fastapi/execution_api/routes/jobs.py: ########## @@ -0,0 +1,295 @@ +# 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. + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import svcs +from cadwyn import VersionedAPIRouter +from fastapi import HTTPException, Security, status +from sqlalchemy import insert, select, update +from sqlalchemy.exc import IntegrityError + +from airflow._shared.timezones import timezone +from airflow.api_fastapi.auth.tokens import JWTGenerator +from airflow.api_fastapi.common.db.common import SessionDep +from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc +from airflow.api_fastapi.execution_api.dag_processor_tokens import ( + ExpiredDagProcessorToken, + generate_dag_parse_token, + generate_dag_processor_token, +) +from airflow.api_fastapi.execution_api.datamodels.job import ( + DagParseTokenBody, + DagParseTokenResponse, + JobCompleteBody, + JobHeartbeatResponse, + JobRegisterBody, + JobRegisterResponse, +) +from airflow.api_fastapi.execution_api.datamodels.token import DagProcessorSessionToken, DagProcessorToken +from airflow.api_fastapi.execution_api.deps import DepContainer +from airflow.api_fastapi.execution_api.security import ( + JOB_UNCHECKED_SCOPE, + CurrentDagProcessorSessionToken, + CurrentDagProcessorToken, + ExecutionAPIRoute, + require_auth, +) +from airflow.configuration import conf +from airflow.jobs.dag_processor_job_runner import DagProcessorJobRunner +from airflow.jobs.job import Job, JobState +from airflow.models.dagbundle import DagBundleModel +from airflow.models.team import JobTeam + +if TYPE_CHECKING: + from uuid import UUID + + from sqlalchemy.orm import Session + +router = VersionedAPIRouter(route_class=ExecutionAPIRoute) + +_JOB_NOT_FOUND = create_openapi_http_exception_doc( + [(status.HTTP_404_NOT_FOUND, "Job not found for this token")] +) + + +def _check_token_job(job_id: int, token: DagProcessorToken) -> None: + if job_id != token.claims.job_id: + raise HTTPException( + status.HTTP_404_NOT_FOUND, + detail={"reason": "not_found", "message": f"Job {job_id} not found for this token"}, + ) + + +def _get_registered_job(registration_id: UUID, *, session: Session) -> Job | None: + return session.scalar(select(Job).where(Job.registration_id == registration_id).with_for_update()) + + +def _build_registration_response( + job_id: int, bundle_names: list[str], token: DagProcessorSessionToken, services: svcs.Container +) -> JobRegisterResponse: + try: + credential = generate_dag_processor_token( + services.get(JWTGenerator), + session_id=token.id, + job_id=job_id, + bundle_names=bundle_names, + session_expiry=token.claims.exp, + ) + except ExpiredDagProcessorToken as error: + raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(error)) from error + return JobRegisterResponse(job_id=job_id, token=credential) + + +def _resume_registration( + job: Job, + body: JobRegisterBody, + bundle_names: list[str], + token: DagProcessorSessionToken, + services: svcs.Container, +) -> JobRegisterResponse: + """Return the Job an earlier registration created, with a fresh token, if that Job is still the session's.""" + if job.session_id != token.id or job.end_date is not None: + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_retired", + "message": f"Registration {body.registration_id} has ended; a restarted processor uses a new one", + }, + ) + if (job.hostname, job.unixname, job.bundle_names) != (body.hostname, body.unixname, bundle_names): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_conflict", + "message": f"Registration {body.registration_id} was made with different details", + }, + ) + return _build_registration_response(job.id, bundle_names, token, services) + + [email protected]( + "", + status_code=status.HTTP_201_CREATED, + dependencies=[Security(require_auth, scopes=["token:dag_processor_session"])], + responses=create_openapi_http_exception_doc( + [ + ( + status.HTTP_403_FORBIDDEN, + "A requested bundle is not granted to the session, or its credential has expired", + ), + ( + status.HTTP_409_CONFLICT, + "The registration has ended, conflicts with an earlier one, or another process's Job is running", + ), + ] + ), +) +def register_job( + body: JobRegisterBody, + session: SessionDep, + token: DagProcessorSessionToken = CurrentDagProcessorSessionToken, + services: svcs.Container = DepContainer, +) -> JobRegisterResponse: + """ + Register the Job of a Dag processor session in exchange for its management credential. + + A registration creates one Job. Repeating it while that Job is open returns the Job with a fresh token; + once the Job completes or is replaced, the registration is refused for good. A new registration is + refused while the session's Job is alive, and otherwise ends and replaces it, which ends every token + issued for the replaced Job. + """ + requested = token.claims.dag_bundles if body.bundle_names is None else set(body.bundle_names) + if ungranted := requested - token.claims.dag_bundles: + raise HTTPException( + status.HTTP_403_FORBIDDEN, + detail={"reason": "bundle_not_granted", "message": f"Bundles not granted: {sorted(ungranted)}"}, + ) + bundle_names = sorted(requested) + + if registered := _get_registered_job(body.registration_id, session=session): + return _resume_registration(registered, body, bundle_names, token, services) + + now = timezone.utcnow() + previous = session.scalar(select(Job).where(Job.session_id == token.id).with_for_update()) + if previous is not None: + if previous.end_date is None and previous.is_alive(): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={"reason": "job_running", "message": f"Session already has running Job {previous.id}"}, + ) + if previous.end_date is None: + previous.state = JobState.FAILED + previous.end_date = now + previous.session_id = None + session.flush() + + try: + with session.begin_nested(): + # A core insert: Job.__init__ would stamp the API server's host and user and fire listeners. + session.execute( + insert(Job).values( + job_type=DagProcessorJobRunner.job_type, + state=JobState.RUNNING, + start_date=now, + latest_heartbeat=now, + hostname=body.hostname, + unixname=body.unixname, + bundle_names=bundle_names, + session_id=token.id, + registration_id=body.registration_id, + ) + ) + job_id = session.scalars(select(Job.id).where(Job.session_id == token.id)).one() + # Issue the credential before releasing the savepoint so an expiry failure + # also rolls back the new Job under SQLite's legacy transaction control. + response = _build_registration_response(job_id, bundle_names, token, services) + except IntegrityError: + # A retry sent before the original request finished can lose the race to it; resume the winner. + if registered := _get_registered_job(body.registration_id, session=session): + return _resume_registration(registered, body, bundle_names, token, services) + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={"reason": "job_running", "message": "Session registered another Job concurrently"}, + ) + if conf.getboolean("core", "multi_team"): + team_names = DagBundleModel.get_team_names(bundle_names, session=session) + session.add_all( + JobTeam(job_id=job_id, team_name=team) + for team in sorted({team for team in team_names.values() if team}) + ) + return response + + [email protected]( + "/{job_id}/heartbeat", + dependencies=[Security(require_auth, scopes=["token:dag_processor"])], + responses=_JOB_NOT_FOUND, +) +def heartbeat_job( + job_id: int, session: SessionDep, token: DagProcessorToken = CurrentDagProcessorToken +) -> JobHeartbeatResponse: + """Record a heartbeat and return the Job state, which tells the processor whether to stop.""" + _check_token_job(job_id, token) + job = session.scalars(select(Job).where(Job.id == job_id)).one() + job.latest_heartbeat = timezone.utcnow() + return JobHeartbeatResponse(state=JobState(job.state)) + + [email protected]( + "/{job_id}/complete", + status_code=status.HTTP_204_NO_CONTENT, + dependencies=[Security(require_auth, scopes=["token:dag_processor", JOB_UNCHECKED_SCOPE])], + responses=_JOB_NOT_FOUND, +) +def complete_job( + job_id: int, + body: JobCompleteBody, + session: SessionDep, + token: DagProcessorToken = CurrentDagProcessorToken, +) -> None: + """ + Record the final state of the Job, which ends every token issued for it. + + The first completion wins. Repeating it succeeds without changing the Job, so a retry after a lost + response is safe. + """ + _check_token_job(job_id, token) + owned_by_token = (Job.id == job_id) & (Job.session_id == token.id) + # One conditional statement, so overlapping completions cannot overwrite the first accepted outcome. + session.execute( + update(Job) + .where(owned_by_token, Job.end_date.is_(None)) + .values(end_date=timezone.utcnow(), state=JobState(body.state.value)) + ) + if session.scalar(select(Job.id).where(owned_by_token)) is None: + raise HTTPException( + status.HTTP_404_NOT_FOUND, + detail={"reason": "not_found", "message": f"Job {job_id} not found for this token"}, + ) + + [email protected]( + "/{job_id}/parse-token", + dependencies=[Security(require_auth, scopes=["token:dag_processor"])], + responses=_JOB_NOT_FOUND, +) +def exchange_parse_token( + job_id: int, + body: DagParseTokenBody, + token: DagProcessorToken = CurrentDagProcessorToken, + services: svcs.Container = DepContainer, +) -> DagParseTokenResponse: + """Bind runtime access to the bundle and file selected by the trusted processor manager.""" + _check_token_job(job_id, token) + if body.bundle_name not in token.claims.dag_bundles: + raise HTTPException(status.HTTP_403_FORBIDDEN, detail="Token is not granted this Dag bundle") Review Comment: This 403 and the two `ExpiredDagProcessorToken` ones (L96, L294) use a bare string `detail`, while the rest of this file uses `{"reason", "message"}` and the client branches on `reason`. The same bundle-not-granted condition is a dict at L162. Worth making these consistent before the contract ships. ########## airflow-core/src/airflow/cli/commands/dag_processor_token_command.py: ########## @@ -0,0 +1,76 @@ +# 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. +"""Dag processor token command.""" + +from __future__ import annotations + +import contextlib +import logging +import os +import tempfile +import time +from pathlib import Path + +import uuid6 + +from airflow.api_fastapi.execution_api.app import create_jwt_generator +from airflow.api_fastapi.execution_api.dag_processor_tokens import generate_dag_processor_session_token +from airflow.configuration import conf +from airflow.dag_processing.bundles.manager import _load_bundle_config_snapshot Review Comment: `_load_bundle_config_snapshot` is private to the bundles manager. A small public helper for the configured bundle names would keep the CLI off a private import. ########## airflow-core/src/airflow/api_fastapi/execution_api/security.py: ########## @@ -279,7 +330,122 @@ async def _require_live_attempt(token: TIToken, *, allow_callback: bool) -> None ) -CurrentTIToken: TIToken = Depends(require_auth) +async def _require_open_dag_processor_job( + request: Request, token_id: UUID, claims: DagProcessorClaims | DagParseClaims +) -> None: + """Refuse processor or parsing access after the Job ends or is replaced.""" + if request.scope.get(_REQUEST_SCOPE_JOB_KEY): + return + + from airflow.jobs.job import Job Review Comment: There's no cycle to avoid here: `airflow.jobs.job` doesn't import this module, and routes/jobs.py already imports it at the top. Same for the new `DagModel` / `DagBundleModel` / `Team` imports further down. Can these move to the top next to the `TaskInstance` and `Callback` imports? ########## airflow-core/src/airflow/api_fastapi/execution_api/routes/jobs.py: ########## @@ -0,0 +1,295 @@ +# 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. + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import svcs +from cadwyn import VersionedAPIRouter +from fastapi import HTTPException, Security, status +from sqlalchemy import insert, select, update +from sqlalchemy.exc import IntegrityError + +from airflow._shared.timezones import timezone +from airflow.api_fastapi.auth.tokens import JWTGenerator +from airflow.api_fastapi.common.db.common import SessionDep +from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc +from airflow.api_fastapi.execution_api.dag_processor_tokens import ( + ExpiredDagProcessorToken, + generate_dag_parse_token, + generate_dag_processor_token, +) +from airflow.api_fastapi.execution_api.datamodels.job import ( + DagParseTokenBody, + DagParseTokenResponse, + JobCompleteBody, + JobHeartbeatResponse, + JobRegisterBody, + JobRegisterResponse, +) +from airflow.api_fastapi.execution_api.datamodels.token import DagProcessorSessionToken, DagProcessorToken +from airflow.api_fastapi.execution_api.deps import DepContainer +from airflow.api_fastapi.execution_api.security import ( + JOB_UNCHECKED_SCOPE, + CurrentDagProcessorSessionToken, + CurrentDagProcessorToken, + ExecutionAPIRoute, + require_auth, +) +from airflow.configuration import conf +from airflow.jobs.dag_processor_job_runner import DagProcessorJobRunner +from airflow.jobs.job import Job, JobState +from airflow.models.dagbundle import DagBundleModel +from airflow.models.team import JobTeam + +if TYPE_CHECKING: + from uuid import UUID + + from sqlalchemy.orm import Session + +router = VersionedAPIRouter(route_class=ExecutionAPIRoute) + +_JOB_NOT_FOUND = create_openapi_http_exception_doc( + [(status.HTTP_404_NOT_FOUND, "Job not found for this token")] +) + + +def _check_token_job(job_id: int, token: DagProcessorToken) -> None: + if job_id != token.claims.job_id: + raise HTTPException( + status.HTTP_404_NOT_FOUND, + detail={"reason": "not_found", "message": f"Job {job_id} not found for this token"}, + ) + + +def _get_registered_job(registration_id: UUID, *, session: Session) -> Job | None: + return session.scalar(select(Job).where(Job.registration_id == registration_id).with_for_update()) + + +def _build_registration_response( + job_id: int, bundle_names: list[str], token: DagProcessorSessionToken, services: svcs.Container +) -> JobRegisterResponse: + try: + credential = generate_dag_processor_token( + services.get(JWTGenerator), + session_id=token.id, + job_id=job_id, + bundle_names=bundle_names, + session_expiry=token.claims.exp, + ) + except ExpiredDagProcessorToken as error: + raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(error)) from error + return JobRegisterResponse(job_id=job_id, token=credential) + + +def _resume_registration( + job: Job, + body: JobRegisterBody, + bundle_names: list[str], + token: DagProcessorSessionToken, + services: svcs.Container, +) -> JobRegisterResponse: + """Return the Job an earlier registration created, with a fresh token, if that Job is still the session's.""" + if job.session_id != token.id or job.end_date is not None: + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_retired", + "message": f"Registration {body.registration_id} has ended; a restarted processor uses a new one", + }, + ) + if (job.hostname, job.unixname, job.bundle_names) != (body.hostname, body.unixname, bundle_names): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_conflict", + "message": f"Registration {body.registration_id} was made with different details", + }, + ) + return _build_registration_response(job.id, bundle_names, token, services) + + [email protected]( + "", + status_code=status.HTTP_201_CREATED, + dependencies=[Security(require_auth, scopes=["token:dag_processor_session"])], + responses=create_openapi_http_exception_doc( + [ + ( + status.HTTP_403_FORBIDDEN, + "A requested bundle is not granted to the session, or its credential has expired", + ), + ( + status.HTTP_409_CONFLICT, + "The registration has ended, conflicts with an earlier one, or another process's Job is running", + ), + ] + ), +) +def register_job( + body: JobRegisterBody, + session: SessionDep, + token: DagProcessorSessionToken = CurrentDagProcessorSessionToken, + services: svcs.Container = DepContainer, +) -> JobRegisterResponse: + """ + Register the Job of a Dag processor session in exchange for its management credential. + + A registration creates one Job. Repeating it while that Job is open returns the Job with a fresh token; + once the Job completes or is replaced, the registration is refused for good. A new registration is + refused while the session's Job is alive, and otherwise ends and replaces it, which ends every token + issued for the replaced Job. + """ + requested = token.claims.dag_bundles if body.bundle_names is None else set(body.bundle_names) + if ungranted := requested - token.claims.dag_bundles: + raise HTTPException( + status.HTTP_403_FORBIDDEN, + detail={"reason": "bundle_not_granted", "message": f"Bundles not granted: {sorted(ungranted)}"}, + ) + bundle_names = sorted(requested) + + if registered := _get_registered_job(body.registration_id, session=session): + return _resume_registration(registered, body, bundle_names, token, services) + + now = timezone.utcnow() + previous = session.scalar(select(Job).where(Job.session_id == token.id).with_for_update()) + if previous is not None: + if previous.end_date is None and previous.is_alive(): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={"reason": "job_running", "message": f"Session already has running Job {previous.id}"}, + ) + if previous.end_date is None: + previous.state = JobState.FAILED + previous.end_date = now + previous.session_id = None + session.flush() + + try: + with session.begin_nested(): + # A core insert: Job.__init__ would stamp the API server's host and user and fire listeners. + session.execute( + insert(Job).values( + job_type=DagProcessorJobRunner.job_type, + state=JobState.RUNNING, + start_date=now, + latest_heartbeat=now, + hostname=body.hostname, + unixname=body.unixname, + bundle_names=bundle_names, + session_id=token.id, + registration_id=body.registration_id, + ) + ) + job_id = session.scalars(select(Job.id).where(Job.session_id == token.id)).one() Review Comment: The insert's result already has `inserted_primary_key`, so this second SELECT isn't needed. ########## airflow-core/src/airflow/api_fastapi/execution_api/security.py: ########## @@ -91,14 +125,23 @@ log = structlog.get_logger(logger_name=__name__) VALID_TOKEN_TYPES: frozenset[str] = frozenset(get_args(TokenScope)) +_TOKEN_MODELS: dict[str, type[ExecutionToken]] = { + **dict.fromkeys(get_args(TaskTokenScope), TIToken), + "dag_processor_session": DagProcessorSessionToken, + "dag_processor": DagProcessorToken, + "dag_parse": ExecutionToken, Review Comment: Every other scope gets its own token class, but `dag_parse` maps to plain `ExecutionToken`. So the processor checks in `get_selected_dag_bundle`, `require_dag_in_granted_bundle`, `get_team_name_dep` and `has_xcom_access` compare `token.claims.scope` strings, and `token.claims.dag_bundles` isn't type-checked because those dependencies take an untyped `token`. A `DagParseToken` with `claims: DagParseClaims` would let them narrow with isinstance and give mypy something to check. ########## airflow-core/src/airflow/api_fastapi/execution_api/routes/jobs.py: ########## @@ -0,0 +1,295 @@ +# 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. + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import svcs +from cadwyn import VersionedAPIRouter +from fastapi import HTTPException, Security, status +from sqlalchemy import insert, select, update +from sqlalchemy.exc import IntegrityError + +from airflow._shared.timezones import timezone +from airflow.api_fastapi.auth.tokens import JWTGenerator +from airflow.api_fastapi.common.db.common import SessionDep +from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc +from airflow.api_fastapi.execution_api.dag_processor_tokens import ( + ExpiredDagProcessorToken, + generate_dag_parse_token, + generate_dag_processor_token, +) +from airflow.api_fastapi.execution_api.datamodels.job import ( + DagParseTokenBody, + DagParseTokenResponse, + JobCompleteBody, + JobHeartbeatResponse, + JobRegisterBody, + JobRegisterResponse, +) +from airflow.api_fastapi.execution_api.datamodels.token import DagProcessorSessionToken, DagProcessorToken +from airflow.api_fastapi.execution_api.deps import DepContainer +from airflow.api_fastapi.execution_api.security import ( + JOB_UNCHECKED_SCOPE, + CurrentDagProcessorSessionToken, + CurrentDagProcessorToken, + ExecutionAPIRoute, + require_auth, +) +from airflow.configuration import conf +from airflow.jobs.dag_processor_job_runner import DagProcessorJobRunner +from airflow.jobs.job import Job, JobState +from airflow.models.dagbundle import DagBundleModel +from airflow.models.team import JobTeam + +if TYPE_CHECKING: + from uuid import UUID + + from sqlalchemy.orm import Session + +router = VersionedAPIRouter(route_class=ExecutionAPIRoute) + +_JOB_NOT_FOUND = create_openapi_http_exception_doc( + [(status.HTTP_404_NOT_FOUND, "Job not found for this token")] +) + + +def _check_token_job(job_id: int, token: DagProcessorToken) -> None: + if job_id != token.claims.job_id: + raise HTTPException( + status.HTTP_404_NOT_FOUND, + detail={"reason": "not_found", "message": f"Job {job_id} not found for this token"}, + ) + + +def _get_registered_job(registration_id: UUID, *, session: Session) -> Job | None: + return session.scalar(select(Job).where(Job.registration_id == registration_id).with_for_update()) + + +def _build_registration_response( + job_id: int, bundle_names: list[str], token: DagProcessorSessionToken, services: svcs.Container +) -> JobRegisterResponse: + try: + credential = generate_dag_processor_token( + services.get(JWTGenerator), + session_id=token.id, + job_id=job_id, + bundle_names=bundle_names, + session_expiry=token.claims.exp, + ) + except ExpiredDagProcessorToken as error: + raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(error)) from error + return JobRegisterResponse(job_id=job_id, token=credential) + + +def _resume_registration( + job: Job, + body: JobRegisterBody, + bundle_names: list[str], + token: DagProcessorSessionToken, + services: svcs.Container, +) -> JobRegisterResponse: + """Return the Job an earlier registration created, with a fresh token, if that Job is still the session's.""" + if job.session_id != token.id or job.end_date is not None: + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_retired", + "message": f"Registration {body.registration_id} has ended; a restarted processor uses a new one", + }, + ) + if (job.hostname, job.unixname, job.bundle_names) != (body.hostname, body.unixname, bundle_names): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={ + "reason": "registration_conflict", + "message": f"Registration {body.registration_id} was made with different details", + }, + ) + return _build_registration_response(job.id, bundle_names, token, services) + + [email protected]( + "", + status_code=status.HTTP_201_CREATED, + dependencies=[Security(require_auth, scopes=["token:dag_processor_session"])], + responses=create_openapi_http_exception_doc( + [ + ( + status.HTTP_403_FORBIDDEN, + "A requested bundle is not granted to the session, or its credential has expired", + ), + ( + status.HTTP_409_CONFLICT, + "The registration has ended, conflicts with an earlier one, or another process's Job is running", + ), + ] + ), +) +def register_job( + body: JobRegisterBody, + session: SessionDep, + token: DagProcessorSessionToken = CurrentDagProcessorSessionToken, + services: svcs.Container = DepContainer, +) -> JobRegisterResponse: + """ + Register the Job of a Dag processor session in exchange for its management credential. + + A registration creates one Job. Repeating it while that Job is open returns the Job with a fresh token; + once the Job completes or is replaced, the registration is refused for good. A new registration is + refused while the session's Job is alive, and otherwise ends and replaces it, which ends every token + issued for the replaced Job. + """ + requested = token.claims.dag_bundles if body.bundle_names is None else set(body.bundle_names) + if ungranted := requested - token.claims.dag_bundles: + raise HTTPException( + status.HTTP_403_FORBIDDEN, + detail={"reason": "bundle_not_granted", "message": f"Bundles not granted: {sorted(ungranted)}"}, + ) + bundle_names = sorted(requested) + + if registered := _get_registered_job(body.registration_id, session=session): + return _resume_registration(registered, body, bundle_names, token, services) + + now = timezone.utcnow() + previous = session.scalar(select(Job).where(Job.session_id == token.id).with_for_update()) + if previous is not None: + if previous.end_date is None and previous.is_alive(): + raise HTTPException( + status.HTTP_409_CONFLICT, + detail={"reason": "job_running", "message": f"Session already has running Job {previous.id}"}, + ) + if previous.end_date is None: + previous.state = JobState.FAILED + previous.end_date = now + previous.session_id = None + session.flush() + + try: + with session.begin_nested(): + # A core insert: Job.__init__ would stamp the API server's host and user and fire listeners. + session.execute( + insert(Job).values( + job_type=DagProcessorJobRunner.job_type, Review Comment: `POST /jobs` reads as generic, but it only creates a DagProcessorJob, takes `bundle_names`, and only accepts session tokens. If the triggerer or another component registers over the API later, this path is already taken. With the open thread on token.py about renaming `dag_processor_session` to `dag_processor_registration`, which would then sit next to the `session_id` / `registration_id` Job columns from 0143, might be worth settling the route path, scope name and column names together before v2027_02_28 and the migration are released. ########## airflow-core/src/airflow/api_fastapi/execution_api/datamodels/job.py: ########## @@ -0,0 +1,92 @@ +# 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. + +from __future__ import annotations + +from enum import Enum +from uuid import UUID + +from pydantic import Field, field_validator + +from airflow.api_fastapi.core_api.base import StrictBaseModel +from airflow.jobs.job import JobState + + +class TerminalJobState(str, Enum): + """States a Job can finish in.""" + + SUCCESS = JobState.SUCCESS.value + FAILED = JobState.FAILED.value + + +class JobRegisterBody(StrictBaseModel): + """Request body a Dag processor sends to register the Job of its session.""" + + registration_id: UUID = Field( + description=( + "Chosen by the processor once per process start. Registering again with the same id while its Job " + "is open returns that Job with a fresh token, which recovers a lost response and renews the " + "token. Once the Job completes or is replaced the id is refused, so a restart chooses a new one." + ) + ) + hostname: str = Field(min_length=1, max_length=500) + unixname: str | None = Field(default=None, max_length=1000) + bundle_names: list[str] | None = Field( + default=None, + min_length=1, + description="Bundles the processor parses; defaults to every bundle its token grants.", + ) + + +class JobRegisterResponse(StrictBaseModel): + """The registered Job and its management credential.""" + + job_id: int + token: str = Field(description="A ``dag_processor`` token valid while the Job is open.") Review Comment: This token expires after `jwt_expiration_time` (or sooner, capped by the session token), not when the Job closes. Maybe "valid until it expires or the Job ends"? Returning an `expires_in` would also let the client drop the two unverified `jwt.decode` blocks it uses only to schedule renewal. ########## airflow-core/src/airflow/api_fastapi/execution_api/app.py: ########## @@ -149,10 +151,11 @@ async def dispatch(self, request: Request, call_next): validator: JWTValidator = await services.aget(JWTValidator) claims = await validator.avalidated_claims(token, {}) - # Workload and callback tokens are long-lived and meant to survive - # queue wait times so avoid refreshing them. If avalidated_claims - # raises for such a token, the outer except handles it. - if claims.get("scope") in ("workload", "callback"): + # Only short-lived execution tokens are renewed here. Any other type has a + # lifetime set by its issuer (workload and callback tokens outlive queue waits, + # dag_processor tokens are rotated by provisioning), so a new type must not Review Comment: Only the session token is rotated by provisioning. `dag_processor` and `dag_parse` tokens are renewed by re-registering and re-exchanging. ########## airflow-core/tests/unit/dag_processing/test_api_client.py: ########## @@ -0,0 +1,958 @@ +# 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. +from __future__ import annotations + +import json +import runpy +from contextvars import Context +from unittest import mock + +import httpx +import jwt +import pytest +from tenacity import wait_none + +from airflow.api_fastapi.execution_api.datamodels.job import DagParseTokenBody, JobState, TerminalJobState +from airflow.dag_processing import api_client +from airflow.dag_processing.api_client import ( + DagParseContext, + DagProcessorAPIClient, + DagProcessorJobAlreadyRunning, + DagProcessorRegistrationRetired, + DagProcessorSecretsComms, +) +from airflow.sdk.api.client import API_RETRIES, Client +from airflow.sdk.api.datamodels import _generated +from airflow.sdk.execution_time.comms import GetVariable, GetXCom, MaskSecret, VariableResult + +SECRET = "processor-client-unit-test-signing-key" + + +def make_job_token(**claims): + return jwt.encode( + {"iat": 0, "exp": 300, "scope": "dag_processor", "job_id": 1, **claims}, + SECRET, + ) + + +def make_registration_response(*, job_id=1, token=None, **kwargs): + return httpx.Response( + 201, json={"job_id": job_id, "token": token or make_job_token(job_id=job_id)}, **kwargs + ) + + [email protected] +def token_file(tmp_path): + path = tmp_path / "processor.jwt" + path.write_text("session-1\n") + return path + + [email protected](autouse=True) +def no_retry_wait(monkeypatch): + monkeypatch.setattr(Client._request_with_retry.retry, "wait", wait_none()) + + [email protected] +def clock(): + with mock.patch("airflow.dag_processing.api_client.monotonic", autospec=True, return_value=0) as clock: + yield clock + + [email protected] +def make_client(token_file, clock): + clients = [] + + def create(*responses, **kwargs): + pending = iter(responses) + requests = [] + + def handle(request): + requests.append(request) + response = next(pending) + if isinstance(response, Exception): + raise response + return response(request) if callable(response) else response + + client = DagProcessorAPIClient( + base_url="http://api/execution/", + token_file=token_file, + hostname="processor-1", + transport=httpx.MockTransport(handle), + **kwargs, + ) + clients.append(client) + return client, requests + + yield create + for client in clients: + client.close() + + +def test_registration_and_runtime_use_distinct_credentials(make_client): + client, requests = make_client( + make_registration_response(headers={"Refreshed-API-Token": "wrong-token"}), + httpx.Response(200, json={"state": "restarting"}, headers={"Refreshed-API-Token": "wrong-token"}), + httpx.Response(200, json={"state": "running"}), + unixname="airflow", + bundle_names=["bundle-a"], + ) + + assert client.register_job() == client.job_id == 1 + assert client.heartbeat() == JobState.RESTARTING + assert client.heartbeat() == JobState.RUNNING + assert json.loads(requests[0].content) == { + "registration_id": str(client.registration_id), + "hostname": "processor-1", + "unixname": "airflow", + "bundle_names": ["bundle-a"], + } + assert [request.headers["Authorization"] for request in requests] == [ + "Bearer session-1", + f"Bearer {make_job_token()}", + f"Bearer {make_job_token()}", + ] + assert requests[1].url.path == "/execution/jobs/1/heartbeat" + + +def test_registration_identity_survives_response_loss_but_not_a_process_restart(make_client): + first, requests = make_client(httpx.ReadError("Lost acknowledgment"), make_registration_response()) + second, _ = make_client() + + first.register_job() + + assert first.registration_id != second.registration_id + assert len(requests) == 2 + assert requests[0].content == requests[1].content + + [email protected]("succeeds", [False, True]) [email protected]("status", [401, 403]) +def test_registration_reloads_a_rotated_credential_after_authentication_failure( + make_client, token_file, status, succeeds +): + def rotate(request): + token_file.write_text("session-2") + return httpx.Response(status) + + client, requests = make_client( + rotate, make_registration_response() if succeeds else httpx.Response(status) + ) + + if succeeds: + client.register_job() + else: + with pytest.raises(httpx.HTTPStatusError): + client.register_job() + + assert [request.headers["Authorization"] for request in requests] == [ + "Bearer session-1", + "Bearer session-2", + ] + assert requests[0].content == requests[1].content + + [email protected]("status", [401, 403, 404, 409]) +def test_registration_errors_do_not_start_another_job(make_client, status): + client, requests = make_client(httpx.Response(status)) + registration_id = client.registration_id + + with pytest.raises(httpx.HTTPStatusError) as error: + client.register_job() + + assert error.value.response.status_code == status + assert client.job_id is None + assert client.registration_id == registration_id + assert len(requests) == 1 + + +def test_registration_retry_after_a_crashed_job_keeps_its_identity(make_client): + client, requests = make_client( + httpx.Response(409, json={"detail": {"reason": "job_running"}}), + make_registration_response(), + ) + registration_id = client.registration_id + + with pytest.raises(DagProcessorJobAlreadyRunning): + client.register_job() + assert client.job_id is None + assert client.register_job() == 1 + + assert client.registration_id == registration_id + assert requests[0].content == requests[1].content + + +def test_rotation_is_detected_within_the_cache_interval(make_client, token_file, clock): + client, requests = make_client( + make_registration_response(), + httpx.Response(204), + make_registration_response(token=make_job_token(iat=30, exp=330)), + httpx.Response(204), + ) + client.register_job() + token_file.write_text("session-2") + clock.return_value = 29 + client.get("resource") + clock.return_value = 30 + client.get("resource") + + assert [request.url.path for request in requests] == [ + "/execution/jobs", + "/execution/resource", + "/execution/jobs", + "/execution/resource", + ] + assert requests[2].headers["Authorization"] == "Bearer session-2" + assert requests[3].headers["Authorization"] == f"Bearer {make_job_token(iat=30, exp=330)}" + + +def test_job_token_renewal_uses_monotonic_time(make_client, clock, time_machine): + client, requests = make_client( + make_registration_response(), httpx.Response(204), make_registration_response(), httpx.Response(204) + ) + client.register_job() + time_machine.move_to("2030-01-01", tick=False) + clock.return_value = 239 + client.get("resource") + time_machine.move_to("2020-01-01", tick=False) + clock.return_value = 240 + client.get("resource") + + assert [request.url.path for request in requests].count("/execution/jobs") == 2 + assert requests[0].content == requests[2].content + + [email protected]("status", [401, 403, 404, 409]) +def test_runtime_errors_do_not_renew_credentials(make_client, status): + client, requests = make_client(make_registration_response(), httpx.Response(status)) + client.register_job() + + with pytest.raises(httpx.HTTPStatusError) as error: + client.get("resource") + + assert error.value.response.status_code == status + assert len(requests) == 2 + + [email protected]("expired", [False, True]) [email protected]("failure", [httpx.ReadError("Unavailable"), httpx.Response(503)]) +def test_heartbeat_does_not_retry_transport_failures(make_client, clock, expired, failure): + client, requests = make_client(make_registration_response(), failure) + client.register_job() + if expired: + clock.return_value = 300 + + with pytest.raises(httpx.HTTPError): + client.heartbeat() + + assert len(requests) == 2 + + [email protected]("registered", [False, True]) [email protected]("missing", [False, True]) +def test_unreadable_credentials_fail_closed(make_client, token_file, clock, registered, missing): + client, requests = make_client(make_registration_response()) + if registered: + client.register_job() + if missing: + token_file.unlink() + else: + token_file.write_text(" \n") + clock.return_value = 300 + + with pytest.raises(FileNotFoundError if missing else ValueError): + client.heartbeat() if registered else client.register_job() + + assert len(requests) == int(registered) + + +def test_renewal_cannot_replace_the_job(make_client): + client, _ = make_client(make_registration_response(), make_registration_response(job_id=2)) + client.register_job() + + with pytest.raises(RuntimeError, match="different Dag processor Job"): + client.register_job() + + assert client.job_id == 1 + assert client.auth.token == make_job_token() + + [email protected]( + "token", + [ + "not-a-jwt", + make_job_token(scope="execution"), + make_job_token(job_id=2), + make_job_token(exp=0), + make_job_token(exp="not-a-timestamp"), + make_job_token(exp=float("inf")), + jwt.encode({"scope": "dag_processor"}, SECRET), + ], +) +def test_invalid_job_tokens_are_not_adopted(make_client, token): + client, _ = make_client(make_registration_response(token=token)) + + with pytest.raises(ValueError, match="invalid Dag processor Job token"): + client.register_job() + + assert client.job_id is None + assert client.auth.token == "" + + [email protected]( + "operation", + [ + pytest.param(lambda client: client.get("resource"), id="request"), + pytest.param(lambda client: client.heartbeat(), id="heartbeat"), + pytest.param(lambda client: client.complete_job(TerminalJobState.SUCCESS), id="complete"), + ], +) +def test_runtime_operations_require_registration(make_client, operation): + client, requests = make_client() + + with pytest.raises(RuntimeError, match="Register the Dag processor Job"): + operation(client) + + assert requests == [] + + +def test_completion_recovers_a_lost_acknowledgment_with_the_same_token_and_outcome(make_client): + client, requests = make_client( + make_registration_response(), httpx.ReadError("Lost acknowledgment"), httpx.Response(204) + ) + client.register_job() + + client.complete_job(TerminalJobState.SUCCESS) + client.complete_job(TerminalJobState.SUCCESS) + + assert len(requests) == 3 + assert requests[1].url.path == "/execution/jobs/1/complete" + assert requests[1].content == requests[2].content == b'{"state":"success"}' + assert requests[1].headers["Authorization"] == requests[2].headers["Authorization"] + + +def test_completion_retries_remain_possible_after_transport_retries_are_exhausted( + make_client, token_file, clock +): + client, requests = make_client( + make_registration_response(), + *[httpx.ReadError("Lost acknowledgment") for _ in range(API_RETRIES)], + httpx.Response(204), + ) + client.register_job() + with pytest.raises(httpx.ReadError): + client.complete_job(TerminalJobState.FAILED) + token_file.unlink() + clock.return_value = 290 + + client.complete_job(TerminalJobState.FAILED) + + assert len(requests) == API_RETRIES + 2 + assert {request.content for request in requests[1:]} == {b'{"state":"failed"}'} + assert {request.headers["Authorization"] for request in requests[1:]} == {f"Bearer {make_job_token()}"} + + [email protected]("status", [204, 401, 403, 404]) +def test_completion_stops_normal_work_even_when_acknowledgment_is_uncertain(make_client, status): + client, requests = make_client(make_registration_response(), httpx.Response(status)) + client.register_job() + if status == 204: + client.complete_job(TerminalJobState.SUCCESS) + else: + with pytest.raises(httpx.HTTPStatusError): + client.complete_job(TerminalJobState.SUCCESS) + + with pytest.raises(RuntimeError, match="only completion retries"): + client.heartbeat() + with pytest.raises(RuntimeError, match="only completion retries"): + client.register_job() + with pytest.raises(ValueError, match="original outcome"): + client.complete_job(TerminalJobState.FAILED) + assert len(requests) == 2 + + [email protected]("interval", [-1, float("inf"), float("nan")]) +def test_invalid_reload_interval_is_rejected(make_client, interval): + with pytest.raises(ValueError, match="token_reload_interval"): + make_client(token_reload_interval=interval) + + [email protected]("operation", ["runtime", "heartbeat"]) [email protected]( + "failure", [httpx.ReadTimeout("Unavailable"), httpx.Response(503), httpx.Response(403)] +) +def test_failed_early_renewal_keeps_using_the_valid_token(make_client, clock, operation, failure): + response_body = {"state": "running"} if operation == "heartbeat" else {"key": "key", "value": "value"} + client, requests = make_client( + make_registration_response(), + failure, + httpx.Response(200, json=response_body), + httpx.Response(200, json=response_body), + make_registration_response(), + httpx.Response(204), + ) + client.register_job() + clock.return_value = 240 + + for _ in range(2): + if operation == "runtime": + assert client.variables.get("key").value == "value" + else: + assert client.heartbeat() == JobState.RUNNING + + assert len(requests) == 4 + assert requests[1].extensions["timeout"] == {"connect": 1, "read": 1, "write": 1, "pool": 1} + assert [request.headers["Authorization"] for request in requests[2:]] == [ + f"Bearer {make_job_token()}" + ] * 2 + clock.return_value = 270 + client.get("resource") + assert requests[4].url.path == "/execution/jobs" + + [email protected]("missing", [False, True]) +def test_unreadable_rotated_file_does_not_discard_a_valid_token(make_client, token_file, clock, missing): + client, requests = make_client( + make_registration_response(), httpx.Response(200, json={"state": "running"}) + ) + client.register_job() + token_file.unlink() if missing else token_file.write_text("") + clock.return_value = 30 + + assert client.heartbeat() == JobState.RUNNING + assert len(requests) == 2 + + +def test_renewal_that_outlasts_the_token_requires_synchronous_recovery(make_client, clock): + def time_out(request): + clock.return_value = 300 + raise httpx.ReadTimeout("Outlasted the token") + + client, requests = make_client( + make_registration_response(), time_out, make_registration_response(), httpx.Response(204) + ) + client.register_job() + clock.return_value = 299 + + client.get("resource") + + assert [request.url.path for request in requests] == ["/execution/jobs"] * 3 + ["/execution/resource"] + + [email protected]("complete", [False, True]) +def test_retired_registration_signals_restart_without_blocking_valid_requests(make_client, clock, complete): + client, requests = make_client( + make_registration_response(), + httpx.Response(409, json={"detail": {"reason": "registration_retired"}}), + httpx.Response(204), + httpx.Response(204), + httpx.Response(204), + ) + client.register_job() + clock.return_value = 240 + client.get("resource") + assert client.restart_required + client.get("resource") + if complete: + client.complete_job(TerminalJobState.SUCCESS) + assert requests[-1].url.path == "/execution/jobs/1/complete" + assert requests[-1].headers["Authorization"] == f"Bearer {make_job_token()}" + else: + clock.return_value = 300 + with pytest.raises(DagProcessorRegistrationRetired): + client.get("resource") + assert len(requests) == 4 + complete + + [email protected]( + "payload", [{"detail": "conflict"}, ["unexpected"], {"detail": {"reason": "registration_conflict"}}] +) +def test_other_conflicts_are_not_mistaken_for_retirement(make_client, payload): + client, _ = make_client(httpx.Response(409, json=payload)) + with pytest.raises(httpx.HTTPStatusError): + client.register_job() + assert not client.restart_required + + +def test_closed_job_signals_restart(make_client): + client, _ = make_client( + make_registration_response(), + httpx.Response(403, json={"detail": {"reason": "job_closed", "message": "replaced"}}), + ) + client.register_job() + + with pytest.raises(DagProcessorRegistrationRetired): + client.heartbeat() + + assert client.restart_required + + [email protected]( + "payload", [{"detail": "Invalid auth token"}, {"detail": {"reason": "bundle_not_granted"}}] +) +def test_other_authorization_failures_are_not_mistaken_for_a_closed_job(make_client, payload): + client, _ = make_client(make_registration_response(), httpx.Response(403, json=payload)) + client.register_job() + + with pytest.raises(httpx.HTTPStatusError): + client.get("resource") + + assert not client.restart_required + + +def test_sdk_helpers_keep_bundle_context_for_reads_and_writes(make_client): + client, requests = make_client( + make_registration_response(), + httpx.Response(200, json={"key": "key", "value": "value"}), + httpx.Response(204), + httpx.Response( + 200, + json={ + "conn_id": "connection", + "conn_type": "generic", + "host": None, + "schema": None, + "login": None, + "password": None, + "port": None, + "extra": None, + }, + ), + ) + client.register_job() + with client.use_bundle("bundle-a"): + assert client.variables.get("key").value == "value" + assert client.variables.set("key", "new").ok + assert client.connections.get("connection").conn_id == "connection" + + assert [request.headers["Airflow-Dag-Bundle"] for request in requests[1:]] == ["bundle-a"] * 3 + assert requests[2].headers["Content-Type"] == "application/json" + assert json.loads(requests[2].content)["val"] == "new" + + +def test_bundle_context_preserves_headers_and_restores_context_on_error(make_client): + client, requests = make_client( + make_registration_response(), + httpx.Response(204), + httpx.ReadError("Subprocess failure"), + *[httpx.Response(204) for _ in range(3)], + ) + client.register_job() + with client.use_bundle("bundle-a"): + client.put("resource", content=b"data", headers={"content-type": "text/plain", "custom": "value"}) + with pytest.raises(httpx.ReadError), client.use_bundle("bundle-b"): + client.request( + "PUT", + "resource", + content=b"{}", + headers={"custom": "value", "airflow-dag-bundle": "ignored"}, + retry=False, + ) + client.get("resource") + Context().run(client.get, "resource") + client.get("resource") + + assert [request.headers.get("Airflow-Dag-Bundle") for request in requests[1:]] == [ + "bundle-a", + "bundle-b", + "bundle-a", + None, + None, + ] + assert requests[1].headers["Content-Type"] == "text/plain" + assert requests[2].headers["Content-Type"] == "application/json" + assert requests[1].headers["custom"] == requests[2].headers["custom"] == "value" + + +def test_empty_bundle_context_is_rejected(make_client): + client, requests = make_client() + with pytest.raises(ValueError, match="nonempty bundle"), client.use_bundle(""): + pytest.fail("Empty bundle was accepted") + assert not requests + + [email protected]("lost_acknowledgment", [False, True]) +def test_completion_renews_an_expired_token_if_the_job_is_still_open(make_client, clock, lost_acknowledgment): + failures = ( + [httpx.ReadError("Lost acknowledgment") for _ in range(API_RETRIES)] if lost_acknowledgment else [] + ) + renewed_token = make_job_token(iat=300, exp=600) + client, requests = make_client( + make_registration_response(), + *failures, + make_registration_response(token=renewed_token), + httpx.Response(204), + ) + client.register_job() + if lost_acknowledgment: + with pytest.raises(httpx.ReadError): + client.complete_job(TerminalJobState.SUCCESS) + clock.return_value = 300 + + client.complete_job(TerminalJobState.SUCCESS) + + assert requests[0].content == requests[-2].content + assert requests[-1].headers["Authorization"] == f"Bearer {renewed_token}" + assert {request.content for request in requests if request.url.path.endswith("/complete")} == { + b'{"state":"success"}' + } + + +def test_completion_does_not_claim_success_when_expired_registration_is_retired(make_client, clock): + client, requests = make_client( + make_registration_response(), httpx.Response(409, json={"detail": {"reason": "registration_retired"}}) + ) + client.register_job() + clock.return_value = 300 + + for _ in range(2): + with pytest.raises(DagProcessorRegistrationRetired): + client.complete_job(TerminalJobState.SUCCESS) + + assert client.restart_required + assert len(requests) == 2 + + [email protected]("succeeds", [False, True]) [email protected]("status", [401, 403]) +def test_completion_recovers_if_the_token_expires_during_the_request(make_client, clock, status, succeeds): + def expire(request): + clock.return_value += 300 + return httpx.Response(status) + + client, requests = make_client( + make_registration_response(), + expire, + make_registration_response(), + httpx.Response(204) if succeeds else expire, + ) + client.register_job() + if succeeds: + client.complete_job(TerminalJobState.SUCCESS) + else: + with pytest.raises(httpx.HTTPStatusError): + client.complete_job(TerminalJobState.SUCCESS) + + assert requests[0].content == requests[2].content + assert requests[1].content == requests[3].content + + +def test_control_contracts_do_not_require_new_sdk_models(monkeypatch, token_file): + for name in ( + "JobRegisterBody", Review Comment: api_client also imports `DagParseTokenBody` and `DagParseTokenResponse` from the core datamodels, and both exist in `_generated`. Adding them to this list keeps the test covering every contract the client uses. -- 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]
