kaxil commented on code in PR #74351:
URL: https://github.com/apache/airflow/pull/74351#discussion_r4198266849


##########
providers/common/ai/src/airflow/providers/common/ai/plugins/hitl_review.py:
##########
@@ -87,31 +88,51 @@ def _get_bundle_url() -> str:
     from airflow.sdk import TaskInstanceState
     from airflow.utils.session import create_session
 
+    if AIRFLOW_V_3_4_PLUS:
+        from airflow.api_fastapi.core_api.services.public.task_coordinates 
import resolve_task_scope
+        from airflow.models.dagbag import DBDagBag
+
     def _get_session():
         with create_session(scoped=False) as session:
             yield session
 
     SessionDep = Annotated[Session, Depends(_get_session)]
 
     def _read_xcom(
-        session: Session, *, dag_id: str, run_id: str, task_id: str, 
map_index: int = -1, key: str
+        session: Session,
+        *,
+        dag_id: str,
+        run_id: str,
+        task_id: str,
+        map_index: int = -1,
+        region_id: UUID | None = None,
+        key: str,
     ):
         """Read a single XCom value from the database."""
+        scope = {"region_id": region_id} if AIRFLOW_V_3_4_PLUS else {}

Review Comment:
   When the scope resolves to the sentinel (a client passes 
`region_id=00000000-0000-0000-0000-000000000000` explicitly for a top-level 
mapped task), `get_many` keeps its default `include_node_regions=True`, so this 
read also admits the task's own node region and returns slot 2's session. 
`_write_xcom` and `_is_task_completed` match `TI.region_id` exactly, so the 
same request reports `task_completed: true` and a follow-up approve 404s. Since 
the scope is already resolved by this point, should both `get_many` calls here 
pass `include_node_regions=False`? The `region_id=None` default also means 
different things across the four helpers (every region for the reads, the 
sentinel for the write and the completion check), so making it required would 
stop them drifting apart.



##########
providers/common/ai/src/airflow/providers/common/ai/plugins/www/tests/main.test.mjs:
##########
@@ -0,0 +1,43 @@
+/*!
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *   http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied.  See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+import assert from "node:assert/strict";
+import { after, test } from "node:test";
+import { createElement } from "react";
+import { renderToStaticMarkup } from "react-dom/server";
+
+import { createServer } from "vite";
+
+const server = await createServer({ server: { middlewareMode: true } });

Review Comment:
   As far as I can tell nothing runs these: `package.json` has no `test` 
script, and the prek hook (`ts_compile_lint_common_ai.py`) only runs eslint, 
prettier and tsc. This one also needs Vite started from `www/`, since 
`ssrLoadModule("src/main.tsx")` resolves against the cwd. Could you add a 
`test` script (`node --test tests/`) and call it from the existing hook or a CI 
step, so the unresolved-host guard in `main.tsx` stays covered?



##########
providers/common/ai/src/airflow/providers/common/ai/plugins/hitl_review.py:
##########
@@ -273,6 +325,33 @@ def _get_map_index(q: str = Query("-1", 
alias="map_index")) -> int:
 
     MapIndexDep = Annotated[int, Depends(_get_map_index)]
 
+    def _get_task_scope(
+        db: SessionDep,
+        dag_id: str,
+        run_id: str,
+        task_id: str,
+        map_index: MapIndexDep,
+        region_id: UUID | None = None,

Review Comment:
   The REST API section in `providers/common/ai/docs/hitl_review.rst` still 
lists only `dag_id`, `run_id`, `task_id` and `map_index` as the common query 
parameters. Could you add `region_id` / `region_index` there (sent together for 
a loop pass, only on hosts that have regions), along with the new 400 for a 
loop task addressed without them and the 409 for an ambiguous selection?



##########
devel-common/src/tests_common/test_utils/taskinstance.py:
##########
@@ -146,13 +146,23 @@ def run_task_instance(
     from airflow.sdk.definitions.dag import _run_task
 
     # Session handling is a mess in tests; use a fresh ti to run the task.
-    new_ti = TaskInstance.get_task_instance(
-        dag_id=ti.dag_id,
-        run_id=ti.run_id,
-        task_id=ti.task_id,
-        map_index=ti.map_index,
-        **session_kwargs,
-    )
+    if AIRFLOW_V_3_4_PLUS:
+        new_ti = TaskInstance.get_task_instance(
+            dag_id=ti.dag_id,
+            run_id=ti.run_id,
+            task_id=ti.task_id,
+            map_index=ti.region_index,
+            region_id=ti.region_id,
+            **session_kwargs,
+        )
+    else:
+        new_ti = TaskInstance.get_task_instance(
+            dag_id=ti.dag_id,
+            run_id=ti.run_id,
+            task_id=ti.task_id,
+            map_index=ti.map_index,
+            **session_kwargs,
+        )

Review Comment:
   `TaskInstance.map_index` is a synonym for `region_index`, so these two 
branches only differ by `region_id`. One call with `map_index=ti.map_index` and 
a `region_id=ti.region_id` kwarg added only when the host has regions would do. 
Small note on the rationale too: without a region, `get_task_instance` applies 
`public_region_filter`, which excludes loop regions, so the old call returned 
`None` and the helper fell back to the caller's `ti` rather than running 
another pass's row.



##########
providers/common/ai/src/airflow/providers/common/ai/plugins/www/src/main.tsx:
##########
@@ -40,27 +41,36 @@ export interface PluginComponentProps {
  * host (from route params). Renders ChatPage when all params are present,
  * otherwise shows NoSession fallback.
  */
