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
Original file line number Diff line number Diff line change
Expand Up @@ -181,7 +181,7 @@ paths:
application/json:
schema:
$ref: '#/components/schemas/HTTPValidationError'
/ui/pending_partitioned_dag_run/{dag_id}/{partition_key}:
/ui/pending_partitioned_dag_run/{dag_id}:
get:
tags:
- PartitionedDagRun
Expand All @@ -199,7 +199,7 @@ paths:
type: string
title: Dag Id
- name: partition_key
in: path
in: query
required: true
schema:
type: string
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -342,7 +342,7 @@ def get_partitioned_dag_runs(


@partitioned_dag_runs_router.get(
"/pending_partitioned_dag_run/{dag_id}/{partition_key}",
"/pending_partitioned_dag_run/{dag_id}",
dependencies=[Depends(requires_access_asset(method="GET")), Depends(requires_access_dag(method="GET"))],
)
def get_pending_partitioned_dag_run(
Expand All @@ -351,6 +351,9 @@ def get_pending_partitioned_dag_run(
session: SessionDep,
) -> PartitionedDagRunDetailResponse:
"""Return full details for pending PartitionedDagRun."""
# partition_key is a query param, not a path segment: it is a free-form key
# (up to 250 chars) that may itself contain "/", which would otherwise be
# ambiguous (or mis-routed) as a path segment.
partitioned_dag_run = session.execute(
select(
AssetPartitionDagRun.id,
Expand All @@ -366,7 +369,11 @@ def get_pending_partitioned_dag_run(
AssetPartitionDagRun.partition_key == partition_key,
AssetPartitionDagRun.created_dag_run_id.is_(None),
)
).one_or_none()
# Duplicate pending rows for the same (dag_id, partition_key) can exist
# after a crash; mirror _get_or_create_apdr and work on the latest one.
.order_by(AssetPartitionDagRun.id.desc())
.limit(1)
).first()

if partitioned_dag_run is None:
raise HTTPException(
Expand Down Expand Up @@ -440,9 +447,15 @@ def get_pending_partitioned_dag_run(
required_keys = []
received_count = 0
required_count = 1
else:
elif is_rollup:
received_count = len(received_keys)
required_count = len(required_keys)
else:
# Match the list route's _compute_received_count: a non-rollup asset is
# satisfied by any single received event, so credit caps at 1 even if
# several distinct upstream keys mapped onto this one target key.
required_count = len(required_keys)
received_count = 1 if received_keys else 0
assets.append(
PartitionedDagRunAssetResponse(
asset_id=asset_row.id,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4696,9 +4696,11 @@ export class PartitionedDagRunService {
public static getPendingPartitionedDagRun(data: GetPendingPartitionedDagRunData): CancelablePromise<GetPendingPartitionedDagRunResponse> {
return __request(OpenAPI, {
method: 'GET',
url: '/ui/pending_partitioned_dag_run/{dag_id}/{partition_key}',
url: '/ui/pending_partitioned_dag_run/{dag_id}',
path: {
dag_id: data.dagId,
dag_id: data.dagId
},
query: {
partition_key: data.partitionKey
},
errors: {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8238,7 +8238,7 @@ export type $OpenApiTs = {
};
};
};
'/ui/pending_partitioned_dag_run/{dag_id}/{partition_key}': {
'/ui/pending_partitioned_dag_run/{dag_id}': {
get: {
req: GetPendingPartitionedDagRunData;
res: {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -544,11 +544,13 @@ def test_list_route_total_required_includes_inactive_asset(self, test_client, da

class TestGetPendingPartitionedDagRun:
def test_should_response_401(self, unauthenticated_test_client):
response = unauthenticated_test_client.get("/pending_partitioned_dag_run/any_dag/any_key")
response = unauthenticated_test_client.get(
"/pending_partitioned_dag_run/any_dag?partition_key=any_key"
)
assert response.status_code == 401

def test_should_response_403(self, unauthorized_test_client):
response = unauthorized_test_client.get("/pending_partitioned_dag_run/any_dag/any_key")
response = unauthorized_test_client.get("/pending_partitioned_dag_run/any_dag?partition_key=any_key")
assert response.status_code == 403

@pytest.mark.parametrize(
Expand Down Expand Up @@ -583,7 +585,15 @@ def test_should_response_404(self, test_client, dag_maker, session, dag_id, part
)
session.commit()

resp = test_client.get(f"/pending_partitioned_dag_run/{dag_id}/{partition_key}")
resp = test_client.get(f"/pending_partitioned_dag_run/{dag_id}?partition_key={partition_key}")
assert resp.status_code == 404

def test_missing_partition_key_query_param_returns_422(self, test_client):
resp = test_client.get("/pending_partitioned_dag_run/any_dag")
assert resp.status_code == 422

def test_empty_partition_key_query_param_returns_404(self, test_client):
resp = test_client.get("/pending_partitioned_dag_run/any_dag?partition_key=")
assert resp.status_code == 404

@pytest.mark.parametrize(
Expand Down Expand Up @@ -644,7 +654,7 @@ def test_should_response_200(self, test_client, dag_maker, session, num_assets,
)
session.commit()

resp = test_client.get("/pending_partitioned_dag_run/detail_dag/2024-07-01")
resp = test_client.get("/pending_partitioned_dag_run/detail_dag?partition_key=2024-07-01")
assert resp.status_code == 200
body = resp.json()
assert body["dag_id"] == "detail_dag"
Expand Down Expand Up @@ -677,7 +687,7 @@ def test_is_rollup_false_for_non_rollup_asset(self, test_client, dag_maker, sess
session.add(AssetPartitionDagRun(target_dag_id="nr_detail_dag", partition_key="2024-07-01"))
session.commit()

resp = test_client.get("/pending_partitioned_dag_run/nr_detail_dag/2024-07-01")
resp = test_client.get("/pending_partitioned_dag_run/nr_detail_dag?partition_key=2024-07-01")
assert resp.status_code == 200
assets = resp.json()["assets"]
assert len(assets) == 2
Expand Down Expand Up @@ -713,7 +723,7 @@ def test_is_rollup_true_for_default_rollup_mapper(self, test_client, dag_maker,
session.add(AssetPartitionDagRun(target_dag_id="rollup_default_dag", partition_key="2024-06-03"))
session.commit()

resp = test_client.get("/pending_partitioned_dag_run/rollup_default_dag/2024-06-03")
resp = test_client.get("/pending_partitioned_dag_run/rollup_default_dag?partition_key=2024-06-03")
assert resp.status_code == 200
body = resp.json()
assert body["total_required"] == 14
Expand Down Expand Up @@ -761,7 +771,7 @@ def test_is_rollup_true_for_rollup_asset(self, test_client, dag_maker, session):
)
session.commit()

resp = test_client.get("/pending_partitioned_dag_run/rollup_detail_dag/2024-06-03")
resp = test_client.get("/pending_partitioned_dag_run/rollup_detail_dag?partition_key=2024-06-03")
assert resp.status_code == 200
body = resp.json()
assert body["total_required"] == 7
Expand Down Expand Up @@ -832,7 +842,9 @@ def test_rollup_mapper_failure_treats_asset_as_not_satisfied(self, test_client,
),
mock.patch("airflow.api_fastapi.core_api.routes.ui.partitioned_dag_runs.log") as mock_log,
):
resp = test_client.get("/pending_partitioned_dag_run/rollup_warn_detail_dag/2024-06-03")
resp = test_client.get(
"/pending_partitioned_dag_run/rollup_warn_detail_dag?partition_key=2024-06-03"
)

assert resp.status_code == 200
a = resp.json()["assets"][0]
Expand Down Expand Up @@ -873,7 +885,7 @@ def test_partitioned_dag_runs_asset_inactive_true_when_deactivated(self, test_cl
session.delete(asset_active_row)
session.commit()

resp = test_client.get("/pending_partitioned_dag_run/inactive_detail_dag/2024-07-01")
resp = test_client.get("/pending_partitioned_dag_run/inactive_detail_dag?partition_key=2024-07-01")
assert resp.status_code == 200
assets = resp.json()["assets"]
assert len(assets) == 1
Expand All @@ -896,8 +908,137 @@ def test_partitioned_dag_runs_asset_inactive_false_for_active_asset(
session.add(AssetPartitionDagRun(target_dag_id="active_detail_dag", partition_key="2024-07-01"))
session.commit()

resp = test_client.get("/pending_partitioned_dag_run/active_detail_dag/2024-07-01")
resp = test_client.get("/pending_partitioned_dag_run/active_detail_dag?partition_key=2024-07-01")
assert resp.status_code == 200
assets = resp.json()["assets"]
assert len(assets) == 1
assert assets[0]["asset_inactive"] is False

def test_partition_key_containing_slash_round_trips(self, test_client, dag_maker, session):
"""
A partition key containing ``/`` must survive as a query parameter.

As a path segment, a URL-encoded ``/`` (``%2F``) is decoded before Starlette
routing sees it, so any key containing a literal ``/`` would 404. Sending it
as a query parameter instead avoids that ambiguity entirely.
"""
asset = Asset(uri="s3://bucket/slash_key", name="slash_key")
with dag_maker(
dag_id="slash_key_dag",
schedule=PartitionedAssetTimetable(assets=asset),
serialized=True,
):
EmptyOperator(task_id="t")
dag_maker.create_dagrun()
dag_maker.sync_dagbag_to_db()

session.add(AssetPartitionDagRun(target_dag_id="slash_key_dag", partition_key="region/us"))
session.commit()

resp = test_client.get(
"/pending_partitioned_dag_run/slash_key_dag", params={"partition_key": "region/us"}
)
assert resp.status_code == 200
body = resp.json()
assert body["dag_id"] == "slash_key_dag"
assert body["partition_key"] == "region/us"

def test_duplicate_pending_apdr_rows_return_latest(self, test_client, dag_maker, session):
"""
Duplicate pending APDR rows for the same (dag_id, partition_key) must not 500.

The model docstring for ``AssetPartitionDagRun`` says callers should always
work on the latest row when duplicates exist; the route must do the same
instead of raising ``MultipleResultsFound`` from ``.one_or_none()``.
"""
asset = Asset(uri="s3://bucket/dup_apdr", name="dup_apdr")
with dag_maker(
dag_id="dup_apdr_dag",
schedule=PartitionedAssetTimetable(assets=asset),
serialized=True,
):
EmptyOperator(task_id="t")
dag_maker.create_dagrun()
dag_maker.sync_dagbag_to_db()

asset = session.scalar(select(AssetModel).where(AssetModel.uri == "s3://bucket/dup_apdr"))

# Older duplicate row: no received events.
stale_pdr = AssetPartitionDagRun(target_dag_id="dup_apdr_dag", partition_key="2024-08-01")
session.add(stale_pdr)
session.flush()

# Newer duplicate row (higher id): has a received event.
latest_pdr = AssetPartitionDagRun(target_dag_id="dup_apdr_dag", partition_key="2024-08-01")
session.add(latest_pdr)
session.flush()
event = AssetEvent(asset_id=asset.id, timestamp=pendulum.now())
session.add(event)
session.flush()
session.add(
PartitionedAssetKeyLog(
asset_id=asset.id,
asset_event_id=event.id,
asset_partition_dag_run_id=latest_pdr.id,
source_partition_key="2024-08-01",
target_dag_id="dup_apdr_dag",
target_partition_key="2024-08-01",
)
)
session.commit()

resp = test_client.get("/pending_partitioned_dag_run/dup_apdr_dag?partition_key=2024-08-01")
assert resp.status_code == 200
body = resp.json()
assert body["id"] == latest_pdr.id
assert body["total_received"] == 1

def test_non_rollup_many_to_one_received_capped_at_one(self, test_client, dag_maker, session):
"""
Detail route must cap non-rollup received credit at 1, matching the list route.

Multiple distinct upstream ``source_partition_key`` values logged against the
same non-rollup asset (a many-to-one mapper) must not inflate ``received_count``
past ``required_count`` — otherwise the detail view shows e.g. "2/1" while the
list view shows "1/1" for the same APDR.
"""
asset = Asset(uri="s3://bucket/many_to_one", name="many_to_one")
with dag_maker(
dag_id="many_to_one_dag",
schedule=PartitionedAssetTimetable(assets=asset),
serialized=True,
):
EmptyOperator(task_id="t")
dag_maker.create_dagrun()
dag_maker.sync_dagbag_to_db()

asset = session.scalar(select(AssetModel).where(AssetModel.uri == "s3://bucket/many_to_one"))
pdr = AssetPartitionDagRun(target_dag_id="many_to_one_dag", partition_key="2024-08-01")
session.add(pdr)
session.flush()

for source_key in ("2024-08-01", "different-key"):
event = AssetEvent(asset_id=asset.id, timestamp=pendulum.now())
session.add(event)
session.flush()
session.add(
PartitionedAssetKeyLog(
asset_id=asset.id,
asset_event_id=event.id,
asset_partition_dag_run_id=pdr.id,
source_partition_key=source_key,
target_dag_id="many_to_one_dag",
target_partition_key="2024-08-01",
)
)
session.commit()

resp = test_client.get("/pending_partitioned_dag_run/many_to_one_dag?partition_key=2024-08-01")
assert resp.status_code == 200
body = resp.json()
assert body["total_required"] == 1
assert body["total_received"] == 1
asset_resp = body["assets"][0]
assert asset_resp["required_count"] == 1
assert asset_resp["received_count"] == 1
assert asset_resp["received"] is True
Loading