1616# under the License.
1717from __future__ import annotations
1818
19- from typing import TYPE_CHECKING , Any , Protocol
19+ from typing import TYPE_CHECKING , Protocol
2020from uuid import UUID
2121
2222import attrs
4949 from airflow .serialization .definitions .dag import SerializedDAG , SerializedOperator
5050
5151
52+ LOOP_GATE_OPERATOR = "LoopGateOperator"
53+ """Operator class name the Task SDK gives the gate task of a loop."""
54+
55+
56+ @attrs .frozen (kw_only = True )
57+ class _ProducerRequest :
58+ dag_id : str
59+ run_id : str
60+ task_id : str
61+ is_mapped : bool
62+ context : ProducerContext | None
63+ map_indexes : int | Collection [int ] | None
64+ region_id : UUID | None
65+ region_index : int | None
66+
67+
5268class TaskCoordinate (Protocol ):
5369 """Stored task identity shared by live and retained task data."""
5470
@@ -249,7 +265,12 @@ def has_regions(self, dag_id: str, run_id: str | None, task_id: str) -> bool:
249265 return bool (self .session .scalar (select (query .exists ())))
250266
251267 def resolve_dependency (self , caller : TaskInstance , task_id : str ) -> tuple [TaskInstance , ...]:
252- producer = self .get_task (caller .dag_id , caller .run_id , task_id , dag_version_id = caller .dag_version_id )
268+ try :
269+ producer = self .get_task (
270+ caller .dag_id , caller .run_id , task_id , dag_version_id = caller .dag_version_id
271+ )
272+ except TaskNotFound :
273+ return self .resolve (dag_id = caller .dag_id , run_id = caller .run_id , task_id = task_id , caller = caller )
253274 loop = enclosing_loop (producer )
254275 caller_task = self .get_task (
255276 caller .dag_id , caller .run_id , caller .task_id , dag_version_id = caller .dag_version_id
@@ -278,6 +299,18 @@ def resolve_dependency(self, caller: TaskInstance, task_id: str) -> tuple[TaskIn
278299 return tuple (gates [:1 ])
279300 return self .resolve (dag_id = caller .dag_id , run_id = caller .run_id , task_id = task_id , caller = caller )
280301
302+ @staticmethod
303+ def _filter_region_indexes (
304+ query : Select , region_index : int | None , map_indexes : int | Collection [int ] | None
305+ ) -> Select :
306+ if region_index is not None :
307+ query = query .where (TaskInstance .region_index == region_index )
308+ if isinstance (map_indexes , int ):
309+ query = query .where (TaskInstance .region_index == map_indexes )
310+ elif map_indexes is not None :
311+ query = query .where (TaskInstance .region_index .in_ (map_indexes ))
312+ return query
313+
281314 def _filter_legacy_producers (
282315 self ,
283316 query : Select ,
@@ -295,13 +328,7 @@ def _filter_legacy_producers(
295328 TaskInstance .task_id == task_id ,
296329 TaskInstance .region_id == SENTINEL_REGION_ID ,
297330 )
298- if region_index is not None :
299- query = query .where (TaskInstance .region_index == region_index )
300- if isinstance (map_indexes , int ):
301- query = query .where (TaskInstance .region_index == map_indexes )
302- elif map_indexes is not None :
303- query = query .where (TaskInstance .region_index .in_ (map_indexes ))
304- return query
331+ return self ._filter_region_indexes (query , region_index , map_indexes )
305332
306333 def _filter_removed_task_producers (
307334 self ,
@@ -324,13 +351,7 @@ def _filter_removed_task_producers(
324351 )
325352 if region_id is not None :
326353 query = query .where (TaskInstance .region_id == region_id )
327- if region_index is not None :
328- query = query .where (TaskInstance .region_index == region_index )
329- if isinstance (map_indexes , int ):
330- query = query .where (TaskInstance .region_index == map_indexes )
331- elif map_indexes is not None :
332- query = query .where (TaskInstance .region_index .in_ (map_indexes ))
333- return query
354+ return self ._filter_region_indexes (query , region_index , map_indexes )
334355
335356 def _is_legacy_lookup (
336357 self , dag_id : str , run_id : str , task_id : str , region_id : UUID | None , previous_iteration : bool
@@ -351,7 +372,7 @@ def _build_producer_request(
351372 region_index : int | None ,
352373 map_indexes : int | Collection [int ] | None ,
353374 previous_iteration : bool ,
354- ) -> dict [ str , Any ] | None :
375+ ) -> _ProducerRequest | None :
355376 shared_run = caller is not None and (caller .dag_id , caller .run_id ) == (dag_id , run_id )
356377 try :
357378 task = (
@@ -378,16 +399,16 @@ def _build_producer_request(
378399 )
379400 if previous_iteration and context is None :
380401 raise ValueError ("Previous-iteration lookup requires a shared loop scope" )
381- return {
382- " dag_id" : dag_id ,
383- " run_id" : run_id ,
384- " task_id" : task_id ,
385- " is_mapped" : task .get_needs_expansion (),
386- " context" : context ,
387- " map_indexes" : map_indexes ,
388- " region_id" : region_id ,
389- " region_index" : region_index ,
390- }
402+ return _ProducerRequest (
403+ dag_id = dag_id ,
404+ run_id = run_id ,
405+ task_id = task_id ,
406+ is_mapped = task .get_needs_expansion (),
407+ context = context ,
408+ map_indexes = map_indexes ,
409+ region_id = region_id ,
410+ region_index = region_index ,
411+ )
391412
392413 def resolve (
393414 self ,
@@ -437,7 +458,7 @@ def resolve(
437458 ).order_by (TaskInstance .region_index )
438459 )
439460 )
440- return resolve_current_producers (** request , session = self .session )
461+ return resolve_current_producers (** attrs . asdict ( request , recurse = False ) , session = self .session )
441462
442463 def select_skip_target_ids (
443464 self , * , caller : TaskInstance , task_id : str , map_indexes : int | None = None
@@ -520,4 +541,4 @@ def select_producer_ids(
520541 region_index = region_index ,
521542 map_indexes = map_indexes ,
522543 )
523- return select_current_producer_ids (** request , session = self .session )
544+ return select_current_producer_ids (** attrs . asdict ( request , recurse = False ) , session = self .session )
0 commit comments