-const PluginComponent: FC<PluginComponentProps> = ({
-  dagId = "",
-  runId = "",
-  taskId = "",
-  mapIndex: mapIndexProp = "-1",
-}) => {
+const PluginComponent: FC<PluginComponentProps> = (props) => {
+  const {
+    dagId = "",
+    runId = "",
+    taskId = "",
+    mapIndex: mapIndexProp = "-1",
+    taskInstance,
+  } = props;
   const mapIndex = /^-?\d+$/.test(String(mapIndexProp)) ? 
parseInt(String(mapIndexProp), 10) : -1;
 
-  if (!dagId || !runId || !taskId) {
+  if (!dagId || !runId || !taskId || ("taskInstance" in props && taskInstance 
== null)) {
     return <NoSession />;
   }
 
-  return <ChatPage dagId={dagId} runId={runId} taskId={taskId} 
mapIndex={mapIndex} />;
+  const region = taskInstance?.region_id !== undefined && 
taskInstance.region_index !== undefined
+    ? { region_id: taskInstance.region_id, region_index: 
taskInstance.region_index }
+    : undefined;
+
+  return <ChatPage
+    key={`${dagId}/${runId}/${taskId}/${mapIndex}/${taskInstance?.id ?? 
""}/${region?.region_id ?? ""}/${region?.region_index ?? ""}`}
+    dagId={dagId} runId={runId} taskId={taskId} mapIndex={mapIndex} 
region={region}
+  />;

Review Comment:
   This key does real work (`useSession` keeps the first `createApi(...)` in a 
`useRef`, so without the remount a selection change would keep sending the 
previous pass's region), but nothing here says so. Could you pull it into a 
named `sessionKey` with a one-line comment, and put the props one per line like 
the rest of the plugin?



##########
providers/common/ai/tests/unit/common/ai/plugins/test_hitl_review.py:
##########
@@ -998,6 +1015,94 @@ def test_react_apps_registered(self):
         assert "main.umd.cjs" in app["bundle_url"]
 
 
[email protected](not AIRFLOW_V_3_4_PLUS, reason="Regional task identity 
requires Airflow 3.4+")
+class TestRegionalReview:
+    @pytest.fixture
+    def regional_review(self, dag_maker, session):
+        _clear_db()
+
+        @task_group
+        def body():
+            EmptyOperator(task_id="review")
+
+        with dag_maker(TEST_DAG_ID, serialized=True):
+            create_loop(body, max_iterations=2)
+        dr = dag_maker.create_dagrun(run_id=TEST_RUN_ID)
+        first = next(ti for ti in dr.task_instances if ti.task_id == 
"body.review")
+        first.state = "success"
+        second = TaskInstance(
+            task=dag_maker.serialized_dag.get_task(first.task_id),
+            run_id=dr.run_id,
+            dag_version_id=first.dag_version_id,
+            region_id=first.region_id,
+            region_index=1,
+            state="deferred",
+        )
+        session.add(second)
+        for ti in (first, second):
+            output = f"pass {ti.region_index}"
+            values = {
+                XCOM_AGENT_SESSION: AgentSessionData(
+                    status=SessionStatus.PENDING_REVIEW,
+                    iteration=1,
+                    max_iterations=5,
+                    current_output=output,
+                ).model_dump(mode="json"),
+                f"{XCOM_AGENT_OUTPUT_PREFIX}1": output,
+            }
+            session.flush()
+            for key, value in values.items():
+                XComModel.set_for_attempt(
+                    task_instance_id=ti.id, key=key, value=value, 
serialize=False, session=session
+                )
+        session.commit()
+        yield second
+        _clear_db()
+
+    @pytest.mark.parametrize("action", ["find", "feedback", "approve", 
"reject"])
+    def test_exact_scope_preserves_other_loop_pass(self, regional_review, 
test_client, session, action):
+        ti = regional_review
+        params = {
+            "dag_id": ti.dag_id,
+            "run_id": ti.run_id,
+            "task_id": ti.task_id,
+            "map_index": -1,
+            "region_id": str(ti.region_id),
+            "region_index": ti.region_index,
+        }
+        if action == "find":
+            response = test_client.get("/hitl-review/sessions/find", 
params=params)
+        else:
+            response = test_client.post(
+                f"/hitl-review/sessions/{action}", params=params, 
json={"feedback": "revise"}
+            )
+
+        assert response.status_code == 200, response.text
+        data = response.json()
+        assert data["current_output"] == "pass 1"
+        assert data["conversation"][0]["content"] == "pass 1"
+        assert data["task_completed"] is False
+        session.expire_all()
+        earlier = session.scalar(
+            select(XComModel.value).where(
+                XComModel.key == XCOM_AGENT_SESSION,
+                XComModel.region_id == ti.region_id,
+                XComModel.region_index == 0,
+            )
+        )
+        assert earlier["status"] == "pending_review"
+        assert earlier["current_output"] == "pass 0"
+
+    def test_loop_scope_is_required(self, regional_review, test_client):

Review Comment:
   These only send the explicit `region_id` + `region_index` pair, which 
returns from `resolve_task_scope` before the resolver runs. A top-level mapped 
slot addressed with just `map_index` (resolver returns the node region, then 
`_write_xcom` has to match `TI.region_id` to it) and the 400 a region-less host 
returns for region selectors have no test, and the file has no mapped task at 
all. Could you add a mapped case (sessions on slots 0 and 1, approve slot 1 
with only `map_index=1`, assert slot 0 is untouched) and one that patches the 
version-compat flag so the host looks region-less and sends `region_id`? 
Asserting `data["status"]` for approve/reject above would also show the write 
landed on pass 1.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to