amoghrajesh commented on code in PR #73706:
URL: https://github.com/apache/airflow/pull/73706#discussion_r4132540706


##########
providers/common/ai/src/airflow/providers/common/ai/operators/agent.py:
##########
@@ -239,16 +256,47 @@ class AgentOperator(CancellableAgentRunMixin, 
BaseOperator, HITLReviewMixin):
 
         A dict that omits ``request_limit`` still gets pydantic-ai's default of
         ``50`` requests -- pass ``"request_limit": None`` explicitly for no
-        request cap. See :ref:`howto/operator:llm` for the full set of caveats,
-        and :ref:`howto/operator:agent` for the ``durable=True`` replay
-        double-counting warning.
+        request cap.
+
+        On Airflow >= 3.3, when this is set, the limit counts usage across
+        every attempt of the task instance combined -- the initial run, every
+        retry, and every HITL regeneration all add to one running total kept
+        in the AIP-103 task state store under the ``__commonai_usage__`` key
+        -- rather than resetting on each attempt. This also applies to the
+        implicit ``request_limit=50`` default, which can now block a retry
+        that used to pass on its own. A step replayed by ``durable=True`` does
+        not count toward the total (see ``durable`` below). Clearing and
+        rerunning a *finished* (failed or
+        succeeded) task instance gets a fresh budget automatically; clearing
+        a *running* task instance does not bump ``max_tries``, so the
+        restarted attempt still sees the prior spend. To reset the budget for
+        a task instance that keeps retrying without a clear of a finished
+        attempt, delete the ``__commonai_usage__`` key via the Task State
+        Store UI. A worker killed with SIGKILL -- including after
+        ``on_kill``'s grace period expires, or an OOM kill -- cannot persist
+        that attempt's usage, so the next attempt's count under-represents
+        actual spend by that amount. To keep
+        the same effective per-attempt headroom this cross-attempt total used
+        to give each attempt on its own, scale each limit by
+        ``retries + 1``, or use ``usage_limits=None`` to opt back out. On
+        Airflow < 3.3, and whenever ``usage_limits`` is ``None``, each attempt
+        is still checked and counted on its own, as before. See
+        :ref:`howto/operator:llm` for the full set of caveats, and
+        :ref:`howto/operator:agent` for more on the cross-attempt budget.

Review Comment:
   Lot of repitition across files:
   
   
providers/common/ai/src/airflow/providers/common/ai/operators/agent.py:260-284
   providers/common/ai/docs/operators/agent.rst:300-322
   providers/common/ai/docs/changelog.rst:28-51
   
   Lets trim it down.



##########
providers/common/ai/tests/unit/common/ai/operators/test_agent.py:
##########
@@ -84,26 +98,93 @@ class Summary(BaseModel):
 
 
 def _make_mock_agent(output, make_mock_run_result, *, cost=None):
-    """Create a mock agent that returns the given output."""
+    """Create a mock agent that returns the given output.
+
+    ``run_sync``'s side effect also increments the ``usage=`` object it was 
called
+    with, mirroring what a real pydantic-ai run does to the ``RunUsage`` the 
operator
+    seeds in -- so a test reading the XCom/log usage the operator reports 
(computed
+    from that object, not from the mock's own ``.usage`` attribute) sees the 
value
+    configured here via ``cost=`` instead of an untouched, all-zero 
``RunUsage()``.
+    """
     mock_agent = MagicMock(spec=["run_sync", "instrument"])
     mock_agent.run_sync.return_value = make_mock_run_result(output, cost=cost)
+
+    def _increment_seeded_usage(*args, **kwargs):
+        # Reads .return_value late (not a captured value) so a test that 
customizes it
+        # after this call (e.g. ``mock_agent.run_sync.return_value.run_id = 
...``) still
+        # gets what it configured.
+        if (seeded := kwargs.get("usage")) is not None:
+            seeded.incr(RunUsage(requests=1, cost=cost))
+        return mock_agent.run_sync.return_value
+
+    mock_agent.run_sync.side_effect = _increment_seeded_usage
     return mock_agent
 
 
-def _make_ti(*, id="ti-1", dag_id="dag", task_id="task", run_id="run", 
map_index=-1, try_number=1):
+def _make_ti(
+    *, id="ti-1", dag_id="dag", task_id="task", run_id="run", map_index=-1, 
try_number=1, max_tries=0
+):
     """Return a task-instance double carrying the identity fields execute() 
reads."""
     ti = MagicMock()
     ti.configure_mock(
-        id=id, dag_id=dag_id, task_id=task_id, run_id=run_id, 
map_index=map_index, try_number=try_number
+        id=id,
+        dag_id=dag_id,
+        task_id=task_id,
+        run_id=run_id,
+        map_index=map_index,
+        try_number=try_number,
+        max_tries=max_tries,
     )
     return ti
 
 
-def _make_context(ti=None):
-    """A context whose ``task_instance`` is a configured ti. Other keys stay 
generic mocks."""
+def _make_task_state_store_accessor():
+    """A ``MagicMock(spec=TaskStateStoreAccessor)`` backed by a plain dict.
+
+    For a test that engages the usage budget (``usage_limits`` set, Airflow >= 
3.3): a
+    bare ``MagicMock()``'s ``.get()`` returns a non-``None``, non-dict value, 
which trips
+    ``TaskStateStoreUsageBudget.load()``'s malformed-record ``ValueError``.
+
+    ``TaskStateStoreAccessor`` doesn't exist below Airflow 3.3; several 
callers of this
+    helper exercise ``execute()`` paths (e.g. ``usage_limits`` forwarding) 
that don't
+    depend on the task state store at all on those cores -- 
``_build_usage_budget``
+    returns ``None`` before ever touching ``context["task_state_store"]``. 
Falling back
+    to a plain method-name spec keeps this helper importable there too, 
instead of
+    forcing every caller to skip on Airflow version for a dependency they 
don't have.
+    """
+    try:
+        from airflow.sdk.execution_time.context import TaskStateStoreAccessor
+    except ImportError:
+        spec = ["get", "set", "delete"]
+    else:
+        spec = TaskStateStoreAccessor
+
+    store = {}
+    accessor = MagicMock(spec=spec)
+    accessor.get.side_effect = lambda key, default=None: store.get(key, 
default)
+    accessor.set.side_effect = lambda key, value, retention=None: 
store.__setitem__(key, value)
+    accessor.delete.side_effect = lambda key: store.pop(key, None)
+    return accessor
+
+
+def _make_context(ti=None, task_state_store=None):
+    """A context whose ``task_instance`` is a configured ti. Other keys stay 
generic mocks.
+
+    :param task_state_store: Backs ``context["task_state_store"]`` when given 
-- pass
+        :func:`_make_task_state_store_accessor` for a test that sets 
``usage_limits``
+        on Airflow >= 3.3 (see that helper's docstring for why the default 
won't do).
+    """
     ti = ti if ti is not None else _make_ti()
+
+    def _getitem(key):
+        if key == "task_instance":
+            return ti
+        if key == "task_state_store" and task_state_store is not None:
+            return task_state_store
+        return MagicMock()

Review Comment:
   Use spec or autospec pls.



-- 
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