Skip to content
Open
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 @@ -37,7 +37,7 @@
from sqlalchemy import select

from airflow._shared.state import AssetScope, AssetStateStoreWriterKind
from airflow.api_fastapi.common.db.common import SessionDep
from airflow.api_fastapi.common.db.common import AsyncSessionDep
from airflow.api_fastapi.execution_api.datamodels.asset_state_store import (
AssetStateStorePutBody,
AssetStateStoreResponse,
Expand All @@ -64,9 +64,9 @@ class _TIWriterFields(NamedTuple):
try_number: int


def _fetch_ti_writer_fields(token: TIToken, session: SessionDep) -> _TIWriterFields:
async def _fetch_ti_writer_fields(token: TIToken, session: AsyncSessionDep) -> _TIWriterFields:
"""Return exact writer attribution for the execution identified by the token."""
row = session.execute(
result_row = await session.execute(
select(
TaskInstance.dag_id,
TaskInstance.run_id,
Expand All @@ -77,7 +77,8 @@ def _fetch_ti_writer_fields(token: TIToken, session: SessionDep) -> _TIWriterFie
TaskInstance.region_index,
TaskInstance.try_number,
).where(TaskInstance.id == token.id)
).one_or_none()
)
row = result_row.one_or_none()
if row is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
Expand All @@ -101,8 +102,10 @@ def _fetch_ti_writer_fields(token: TIToken, session: SessionDep) -> _TIWriterFie
)


def _resolve_asset_id_by_name(name: str, session: SessionDep) -> int:
asset_id = session.scalar(select(AssetModel.id).where(AssetModel.name == name, AssetModel.active.has()))
async def _resolve_asset_id_by_name(name: str, session: AsyncSessionDep) -> int:
asset_id = await session.scalar(
select(AssetModel.id).where(AssetModel.name == name, AssetModel.active.has())
)
if asset_id is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
Expand All @@ -111,8 +114,10 @@ def _resolve_asset_id_by_name(name: str, session: SessionDep) -> int:
return asset_id


def _resolve_asset_id_by_uri(uri: str, session: SessionDep) -> int:
asset_id = session.scalar(select(AssetModel.id).where(AssetModel.uri == uri, AssetModel.active.has()))
async def _resolve_asset_id_by_uri(uri: str, session: AsyncSessionDep) -> int:
asset_id = await session.scalar(
select(AssetModel.id).where(AssetModel.uri == uri, AssetModel.active.has())
)
if asset_id is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
Expand All @@ -122,14 +127,14 @@ def _resolve_asset_id_by_uri(uri: str, session: SessionDep) -> int:


