diff --git a/airflow-core/src/airflow/api_fastapi/execution_api/routes/asset_state_store.py b/airflow-core/src/airflow/api_fastapi/execution_api/routes/asset_state_store.py index f5d387b495ad6..bd99c88d01320 100644 --- a/airflow-core/src/airflow/api_fastapi/execution_api/routes/asset_state_store.py +++ b/airflow-core/src/airflow/api_fastapi/execution_api/routes/asset_state_store.py @@ -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, @@ -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, @@ -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, @@ -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, @@ -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, @@ -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, @@ -138,18 +143,18 @@ 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), @@ -157,9 +162,9 @@ def _put_asset_state_store( 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), @@ -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, @@ -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)