Lee-W commented on code in PR #73706:
URL: https://github.com/apache/airflow/pull/73706#discussion_r4140388684


##########
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:
   Changed it to `MagicMock(spec=dict)` and  `MagicMock(spec=[])`.



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