@router.get("/by-name/value")
def get_asset_state_store_by_name(
async def get_asset_state_store_by_name(
name: Annotated[str, Query(min_length=1)],
key: Annotated[str, Query(min_length=1)],
session: SessionDep,
session: AsyncSessionDep,
) -> AssetStateStoreResponse:
"""Get an asset state store value by asset name."""
asset_id = _resolve_asset_id_by_name(name, session)
value = get_state_backend().get(AssetScope(asset_id=asset_id), key, session=session)
asset_id = await _resolve_asset_id_by_name(name, session)
value = await get_state_backend().aget(AssetScope(asset_id=asset_id), key, session=session)
if value is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
Expand All @@ -138,28 +143,28 @@ def get_asset_state_store_by_name(
return AssetStateStoreResponse(value=json.loads(value))


def _put_asset_state_store(
async def _put_asset_state_store(
scope: AssetScope,
key: str,
body: AssetStateStorePutBody,
token: TIToken,
session: SessionDep,
session: AsyncSessionDep,
) -> None:
backend = get_state_backend()
if isinstance(backend, MetastoreBackend):
if token.id == NULL_UUID:
# Since the asset state store routes do not have `task_instance_id` in their path params, the default kicks in which is"00000000-0000-0000-0000-000000000000"
backend.set_asset_state_store(
await backend.aset_asset_state_store(
scope,
key,
json.dumps(body.value),
kind=AssetStateStoreWriterKind.WATCHER,
session=session,
)
else:
writer = _fetch_ti_writer_fields(token, session)
writer = await _fetch_ti_writer_fields(token, session)

backend.set_asset_state_store(
await backend.aset_asset_state_store(
scope,
key,
json.dumps(body.value),
Expand All @@ -175,53 +180,53 @@ def _put_asset_state_store(
session=session,
)
else:
backend.set(scope, key, json.dumps(body.value), session=session)
await backend.aset(scope, key, json.dumps(body.value), session=session)


@router.put("/by-name/value", status_code=status.HTTP_204_NO_CONTENT)
def set_asset_state_store_by_name(
async def set_asset_state_store_by_name(
name: Annotated[str, Query(min_length=1)],
key: Annotated[str, Query(min_length=1)],
body: AssetStateStorePutBody,
session: SessionDep,
session: AsyncSessionDep,
token: TIToken = CurrentTIToken,
) -> None:
"""Set an asset state store value by asset name."""
_put_asset_state_store(
AssetScope(asset_id=_resolve_asset_id_by_name(name, session)), key, body, token, session
await _put_asset_state_store(
AssetScope(asset_id=await _resolve_asset_id_by_name(name, session)), key, body, token, session
)


@router.delete("/by-name/value", status_code=status.HTTP_204_NO_CONTENT)
def delete_asset_state_store_by_name(
async def delete_asset_state_store_by_name(
name: Annotated[str, Query(min_length=1)],
key: Annotated[str, Query(min_length=1)],
session: SessionDep,
session: AsyncSessionDep,
) -> None:
"""Delete a single asset state store key by asset name."""
asset_id = _resolve_asset_id_by_name(name, session)
get_state_backend().delete(AssetScope(asset_id=asset_id), key, session=session)
asset_id = await _resolve_asset_id_by_name(name, session)
await get_state_backend().adelete(AssetScope(asset_id=asset_id), key, session=session)


@router.delete("/by-name/clear", status_code=status.HTTP_204_NO_CONTENT)
def clear_asset_state_store_by_name(
async def clear_asset_state_store_by_name(
name: Annotated[str, Query(min_length=1)],
session: SessionDep,
session: AsyncSessionDep,
) -> None:
"""Delete all state store keys for an asset by asset name."""
asset_id = _resolve_asset_id_by_name(name, session)
get_state_backend().clear(AssetScope(asset_id=asset_id), session=session)
asset_id = await _resolve_asset_id_by_name(name, session)
await get_state_backend().aclear(AssetScope(asset_id=asset_id), session=session)


@router.get("/by-uri/value")
def get_asset_state_store_by_uri(
async def get_asset_state_store_by_uri(
uri: Annotated[str, Query(min_length=1)],
key: Annotated[str, Query(min_length=1)],
session: SessionDep,
session: AsyncSessionDep,
) -> AssetStateStoreResponse:
"""Get an asset state store value by asset URI."""
asset_id = _resolve_asset_id_by_uri(uri, session)
value = get_state_backend().get(AssetScope(asset_id=asset_id), key, session=session)
asset_id = await _resolve_asset_id_by_uri(uri, session)
value = await get_state_backend().aget(AssetScope(asset_id=asset_id), key, session=session)
if value is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
Expand All @@ -231,35 +236,34 @@ def get_asset_state_store_by_uri(


@router.put("/by-uri/value", status_code=status.HTTP_204_NO_CONTENT)
def set_asset_state_store_by_uri(
async def set_asset_state_store_by_uri(
uri: Annotated[str, Query(min_length=1)],
key: Annotated[str, Query(min_length=1)],
body: AssetStateStorePutBody,
session: SessionDep,
session: AsyncSessionDep,
token: TIToken = CurrentTIToken,
) -> None:
"""Set an asset state store value by asset URI."""
_put_asset_state_store(
AssetScope(asset_id=_resolve_asset_id_by_uri(uri, session)), key, body, token, session
)
asset_id = await _resolve_asset_id_by_uri(uri, session)
await _put_asset_state_store(AssetScope(asset_id=asset_id), key, body, token, session)


@router.delete("/by-uri/value", status_code=status.HTTP_204_NO_CONTENT)
def delete_asset_state_store_by_uri(
async def delete_asset_state_store_by_uri(
uri: Annotated[str, Query(min_length=1)],
key: Annotated[str, Query(min_length=1)],
session: SessionDep,
session: AsyncSessionDep,
) -> None:
"""Delete a single asset state store key by asset URI."""
asset_id = _resolve_asset_id_by_uri(uri, session)
get_state_backend().delete(AssetScope(asset_id=asset_id), key, session=session)
asset_id = await _resolve_asset_id_by_uri(uri, session)
await get_state_backend().adelete(AssetScope(asset_id=asset_id), key, session=session)


@router.delete("/by-uri/clear", status_code=status.HTTP_204_NO_CONTENT)
def clear_asset_state_store_by_uri(
async def clear_asset_state_store_by_uri(
uri: Annotated[str, Query(min_length=1)],
session: SessionDep,
session: AsyncSessionDep,
) -> None:
"""Delete all state store keys for an asset by asset URI."""
asset_id = _resolve_asset_id_by_uri(uri, session)
get_state_backend().clear(AssetScope(asset_id=asset_id), session=session)
asset_id = await _resolve_asset_id_by_uri(uri, session)
await get_state_backend().aclear(AssetScope(asset_id=asset_id), session=session)