Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions airflow/api_connexion/endpoints/task_instance_endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
task_instance_reference_schema,
task_instance_schema,
)
from airflow.api_connexion.security import get_readable_dags
from airflow.models import SlaMiss
from airflow.models.dagrun import DagRun as DR
from airflow.models.operator import needs_expansion
Expand Down Expand Up @@ -342,6 +343,8 @@ def get_task_instances(

if dag_id != "~":
base_query = base_query.where(TI.dag_id == dag_id)
else:
base_query = base_query.where(TI.dag_id.in_(get_readable_dags()))
if dag_run_id != "~":
base_query = base_query.where(TI.run_id == dag_run_id)
base_query = _apply_range_filter(
Expand Down
10 changes: 9 additions & 1 deletion airflow/api_connexion/security.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from functools import wraps
from typing import Callable, Sequence, TypeVar, cast

from flask import Response
from flask import Response, g

from airflow.api_connexion.exceptions import PermissionDenied, Unauthenticated
from airflow.utils.airflow_flask_app import get_airflow_app
Expand Down Expand Up @@ -55,3 +55,11 @@ def decorated(*args, **kwargs):
return cast(T, decorated)

return requires_access_decorator


def get_readable_dags() -> list[str]:
return get_airflow_app().appbuilder.sm.get_accessible_dag_ids(g.user)


def can_read_dag(dag_id: str) -> bool:
return get_airflow_app().appbuilder.sm.can_read_dag(dag_id, g.user)
46 changes: 46 additions & 0 deletions tests/api_connexion/endpoints/test_task_instance_endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -658,6 +658,52 @@ def test_should_respond_200(self, task_instances, update_extras, url, expected_t
assert response.json["total_entries"] == expected_ti
assert len(response.json["task_instances"]) == expected_ti

@pytest.mark.parametrize(
"task_instances, user, expected_ti",
[
pytest.param(
{
"example_python_operator": 2,
"example_skip_dag": 1,
},
"test_read_only_one_dag",
2,
),
pytest.param(
{
"example_python_operator": 1,
"example_skip_dag": 2,
},
"test_read_only_one_dag",
1,
),
pytest.param(
{
"example_python_operator": 1,
"example_skip_dag": 2,
},
"test",
3,
),
],
)
def test_return_TI_only_from_readable_dags(self, task_instances, user, expected_ti, session):
for dag_id in task_instances:
self.create_task_instances(
session,
task_instances=[
{"execution_date": DEFAULT_DATETIME_1 + dt.timedelta(days=i)}
for i in range(task_instances[dag_id])
],
dag_id=dag_id,
)
response = self.client.get(
"/api/v1/dags/~/dagRuns/~/taskInstances", environ_overrides={"REMOTE_USER": user}
)
assert response.status_code == 200
assert response.json["total_entries"] == expected_ti
assert len(response.json["task_instances"]) == expected_ti

def test_should_respond_200_for_dag_id_filter(self, session):
self.create_task_instances(session)
self.create_task_instances(session, dag_id="example_skip_dag")
Expand Down