Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/core_api/security.py: 80%
424 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
1# Licensed to the Apache Software Foundation (ASF) under one
2# or more contributor license agreements. See the NOTICE file
3# distributed with this work for additional information
4# regarding copyright ownership. The ASF licenses this file
5# to you under the Apache License, Version 2.0 (the
6# "License"); you may not use this file except in compliance
7# with the License. You may obtain a copy of the License at
8#
9# http://www.apache.org/licenses/LICENSE-2.0
10#
11# Unless required by applicable law or agreed to in writing,
12# software distributed under the License is distributed on an
13# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14# KIND, either express or implied. See the License for the
15# specific language governing permissions and limitations
16# under the License.
17from __future__ import annotations
19import posixpath
20from collections.abc import Callable, Coroutine
21from contextlib import suppress
22from json import JSONDecodeError
23from typing import TYPE_CHECKING, Annotated, Any, cast
24from urllib.parse import ParseResult, unquote, urljoin, urlparse
26from fastapi import Depends, HTTPException, Request, status
27from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer, OAuth2PasswordBearer
28from itsdangerous import BadSignature, URLSafeSerializer
29from jwt import ExpiredSignatureError, InvalidTokenError
30from pydantic import NonNegativeInt, TypeAdapter, ValidationError
31from sqlalchemy import or_, select
32from sqlalchemy.orm import Session
34from airflow.api_fastapi.app import get_auth_manager
35from airflow.api_fastapi.auth.managers.base_auth_manager import (
36 COOKIE_NAME_JWT_TOKEN,
37 BaseAuthManager,
38)
39from airflow.api_fastapi.auth.managers.models.base_user import BaseUser
40from airflow.api_fastapi.auth.managers.models.batch_apis import (
41 IsAuthorizedConnectionRequest,
42 IsAuthorizedDagRequest,
43 IsAuthorizedPoolRequest,
44 IsAuthorizedVariableRequest,
45)
46from airflow.api_fastapi.auth.managers.models.resource_details import (
47 AccessView,
48 AssetAliasDetails,
49 AssetDetails,
50 ConfigurationDetails,
51 ConnectionDetails,
52 DagAccessEntity,
53 DagDetails,
54 PoolDetails,
55 VariableDetails,
56)
57from airflow.api_fastapi.common.db.common import SessionDep
58from airflow.api_fastapi.core_api.base import OrmClause
59from airflow.api_fastapi.core_api.datamodels.common import (
60 BulkAction,
61 BulkActionOnExistence,
62 BulkBody,
63 BulkCreateAction,
64 BulkDeleteAction,
65 BulkUpdateAction,
66)
67from airflow.api_fastapi.core_api.datamodels.connections import ConnectionBody
68from airflow.api_fastapi.core_api.datamodels.dag_run import BulkDAGRunBody, BulkDAGRunClearBody
69from airflow.api_fastapi.core_api.datamodels.pools import PoolBody
70from airflow.api_fastapi.core_api.datamodels.variables import VariableBody
71from airflow.configuration import conf
72from airflow.models import Connection, Pool, Variable
73from airflow.models.asset import AssetEvent
74from airflow.models.backfill import Backfill
75from airflow.models.dag import DagModel, DagRun, DagTag
76from airflow.models.dag_version import DagVersion
77from airflow.models.dagwarning import DagWarning
78from airflow.models.log import Log
79from airflow.models.taskinstance import TaskInstance as TI
80from airflow.models.team import Team
81from airflow.models.xcom import XComModel
83if TYPE_CHECKING: 83 ↛ 84line 83 didn't jump to line 84 because the condition on line 83 was never true
84 from sqlalchemy.sql import Select
86 from airflow.api_fastapi.auth.managers.base_auth_manager import ResourceMethod
89def auth_manager_from_app(request: Request) -> BaseAuthManager:
90 """
91 FastAPI dependency resolver that returns the shared AuthManager instance from app.state.
93 This ensures that all API routes using AuthManager via dependency injection receive the same
94 singleton instance that was initialized at app startup.
95 """
96 return request.app.state.auth_manager
99AuthManagerDep = Annotated[BaseAuthManager, Depends(auth_manager_from_app)]
101auth_description = (
102 "To authenticate Airflow API requests, clients must include a JWT (JSON Web Token) in "
103 "the Authorization header of each request. This token is used to verify the identity of "
104 "the client and ensure that they have the appropriate permissions to access the "
105 "requested resources. "
106 "You can use the endpoint ``POST /auth/token`` in order to generate a JWT token. "
107 "Upon successful authentication, the server will issue a JWT token that contains the necessary "
108 "information (such as user identity and scope) to authenticate subsequent requests. "
109 "To learn more about Airflow public API authentication, please read https://airflow.apache.org/docs/apache-airflow/stable/security/api.html."
110)
111oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/auth/token", description=auth_description, auto_error=False)
112bearer_scheme = HTTPBearer(auto_error=False)
114MAP_BULK_ACTION_TO_AUTH_METHOD: dict[BulkAction, ResourceMethod] = {
115 BulkAction.CREATE: "POST",
116 BulkAction.DELETE: "DELETE",
117 BulkAction.UPDATE: "PUT",
118}
121async def resolve_user_from_token(token_str: str | None) -> BaseUser:
122 if not token_str:
123 raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
125 try:
126 return await get_auth_manager().get_user_from_token(token_str)
127 except ExpiredSignatureError:
128 raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Token Expired")
129 except InvalidTokenError:
130 raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Invalid JWT token")
133# Sentinel marker that designates a `request.state.user` value as having come from a
134# trusted, in-tree authentication code path (currently only `JWTRefreshMiddleware`).
135# `get_user()` only honours `request.state.user` when this sentinel is also present
136# at `request.state.user_authenticated_via`. This is defense-in-depth against an
137# accidental `request.state.user = ...` assignment in unrelated middleware (a typo,
138# a third-party plugin, a fixture leaked into production); it does NOT defend against
139# a malicious in-process plugin that imports the sentinel and sets it itself, since
140# plugins are trusted code in Airflow's security model — the goal is to keep an
141# accidental write from silently bypassing JWT validation.
142USER_INJECTED_BY_TRUSTED_MIDDLEWARE = object()
145async def get_user(
146 request: Request,
147 # Kept for the OpenAPI security spec so ``/docs`` still renders the OAuth2 password
148 # login form. It resolves to the same ``Authorization: Bearer`` header
149 # ``bearer_scheme`` reads, so the value is unused at runtime.
150 _oauth_token: str | None = Depends(oauth2_scheme),
151 bearer_credentials: HTTPAuthorizationCredentials | None = Depends(bearer_scheme),
152) -> BaseUser:
153 # An explicitly supplied credential always wins over the ambient session cookie.
154 if bearer_credentials and bearer_credentials.scheme.lower() == "bearer":
155 return await resolve_user_from_token(bearer_credentials.credentials)
157 # No explicit credential on this request, so the cookie is the caller's identity.
158 # A user might have been already built by a trusted in-tree middleware (currently
159 # only `JWTRefreshMiddleware`); if so, it is stored in `request.state.user` AND
160 # `request.state.user_authenticated_via` is set to the trust sentinel above.
161 # Honour the cached user only when both are present, so a stray `state.user`
162 # assignment from unrelated middleware can't bypass JWT validation.
163 user: BaseUser | None = getattr(request.state, "user", None)
164 trust_marker = getattr(request.state, "user_authenticated_via", None)
165 if user and trust_marker is USER_INJECTED_BY_TRUSTED_MIDDLEWARE: 165 ↛ 166line 165 didn't jump to line 166 because the condition on line 165 was never true
166 return user
167 return await resolve_user_from_token(request.cookies.get(COOKIE_NAME_JWT_TOKEN))
170def collect_request_tokens(
171 request: Request,
172 bearer_credentials: HTTPAuthorizationCredentials | None,
173) -> list[str]:
174 """
175 Return every distinct credential presented on this request, in precedence order.
177 Logout uses this rather than reproducing the single-credential choice
178 :func:`get_user` makes. Revoking only the precedence-selected credential would leave
179 any other one the caller presented still valid after they asked to be logged out,
180 and which credential "wins" is a question about *authentication* that should not
181 decide what a logout terminates.
182 """
183 candidates: list[str | None] = []
184 if bearer_credentials and bearer_credentials.scheme.lower() == "bearer":
185 candidates.append(bearer_credentials.credentials)
186 candidates.append(request.cookies.get(COOKIE_NAME_JWT_TOKEN))
188 tokens: list[str] = []
189 for candidate in candidates:
190 if candidate and candidate not in tokens:
191 tokens.append(candidate)
192 return tokens
195GetUserDep = Annotated[BaseUser, Depends(get_user)]
198def requires_access_dag(
199 method: ResourceMethod,
200 access_entity: DagAccessEntity | None = None,
201 param_dag_id: str | None = None,
202) -> Callable[[Request, BaseUser], None]:
203 def inner(
204 request: Request,
205 user: GetUserDep,
206 ) -> None:
207 # Required for the closure to capture the dag_id but still be able to mutate it.
208 # Prevent from using a nonlocal statement causing test failures.
209 dag_id = param_dag_id
210 if dag_id is None:
211 dag_id = request.path_params.get("dag_id") or request.query_params.get("dag_id")
212 dag_id = dag_id if dag_id != "~" else None
214 team_name = DagModel.get_team_name(dag_id) if dag_id else None
216 _requires_access(
217 is_authorized_callback=lambda: get_auth_manager().is_authorized_dag(
218 method=method,
219 access_entity=access_entity,
220 details=DagDetails(id=dag_id, team_name=team_name),
221 user=user,
222 )
223 )
225 return inner
228def requires_access_dag_from_file_token(
229 method: ResourceMethod,
230) -> Callable[[str, Request, BaseUser, Session], None]:
231 """
232 Authorize the caller against the DAGs referenced by a signed ``file_token``.
234 For ``file_token`` based endpoints (such as ``reparse``), the token is resolved to its referenced file, and authorization is performed against exactly the DAGs defined in that file, never against a request parameter.
235 """
237 def inner(
238 file_token: str,
239 request: Request,
240 user: GetUserDep,
241 session: SessionDep,
242 ) -> None:
243 try:
244 payload = URLSafeSerializer(request.app.state.secret_key).loads(file_token)
245 except BadSignature:
246 raise HTTPException(status.HTTP_404_NOT_FOUND, "File not found")
248 dag_ids = list(
249 session.scalars(
250 select(DagModel.dag_id).where(
251 DagModel.bundle_name == payload["bundle_name"],
252 DagModel.relative_fileloc == payload["relative_fileloc"],
253 )
254 )
255 )
256 if not dag_ids: 256 ↛ 257line 256 didn't jump to line 257 because the condition on line 256 was never true
257 raise HTTPException(status.HTTP_404_NOT_FOUND, "File not found")
259 dag_id_to_team = DagModel.get_dag_id_to_team_name_mapping(dag_ids, session=session)
260 requests: list[IsAuthorizedDagRequest] = [
261 {"method": method, "details": DagDetails(id=dag_id, team_name=dag_id_to_team.get(dag_id))}
262 for dag_id in dag_ids
263 ]
264 _requires_access(
265 is_authorized_callback=lambda: get_auth_manager().batch_is_authorized_dag(requests, user=user),
266 )
268 return inner
271class PermittedDagFilter(OrmClause[set[str]]):
272 """A parameter that filters the permitted dags for the user."""
274 def to_orm(self, select: Select) -> Select:
275 # self.value may be None (OrmClause holds Optional), ensure we pass an Iterable to in_
276 return select.where(DagModel.dag_id.in_(self.value or set()))
279class PermittedDagRunFilter(PermittedDagFilter):
280 """A parameter that filters the permitted dag runs for the user."""
282 def to_orm(self, select: Select) -> Select:
283 return select.where(DagRun.dag_id.in_(self.value or set()))
286class PermittedDagWarningFilter(PermittedDagFilter):
287 """A parameter that filters the permitted dag warnings for the user."""
289 def to_orm(self, select: Select) -> Select:
290 return select.where(DagWarning.dag_id.in_(self.value or set()))
293class PermittedEventLogFilter(PermittedDagFilter):
294 """A parameter that filters the permitted even logs for the user."""
296 def to_orm(self, select: Select) -> Select:
297 # Event Logs not related to Dags have dag_id as None and are always returned.
298 # return select.where(Log.dag_id.in_(self.value or set()) or Log.dag_id.is_(None))
299 return select.where(or_(Log.dag_id.in_(self.value or set()), Log.dag_id.is_(None)))
302class PermittedAssetEventFilter(PermittedDagFilter):
303 """A parameter that filters asset events to those produced by Dags the user may read."""
305 def to_orm(self, select: Select) -> Select:
306 # Asset events created through the API, or emitted by a watcher, have no source Dag.
307 # They carry no per-Dag key to authorize on, so they stay visible to any caller who
308 # may read assets; only events produced by a Dag's task are scoped to that Dag's
309 # readability. Filtering here rather than after the fact keeps unauthorized rows out
310 # of the count and pagination too, so their existence does not leak either.
311 return select.where(
312 or_(
313 AssetEvent.source_dag_id.in_(self.value or set()),
314 AssetEvent.source_dag_id.is_(None),
315 )
316 )
319class PermittedTIFilter(PermittedDagFilter):
320 """A parameter that filters the permitted task instances for the user."""
322 def to_orm(self, select: Select) -> Select:
323 return select.where(TI.dag_id.in_(self.value or set()))
326class PermittedXComFilter(PermittedDagFilter):
327 """A parameter that filters the permitted XComs for the user."""
329 def to_orm(self, select: Select) -> Select:
330 return select.where(XComModel.dag_id.in_(self.value or set()))
333class PermittedTagFilter(PermittedDagFilter):
334 """A parameter that filters the permitted dag tags for the user."""
336 def to_orm(self, select: Select) -> Select:
337 return select.where(DagTag.dag_id.in_(self.value or set()))
340class PermittedDagVersionFilter(PermittedDagFilter):
341 """A parameter that filters the permitted dag versions for the user."""
343 def to_orm(self, select: Select) -> Select:
344 return select.where(DagVersion.dag_id.in_(self.value or set()))
347class PermittedBackfillFilter(PermittedDagFilter):
348 """A parameter that filters the permitted backfills for the user."""
350 def to_orm(self, select: Select) -> Select:
351 return select.where(Backfill.dag_id.in_(self.value or set()))
354def permitted_dag_filter_factory(
355 method: ResourceMethod, filter_class=PermittedDagFilter
356) -> Callable[[BaseUser, BaseAuthManager], PermittedDagFilter]:
357 """
358 Create a callable for Depends in FastAPI that returns a filter of the permitted dags for the user.
360 :param method: whether filter readable or writable.
361 :return: The callable that can be used as Depends in FastAPI.
362 """
364 def depends_permitted_dags_filter(
365 user: GetUserDep,
366 auth_manager: AuthManagerDep,
367 ) -> PermittedDagFilter:
368 authorized_dags: set[str] = auth_manager.get_authorized_dag_ids(user=user, method=method)
369 return filter_class(authorized_dags)
371 return depends_permitted_dags_filter
374EditableDagsFilterDep = Annotated[PermittedDagFilter, Depends(permitted_dag_filter_factory("PUT"))]
375ReadableDagsFilterDep = Annotated[PermittedDagFilter, Depends(permitted_dag_filter_factory("GET"))]
376ReadableDagRunsFilterDep = Annotated[
377 PermittedDagRunFilter, Depends(permitted_dag_filter_factory("GET", PermittedDagRunFilter))
378]
379ReadableAssetEventsFilterDep = Annotated[
380 PermittedAssetEventFilter, Depends(permitted_dag_filter_factory("GET", PermittedAssetEventFilter))
381]
382ReadableDagWarningsFilterDep = Annotated[
383 PermittedDagWarningFilter, Depends(permitted_dag_filter_factory("GET", PermittedDagWarningFilter))
384]
385ReadableTIFilterDep = Annotated[
386 PermittedTIFilter, Depends(permitted_dag_filter_factory("GET", PermittedTIFilter))
387]
388ReadableEventLogsFilterDep = Annotated[
389 PermittedTIFilter, Depends(permitted_dag_filter_factory("GET", PermittedEventLogFilter))
390]
391ReadableXComFilterDep = Annotated[
392 PermittedXComFilter, Depends(permitted_dag_filter_factory("GET", PermittedXComFilter))
393]
395ReadableTagsFilterDep = Annotated[
396 PermittedTagFilter, Depends(permitted_dag_filter_factory("GET", PermittedTagFilter))
397]
398ReadableDagVersionsFilterDep = Annotated[
399 PermittedDagVersionFilter, Depends(permitted_dag_filter_factory("GET", PermittedDagVersionFilter))
400]
401ReadableBackfillsFilterDep = Annotated[
402 PermittedBackfillFilter, Depends(permitted_dag_filter_factory("GET", PermittedBackfillFilter))
403]
406# The type the backfill routes declare for the `backfill_id` path parameter. Shared with
407# `requires_access_backfill` so the authorization decision parses the id exactly as the handler
408# does; see the comment there for why any divergence is a cross-Dag authorization bypass.
409_BACKFILL_ID_ADAPTER: TypeAdapter[NonNegativeInt] = TypeAdapter(NonNegativeInt)
411# Shared with the backfill routes: for an id named in the path this dependency answers before
412# the handler does, so the two must not describe the same condition differently.
413BACKFILL_NOT_FOUND = "Backfill not found"
416def _authorize_backfill_in_path(method: ResourceMethod, dag_id: str | None, user: BaseUser) -> None:
417 """Authorize a backfill named by the request path against that backfill's Dag alone."""
418 # ``dag_id`` is None when the id matched no row. Answering 404 there while a backfill on a Dag
419 # the caller may not read answers 403 would tell them which backfill ids exist across Dags.
420 if dag_id is None:
421 raise HTTPException(status.HTTP_404_NOT_FOUND, BACKFILL_NOT_FOUND)
423 details = DagDetails(id=dag_id, team_name=DagModel.get_team_name(dag_id))
424 auth_manager = get_auth_manager()
425 if auth_manager.is_authorized_dag( 425 ↛ 431line 425 didn't jump to line 431 because the condition on line 425 was always true
426 method=method, access_entity=DagAccessEntity.RUN, details=details, user=user
427 ):
428 return
429 # A caller who may read the Dag can already list its backfills, so the id is no secret from
430 # them: hiding it would only cost them the reason their request was refused.
431 if method != "GET" and auth_manager.is_authorized_dag(
432 method="GET", access_entity=DagAccessEntity.RUN, details=details, user=user
433 ):
434 raise HTTPException(status.HTTP_403_FORBIDDEN, "Forbidden")
435 raise HTTPException(status.HTTP_404_NOT_FOUND, BACKFILL_NOT_FOUND)
438def requires_access_backfill(
439 method: ResourceMethod,
440) -> Callable[[Request, BaseUser, Session], Coroutine[Any, Any, None]]:
441 """Wrap ``requires_access_dag`` and extract the dag_id from the backfill_id."""
443 async def inner(
444 request: Request,
445 user: GetUserDep,
446 session: SessionDep,
447 ) -> None:
448 backfill_id_raw = request.path_params.get("backfill_id")
449 try:
450 # Must parse exactly as the handler does (e.g. pydantic's lax mode coerces "1.0" to 1
451 # where int() raises), or the two can authorize and act on different backfills.
452 backfill_id = (
453 _BACKFILL_ID_ADAPTER.validate_python(backfill_id_raw) if backfill_id_raw is not None else None
454 )
455 except ValidationError:
456 # Rejected by the endpoint's parser too, so the handler cannot run: FastAPI answers
457 # 422 before it is reached. Left as None, preserving that response.
458 backfill_id = None
460 if backfill_id is not None:
461 # The path names the backfill, so its row is the only authorization subject: what the
462 # caller supplies alongside must not decide a decision the path already scoped.
463 dag_id = session.scalar(select(Backfill.dag_id).where(Backfill.id == backfill_id))
464 _authorize_backfill_in_path(method, dag_id, user)
465 return
467 # Left: the routes naming their Dag in the body (create, dry run) or in the query string
468 # (list, read by ``requires_access_dag``), and ids the handler's own parser will reject.
469 dag_id = None
470 # Not a json body, ignore
471 with suppress(JSONDecodeError):
472 body = await request.json()
473 if isinstance(body, dict):
474 dag_id = body.get("dag_id")
475 if dag_id is not None and not isinstance(dag_id, str):
476 # Fail closed: reject non-string dag_id before authz decision.
477 raise HTTPException(
478 status_code=status.HTTP_400_BAD_REQUEST,
479 detail="'dag_id' must be a string",
480 )
482 requires_access_dag(method, DagAccessEntity.RUN, dag_id)(
483 request,
484 user,
485 )
487 return inner
490def requires_access_event_log(
491 method: ResourceMethod,
492) -> Callable[[Request, BaseUser, Session], Coroutine[Any, Any, None]]:
493 """Wrap ``requires_access_dag`` and extract the dag_id from the event_log_id."""
495 async def inner(
496 request: Request,
497 user: GetUserDep,
498 session: SessionDep,
499 ) -> None:
500 dag_id = None
502 event_log_id_raw = request.path_params.get("event_log_id")
503 if event_log_id_raw is not None:
504 try:
505 event_log_id = int(event_log_id_raw)
506 except ValueError:
507 raise HTTPException(
508 status_code=status.HTTP_400_BAD_REQUEST,
509 detail="'event_log_id' must be an integer",
510 )
511 dag_id = session.scalar(select(Log.dag_id).where(Log.id == event_log_id))
513 requires_access_dag(method, DagAccessEntity.AUDIT_LOG, dag_id)(
514 request,
515 user,
516 )
518 return inner
521class PermittedPoolFilter(OrmClause[set[str]]):
522 """A parameter that filters the permitted pools for the user."""
524 def to_orm(self, select: Select) -> Select:
525 return select.where(Pool.pool.in_(self.value or set()))
528def permitted_pool_filter_factory(
529 method: ResourceMethod,
530) -> Callable[[BaseUser, BaseAuthManager], PermittedPoolFilter]:
531 """
532 Create a callable for Depends in FastAPI that returns a filter of the permitted pools for the user.
534 :param method: whether filter readable or writable.
535 """
537 def depends_permitted_pools_filter(
538 user: GetUserDep,
539 auth_manager: AuthManagerDep,
540 ) -> PermittedPoolFilter:
541 authorized_pools: set[str] = auth_manager.get_authorized_pools(user=user, method=method)
542 return PermittedPoolFilter(authorized_pools)
544 return depends_permitted_pools_filter
547ReadablePoolsFilterDep = Annotated[PermittedPoolFilter, Depends(permitted_pool_filter_factory("GET"))]
550def requires_access_pool(method: ResourceMethod) -> Callable[[Request, BaseUser], Coroutine[Any, Any, None]]:
551 async def inner(
552 request: Request,
553 user: GetUserDep,
554 ) -> None:
555 pool_name = request.path_params.get("pool_name")
556 for team_name in await _collect_teams_to_check(method, request, pool_name, Pool.get_team_name):
558 def _callback(tn: str | None = team_name) -> bool:
559 return get_auth_manager().is_authorized_pool(
560 method=method, details=PoolDetails(name=pool_name, team_name=tn), user=user
561 )
563 _requires_access(is_authorized_callback=_callback)
565 return inner
568def requires_access_pool_bulk() -> Callable[[BulkBody[PoolBody], BaseUser], None]:
569 def inner(
570 request: BulkBody[PoolBody],
571 user: GetUserDep,
572 ) -> None:
573 multi_team = conf.getboolean("core", "multi_team")
574 # Build the list of pool names provided as part of the request that may correspond to
575 # an existing resource (UPDATE / DELETE, or CREATE+OVERWRITE which may turn into a PUT).
576 existing_pool_names = [
577 cast("str", entity) if action.action == BulkAction.DELETE else cast("PoolBody", entity).pool
578 for action in request.actions
579 for entity in action.entities
580 if _bulk_action_needs_existing_team_lookup(action)
581 ]
582 # For each pool, find its associated team (if it exists)
583 pool_name_to_team = Pool.get_name_to_team_name_mapping(existing_pool_names)
585 requests: list[IsAuthorizedPoolRequest] = []
586 for action in request.actions:
587 methods = _get_resource_methods_from_bulk_request(action)
588 for pool in action.entities:
589 pool_name = (
590 cast("str", pool) if action.action == BulkAction.DELETE else cast("PoolBody", pool).pool
591 )
592 for method in methods:
593 req: IsAuthorizedPoolRequest = {
594 "method": method,
595 "details": PoolDetails(
596 name=pool_name,
597 team_name=pool_name_to_team.get(pool_name),
598 ),
599 }
600 requests.append(req)
601 # Authorize the destination team_name when the entity body requests a team change.
602 if multi_team and _bulk_action_sets_team(action): 602 ↛ 603line 602 didn't jump to line 603 because the condition on line 602 was never true
603 dest_team = cast("PoolBody", pool).team_name
604 if dest_team is not None and dest_team != pool_name_to_team.get(pool_name):
605 for method in methods:
606 requests.append(
607 {
608 "method": method,
609 "details": PoolDetails(name=pool_name, team_name=dest_team),
610 }
611 )
613 _requires_access(
614 is_authorized_callback=lambda: get_auth_manager().batch_is_authorized_pool(
615 requests=requests,
616 user=user,
617 )
618 )
620 return inner
623class PermittedConnectionFilter(OrmClause[set[str]]):
624 """A parameter that filters the permitted connections for the user."""
626 def to_orm(self, select: Select) -> Select:
627 return select.where(Connection.conn_id.in_(self.value or set()))
630def permitted_connection_filter_factory(
631 method: ResourceMethod,
632) -> Callable[[BaseUser, BaseAuthManager], PermittedConnectionFilter]:
633 """
634 Create a callable for Depends in FastAPI that returns a filter of the permitted connections for the user.
636 :param method: whether filter readable or writable.
637 """
639 def depends_permitted_connections_filter(
640 user: GetUserDep,
641 auth_manager: AuthManagerDep,
642 ) -> PermittedConnectionFilter:
643 authorized_connections: set[str] = auth_manager.get_authorized_connections(user=user, method=method)
644 return PermittedConnectionFilter(authorized_connections)
646 return depends_permitted_connections_filter
649ReadableConnectionsFilterDep = Annotated[
650 PermittedConnectionFilter, Depends(permitted_connection_filter_factory("GET"))
651]
654def requires_access_connection(
655 method: ResourceMethod,
656) -> Callable[[Request, BaseUser], Coroutine[Any, Any, None]]:
657 async def inner(
658 request: Request,
659 user: GetUserDep,
660 ) -> None:
661 connection_id = request.path_params.get("connection_id")
662 for team_name in await _collect_teams_to_check(
663 method, request, connection_id, Connection.get_team_name
664 ):
666 def _callback(tn: str | None = team_name) -> bool:
667 return get_auth_manager().is_authorized_connection(
668 method=method,
669 details=ConnectionDetails(conn_id=connection_id, team_name=tn),
670 user=user,
671 )
673 _requires_access(is_authorized_callback=_callback)
675 return inner
678def requires_access_connection_bulk() -> Callable[[BulkBody[ConnectionBody], BaseUser], None]:
679 def inner(
680 request: BulkBody[ConnectionBody],
681 user: GetUserDep,
682 ) -> None:
683 multi_team = conf.getboolean("core", "multi_team")
684 # Build the list of ``conn_id`` provided as part of the request that may correspond to
685 # an existing resource (UPDATE / DELETE, or CREATE+OVERWRITE which may turn into a PUT).
686 existing_connection_ids = [
687 cast("str", entity)
688 if action.action == BulkAction.DELETE
689 else cast("ConnectionBody", entity).connection_id
690 for action in request.actions
691 for entity in action.entities
692 if _bulk_action_needs_existing_team_lookup(action)
693 ]
694 # For each connection, find its associated team (if it exists)
695 conn_id_to_team = Connection.get_conn_id_to_team_name_mapping(existing_connection_ids)
697 requests: list[IsAuthorizedConnectionRequest] = []
698 for action in request.actions:
699 methods = _get_resource_methods_from_bulk_request(action)
700 for connection in action.entities:
701 connection_id = (
702 cast("str", connection)
703 if action.action == BulkAction.DELETE
704 else cast("ConnectionBody", connection).connection_id
705 )
706 for method in methods:
707 req: IsAuthorizedConnectionRequest = {
708 "method": method,
709 "details": ConnectionDetails(
710 conn_id=connection_id,
711 team_name=conn_id_to_team.get(connection_id),
712 ),
713 }
714 requests.append(req)
715 # Authorize the destination team_name when the entity body requests a team change.
716 if multi_team and _bulk_action_sets_team(action): 716 ↛ 717line 716 didn't jump to line 717 because the condition on line 716 was never true
717 dest_team = cast("ConnectionBody", connection).team_name
718 if dest_team is not None and dest_team != conn_id_to_team.get(connection_id):
719 for method in methods:
720 requests.append(
721 {
722 "method": method,
723 "details": ConnectionDetails(conn_id=connection_id, team_name=dest_team),
724 }
725 )
727 _requires_access(
728 is_authorized_callback=lambda: get_auth_manager().batch_is_authorized_connection(
729 requests=requests,
730 user=user,
731 )
732 )
734 return inner
737def requires_access_configuration(method: ResourceMethod) -> Callable[[Request, BaseUser], None]:
738 def inner(
739 request: Request,
740 user: GetUserDep,
741 ) -> None:
742 section: str | None = request.query_params.get("section") or request.path_params.get("section")
744 _requires_access(
745 is_authorized_callback=lambda: get_auth_manager().is_authorized_configuration(
746 method=method,
747 details=ConfigurationDetails(section=section),
748 user=user,
749 )
750 )
752 return inner
755class PermittedTeamFilter(OrmClause[set[str]]):
756 """A parameter that filters the permitted teams for the user."""
758 def to_orm(self, select: Select) -> Select:
759 return select.where(Team.name.in_(self.value or set()))
762def permitted_team_filter_factory() -> Callable[[BaseUser, BaseAuthManager], PermittedTeamFilter]:
763 """Create a callable for Depends in FastAPI that returns a filter of the permitted teams for the user."""
765 def depends_permitted_teams_filter(
766 user: GetUserDep,
767 auth_manager: AuthManagerDep,
768 ) -> PermittedTeamFilter:
769 authorized_teams: set[str] = auth_manager.get_authorized_teams(user=user, method="GET")
770 return PermittedTeamFilter(authorized_teams)
772 return depends_permitted_teams_filter
775ReadableTeamsFilterDep = Annotated[PermittedTeamFilter, Depends(permitted_team_filter_factory())]
778class PermittedVariableFilter(OrmClause[set[str]]):
779 """A parameter that filters the permitted variables for the user."""
781 def to_orm(self, select: Select) -> Select:
782 return select.where(Variable.key.in_(self.value or set()))
785def permitted_variable_filter_factory(
786 method: ResourceMethod,
787) -> Callable[[BaseUser, BaseAuthManager], PermittedVariableFilter]:
788 """
789 Create a callable for Depends in FastAPI that returns a filter of the permitted variables for the user.
791 :param method: whether filter readable or writable.
792 """
794 def depends_permitted_variables_filter(
795 user: GetUserDep,
796 auth_manager: AuthManagerDep,
797 ) -> PermittedVariableFilter:
798 authorized_variables: set[str] = auth_manager.get_authorized_variables(user=user, method=method)
799 return PermittedVariableFilter(authorized_variables)
801 return depends_permitted_variables_filter
804ReadableVariablesFilterDep = Annotated[
805 PermittedVariableFilter, Depends(permitted_variable_filter_factory("GET"))
806]
809def requires_access_variable(
810 method: ResourceMethod,
811) -> Callable[[Request, BaseUser], Coroutine[Any, Any, None]]:
812 async def inner(
813 request: Request,
814 user: GetUserDep,
815 ) -> None:
816 variable_key: str | None = request.path_params.get("variable_key")
817 for team_name in await _collect_teams_to_check(method, request, variable_key, Variable.get_team_name):
819 def _callback(tn: str | None = team_name) -> bool:
820 return get_auth_manager().is_authorized_variable(
821 method=method, details=VariableDetails(key=variable_key, team_name=tn), user=user
822 )
824 _requires_access(is_authorized_callback=_callback)
826 return inner
829def requires_access_variable_bulk() -> Callable[[BulkBody[VariableBody], BaseUser], None]:
830 def inner(
831 request: BulkBody[VariableBody],
832 user: GetUserDep,
833 ) -> None:
834 multi_team = conf.getboolean("core", "multi_team")
835 # Build the list of variable keys provided as part of the request that may correspond to
836 # an existing resource (UPDATE / DELETE, or CREATE+OVERWRITE which may turn into a PUT).
837 existing_variable_keys = [
838 cast("str", entity) if action.action == BulkAction.DELETE else cast("VariableBody", entity).key
839 for action in request.actions
840 for entity in action.entities
841 if _bulk_action_needs_existing_team_lookup(action)
842 ]
843 # For each variable, find its associated team (if it exists)
844 var_key_to_team = Variable.get_key_to_team_name_mapping(existing_variable_keys)
846 requests: list[IsAuthorizedVariableRequest] = []
847 for action in request.actions:
848 methods = _get_resource_methods_from_bulk_request(action)
849 for variable in action.entities:
850 variable_key = (
851 cast("str", variable)
852 if action.action == BulkAction.DELETE
853 else cast("VariableBody", variable).key
854 )
855 for method in methods:
856 req: IsAuthorizedVariableRequest = {
857 "method": method,
858 "details": VariableDetails(
859 key=variable_key,
860 team_name=var_key_to_team.get(variable_key),
861 ),
862 }
863 requests.append(req)
864 # Authorize the destination team_name when the entity body requests a team change.
865 if multi_team and _bulk_action_sets_team(action): 865 ↛ 866line 865 didn't jump to line 866 because the condition on line 865 was never true
866 dest_team = cast("VariableBody", variable).team_name
867 if dest_team is not None and dest_team != var_key_to_team.get(variable_key):
868 for method in methods:
869 requests.append(
870 {
871 "method": method,
872 "details": VariableDetails(key=variable_key, team_name=dest_team),
873 }
874 )
876 _requires_access(
877 is_authorized_callback=lambda: get_auth_manager().batch_is_authorized_variable(
878 requests=requests,
879 user=user,
880 )
881 )
883 return inner
886def _build_dag_run_access_requests(
887 entity_methods: list[tuple[str, ResourceMethod]],
888) -> list[IsAuthorizedDagRequest]:
889 """
890 Build per-entity DagRun authorization requests for a batched access check.
892 ``entity_methods`` is a list of ``(dag_id, method)`` pairs with unresolvable
893 entries (no dag_id or the ``~`` wildcard) already filtered out by the caller.
894 Teams for all Dags are resolved in a single batched query and shared across each
895 Dag's requests.
896 """
897 if not entity_methods:
898 return []
899 resolved_dag_ids = list({dag_id for dag_id, _ in entity_methods})
900 dag_id_to_team = DagModel.get_dag_id_to_team_name_mapping(resolved_dag_ids)
901 return [
902 {
903 "method": method,
904 "access_entity": DagAccessEntity.RUN,
905 "details": DagDetails(id=dag_id, team_name=dag_id_to_team.get(dag_id)),
906 }
907 for dag_id, method in entity_methods
908 ]
911def requires_access_dag_run_bulk() -> Callable[[BulkBody[BulkDAGRunBody], BaseUser, str], None]:
912 def inner(
913 request: BulkBody[BulkDAGRunBody],
914 user: GetUserDep,
915 dag_id: str,
916 ) -> None:
917 entity_methods: list[tuple[str, ResourceMethod]] = []
918 for action in request.actions:
919 methods = _get_resource_methods_from_bulk_request(action)
920 for entity in action.entities:
921 if isinstance(entity, str):
922 entity_dag_id: str | None = dag_id
923 else:
924 entity_dag_id = entity.dag_id or dag_id
925 # Entities that can't be resolved are surfaced as 400 in the service's BulkResponse.
926 if not entity_dag_id or entity_dag_id == "~": 926 ↛ 927line 926 didn't jump to line 927 because the condition on line 926 was never true
927 continue
928 for method in methods:
929 entity_methods.append((entity_dag_id, method))
931 requests = _build_dag_run_access_requests(entity_methods)
932 _requires_access(
933 is_authorized_callback=lambda: get_auth_manager().batch_is_authorized_dag(
934 requests=requests,
935 user=user,
936 )
937 )
939 return inner
942def requires_access_dag_run_clear_bulk() -> Callable[[BulkDAGRunClearBody, BaseUser, str], None]:
943 def inner(
944 body: BulkDAGRunClearBody,
945 user: GetUserDep,
946 dag_id: str,
947 ) -> None:
948 entity_methods: list[tuple[str, ResourceMethod]] = []
949 for run in body.dag_runs:
950 entity_dag_id = run.dag_id or dag_id
951 if not entity_dag_id or entity_dag_id == "~": 951 ↛ 952line 951 didn't jump to line 952 because the condition on line 951 was never true
952 continue
953 entity_methods.append((entity_dag_id, "PUT"))
955 if not body.dag_runs and body.has_partition_selectors:
956 if dag_id and dag_id != "~": 956 ↛ 959line 956 didn't jump to line 959 because the condition on line 956 was always true
957 entity_methods.append((dag_id, "PUT"))
959 requests = _build_dag_run_access_requests(entity_methods)
960 _requires_access(
961 is_authorized_callback=lambda: get_auth_manager().batch_is_authorized_dag(
962 requests=requests,
963 user=user,
964 )
965 )
967 return inner
970def requires_access_asset(method: ResourceMethod) -> Callable[[Request, BaseUser], None]:
971 def inner(
972 request: Request,
973 user: GetUserDep,
974 ) -> None:
975 asset_id = request.path_params.get("asset_id")
977 _requires_access(
978 is_authorized_callback=lambda: get_auth_manager().is_authorized_asset(
979 method=method, details=AssetDetails(id=asset_id), user=user
980 ),
981 )
983 return inner
986def requires_access_view(access_view: AccessView) -> Callable[[Request, BaseUser], None]:
987 def inner(
988 request: Request,
989 user: GetUserDep,
990 ) -> None:
991 _requires_access(
992 is_authorized_callback=lambda: get_auth_manager().is_authorized_view(
993 access_view=access_view, user=user
994 ),
995 )
997 return inner
1000def requires_access_asset_alias(method: ResourceMethod) -> Callable[[Request, BaseUser], None]:
1001 def inner(
1002 request: Request,
1003 user: GetUserDep,
1004 ) -> None:
1005 asset_alias_id: str | None = request.path_params.get("asset_alias_id")
1007 _requires_access(
1008 is_authorized_callback=lambda: get_auth_manager().is_authorized_asset_alias(
1009 method=method, details=AssetAliasDetails(id=asset_alias_id), user=user
1010 ),
1011 )
1013 return inner
1016def requires_authenticated() -> Callable:
1017 """Just ensure the user is authenticated - no need to check any specific permissions."""
1019 def inner(
1020 request: Request,
1021 user: GetUserDep,
1022 ) -> None:
1023 pass
1025 return inner
1028async def _collect_teams_to_check(
1029 method: ResourceMethod,
1030 request: Request,
1031 resource_id: str | None,
1032 get_existing_team: Callable[[str], str | None],
1033) -> set[str | None]:
1034 """Collect validated team names from existing resource (DB) and/or request body."""
1035 if not conf.getboolean("core", "multi_team"): 1035 ↛ 1037line 1035 didn't jump to line 1037 because the condition on line 1035 was always true
1036 return {None}
1037 teams: set[str | None] = set()
1038 if method != "POST":
1039 teams.add(get_existing_team(resource_id) if resource_id else None)
1040 if method in ("POST", "PUT"):
1041 try:
1042 body = await request.json()
1043 except JSONDecodeError:
1044 # Fail closed: reject unparsable bodies before any authz decision.
1045 raise HTTPException(
1046 status_code=status.HTTP_400_BAD_REQUEST,
1047 detail="Request body is not valid JSON",
1048 )
1049 raw = body.get("team_name") if isinstance(body, dict) else None
1050 if raw is not None and not isinstance(raw, str):
1051 # Fail closed: reject non-string team_name before authz / DB lookup.
1052 raise HTTPException(
1053 status_code=status.HTTP_400_BAD_REQUEST,
1054 detail="'team_name' must be a string",
1055 )
1056 if raw and not Team.get_name_if_exists(raw):
1057 raise HTTPException(
1058 status_code=status.HTTP_400_BAD_REQUEST,
1059 detail=f"Team {raw!r} does not exist",
1060 )
1061 teams.add(raw)
1062 return teams
1065def _requires_access(
1066 *,
1067 is_authorized_callback: Callable[[], bool],
1068) -> None:
1069 if not is_authorized_callback(): 1069 ↛ 1070line 1069 didn't jump to line 1070 because the condition on line 1069 was never true
1070 raise HTTPException(status.HTTP_403_FORBIDDEN, "Forbidden")
1073def is_safe_url(target_url: str, request: Request | None = None) -> bool:
1074 """
1075 Check that the URL is safe.
1077 Needs to belong to the same domain as base_url, use HTTP or HTTPS (no JavaScript/data schemes),
1078 is a valid normalized path.
1079 """
1080 parsed_bases: tuple[tuple[str, ParseResult], ...] = ()
1082 # Check if the target URL matches either the configured base URL, or the URL used to make the request
1083 if request is not None: 1083 ↛ 1086line 1083 didn't jump to line 1086 because the condition on line 1083 was always true
1084 url = str(request.base_url)
1085 parsed_bases += ((url, urlparse(url)),)
1086 if base_url := conf.get("api", "base_url", fallback=None): 1086 ↛ 1087line 1086 didn't jump to line 1087 because the condition on line 1086 was never true
1087 parsed_bases += ((base_url, urlparse(base_url)),)
1089 if not parsed_bases: 1089 ↛ 1091line 1089 didn't jump to line 1091 because the condition on line 1089 was never true
1090 # Can't enforce any security check.
1091 return True
1093 # According to WHATWG for http/https /// is interpreted as // whereas urllib doesnt
1094 # this leads to an inconsistency where python returns a target url with /// as a valid url
1095 # The same thing also happens with \ where under WHATWG \ are translated to /, including
1096 # after a scheme, so "https:\\host" is an authority for a browser but a path for urllib.
1097 target_url = unquote(target_url).strip().replace("\\", "/")
1098 if target_url.startswith("//"): 1098 ↛ 1099line 1098 didn't jump to line 1099 because the condition on line 1098 was never true
1099 return False
1100 for base_url, parsed_base in parsed_bases: 1100 ↛ 1117line 1100 didn't jump to line 1117 because the loop on line 1100 didn't complete
1101 parsed_target = urlparse(urljoin(base_url, target_url)) # Resolves relative URLs
1103 base_path = parsed_base.path or "/"
1104 target_path = parsed_target.path or "/"
1106 # Normalize as POSIX paths (URL paths) and ensure target is under base.
1107 norm_base = posixpath.normpath(base_path)
1108 norm_target = posixpath.normpath(target_path)
1110 if norm_base != "/": 1110 ↛ 1111line 1110 didn't jump to line 1111 because the condition on line 1110 was never true
1111 norm_base_with_slash = norm_base if norm_base.endswith("/") else norm_base + "/"
1112 if norm_target != norm_base and not norm_target.startswith(norm_base_with_slash):
1113 continue
1115 if parsed_target.scheme in {"http", "https"} and parsed_target.netloc == parsed_base.netloc: 1115 ↛ 1100line 1115 didn't jump to line 1100 because the condition on line 1115 was always true
1116 return True
1117 return False
1120def _get_resource_methods_from_bulk_request(
1121 action: BulkCreateAction | BulkUpdateAction | BulkDeleteAction,
1122) -> list[ResourceMethod]:
1123 resource_methods: list[ResourceMethod] = [MAP_BULK_ACTION_TO_AUTH_METHOD[action.action]]
1124 # If ``action_on_existence`` == ``overwrite``, we need to check the user has ``PUT`` access as well.
1125 # With ``action_on_existence`` == ``overwrite``, a create request is actually an update request if the
1126 # resource already exists, hence adding this check.
1127 if action.action == BulkAction.CREATE and action.action_on_existence == BulkActionOnExistence.OVERWRITE:
1128 resource_methods.append("PUT")
1129 return resource_methods
1132def _bulk_action_needs_existing_team_lookup(
1133 action: BulkCreateAction | BulkUpdateAction | BulkDeleteAction,
1134) -> bool:
1135 # UPDATE / DELETE always operate on existing resources, so we need the existing team for authz.
1136 # CREATE with action_on_existence=OVERWRITE may turn into a PUT against an existing resource that
1137 # belongs to a team; if we omit it from the lookup, the PUT authz check runs with team_name=None
1138 # and bypasses the per-team membership check that the single-item PUT endpoint enforces.
1139 if action.action != BulkAction.CREATE:
1140 return True
1141 return action.action_on_existence == BulkActionOnExistence.OVERWRITE
1144def _bulk_action_sets_team(
1145 action: BulkCreateAction | BulkUpdateAction | BulkDeleteAction,
1146) -> bool:
1147 """Return True if this action can write a team_name (UPDATE, or CREATE that carries a body)."""
1148 return action.action in (BulkAction.UPDATE, BulkAction.CREATE)