mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-11 06:28:58 +00:00
fix: read run stop_reason from runtime context (#4188)
* fix: read run stop_reason from runtime context * fix: address review feedback for #4188 stop_reason integration - migration 0005: use safe_add_column for consistency and drift detection - worker: clear runtime.context stop_reason at start of each _stream_once turn so a clean continuation doesn't inherit a prior cap reason - tests: replace circular unit test with real middleware integration tests that exercise LoopDetectionMiddleware._apply and TokenBudgetMiddleware._apply through the worker, proving the full middleware → runtime.context → persist pipeline * fix(test): resume conftest * fix: stamp stop_reason in all guard middlewares, fix clearing semantics
This commit is contained in:
parent
de55982c5a
commit
1769b2de0d
@ -126,6 +126,7 @@ class RunResponse(BaseModel):
|
|||||||
subagent_tokens: int = 0
|
subagent_tokens: int = 0
|
||||||
middleware_tokens: int = 0
|
middleware_tokens: int = 0
|
||||||
message_count: int = 0
|
message_count: int = 0
|
||||||
|
stop_reason: str | None = None
|
||||||
|
|
||||||
|
|
||||||
class ThreadTokenUsageModelBreakdown(BaseModel):
|
class ThreadTokenUsageModelBreakdown(BaseModel):
|
||||||
@ -230,6 +231,7 @@ def _record_to_response(record: RunRecord) -> RunResponse:
|
|||||||
subagent_tokens=record.subagent_tokens,
|
subagent_tokens=record.subagent_tokens,
|
||||||
middleware_tokens=record.middleware_tokens,
|
middleware_tokens=record.middleware_tokens,
|
||||||
message_count=record.message_count,
|
message_count=record.message_count,
|
||||||
|
stop_reason=record.stop_reason,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -592,6 +592,11 @@ class LoopDetectionMiddleware(AgentMiddleware[AgentState]):
|
|||||||
run_id = self._get_run_id(runtime)
|
run_id = self._get_run_id(runtime)
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._stop_reason[run_id] = "loop_capped"
|
self._stop_reason[run_id] = "loop_capped"
|
||||||
|
# Also write to runtime.context so the lead worker can read it
|
||||||
|
# without needing a reference to this middleware instance (#4176).
|
||||||
|
ctx = getattr(runtime, "context", None)
|
||||||
|
if isinstance(ctx, dict):
|
||||||
|
ctx["stop_reason"] = "loop_capped"
|
||||||
# Strip tool_calls from the last AIMessage to force text output.
|
# Strip tool_calls from the last AIMessage to force text output.
|
||||||
# Once tool_calls are stripped, the AIMessage no longer requires
|
# Once tool_calls are stripped, the AIMessage no longer requires
|
||||||
# matching ToolMessage responses, so mutating it in place here
|
# matching ToolMessage responses, so mutating it in place here
|
||||||
|
|||||||
@ -283,6 +283,11 @@ class SafetyFinishReasonMiddleware(AgentMiddleware[AgentState]):
|
|||||||
if termination is None:
|
if termination is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
# Stamp stop_reason so the worker can surface this capped completion
|
||||||
|
# alongside loop_capped / token_capped (#4176).
|
||||||
|
ctx = getattr(runtime, "context", None)
|
||||||
|
if isinstance(ctx, dict):
|
||||||
|
ctx["stop_reason"] = "safety_capped"
|
||||||
patched = self._build_suppressed_message(last, termination)
|
patched = self._build_suppressed_message(last, termination)
|
||||||
|
|
||||||
thread_id = None
|
thread_id = None
|
||||||
|
|||||||
@ -154,6 +154,12 @@ class SubagentLimitMiddleware(AgentMiddleware[AgentState]):
|
|||||||
prior_delegation_count,
|
prior_delegation_count,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Stamp stop_reason when the total per-run cap is exhausted so the
|
||||||
|
# worker surfaces this capped completion alongside loop_capped /
|
||||||
|
# token_capped / safety_capped (#4176).
|
||||||
|
if remaining_total == 0 and isinstance(getattr(runtime, "context", None), dict):
|
||||||
|
runtime.context["stop_reason"] = "subagent_limit_capped"
|
||||||
|
|
||||||
# Replace the AIMessage with truncated tool_calls (same id triggers replacement)
|
# Replace the AIMessage with truncated tool_calls (same id triggers replacement)
|
||||||
content = _append_text(last_msg.content, _TOTAL_LIMIT_STOP_MSG) if remaining_total == 0 else None
|
content = _append_text(last_msg.content, _TOTAL_LIMIT_STOP_MSG) if remaining_total == 0 else None
|
||||||
updated_msg = clone_ai_message_with_tool_calls(last_msg, truncated_tool_calls, content=content)
|
updated_msg = clone_ai_message_with_tool_calls(last_msg, truncated_tool_calls, content=content)
|
||||||
|
|||||||
@ -251,6 +251,11 @@ class TokenBudgetMiddleware(AgentMiddleware[AgentState]):
|
|||||||
# returns (the hard stop itself does not raise). See
|
# returns (the hard stop itself does not raise). See
|
||||||
# ``consume_stop_reason``.
|
# ``consume_stop_reason``.
|
||||||
self._stop_reason[run_id] = "token_capped"
|
self._stop_reason[run_id] = "token_capped"
|
||||||
|
# Also write to runtime.context so the lead worker can read it
|
||||||
|
# without needing a reference to this middleware instance (#4176).
|
||||||
|
ctx = getattr(runtime, "context", None)
|
||||||
|
if isinstance(ctx, dict):
|
||||||
|
ctx["stop_reason"] = "token_capped"
|
||||||
stop_text = _BUDGET_EXCEEDED_MSG.format(reason=trigger_reason, used=trigger_used, budget=trigger_budget)
|
stop_text = _BUDGET_EXCEEDED_MSG.format(reason=trigger_reason, used=trigger_used, budget=trigger_budget)
|
||||||
return self._build_hard_stop_update(last_msg, stop_text)
|
return self._build_hard_stop_update(last_msg, stop_text)
|
||||||
|
|
||||||
|
|||||||
@ -0,0 +1,28 @@
|
|||||||
|
"""run stop_reason
|
||||||
|
|
||||||
|
Revision ID: 0005_run_stop_reason
|
||||||
|
Revises: 0004_run_ownership
|
||||||
|
Create Date: 2026-07-15
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "0005_run_stop_reason"
|
||||||
|
down_revision: str | Sequence[str] | None = "0004_run_ownership"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
from deerflow.persistence.migrations._helpers import safe_add_column
|
||||||
|
|
||||||
|
safe_add_column("runs", sa.Column("stop_reason", sa.String(50), nullable=True))
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_column("runs", "stop_reason")
|
||||||
@ -25,6 +25,7 @@ class RunRow(Base):
|
|||||||
metadata_json: Mapped[dict] = mapped_column(JSON, default=dict)
|
metadata_json: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||||
kwargs_json: Mapped[dict] = mapped_column(JSON, default=dict)
|
kwargs_json: Mapped[dict] = mapped_column(JSON, default=dict)
|
||||||
error: Mapped[str | None] = mapped_column(Text)
|
error: Mapped[str | None] = mapped_column(Text)
|
||||||
|
stop_reason: Mapped[str | None] = mapped_column(String(50))
|
||||||
|
|
||||||
# Convenience fields (for listing pages without querying RunEventStore)
|
# Convenience fields (for listing pages without querying RunEventStore)
|
||||||
message_count: Mapped[int] = mapped_column(default=0)
|
message_count: Mapped[int] = mapped_column(default=0)
|
||||||
|
|||||||
@ -96,6 +96,7 @@ class RunRepository(RunStore):
|
|||||||
metadata=None,
|
metadata=None,
|
||||||
kwargs=None,
|
kwargs=None,
|
||||||
error=None,
|
error=None,
|
||||||
|
stop_reason: str | None = None,
|
||||||
created_at=None,
|
created_at=None,
|
||||||
follow_up_to_run_id=None,
|
follow_up_to_run_id=None,
|
||||||
owner_worker_id: str | None = None,
|
owner_worker_id: str | None = None,
|
||||||
@ -121,6 +122,7 @@ class RunRepository(RunStore):
|
|||||||
"metadata_json": self._safe_json(metadata) or {},
|
"metadata_json": self._safe_json(metadata) or {},
|
||||||
"kwargs_json": self._safe_json(kwargs) or {},
|
"kwargs_json": self._safe_json(kwargs) or {},
|
||||||
"error": error,
|
"error": error,
|
||||||
|
"stop_reason": stop_reason,
|
||||||
"follow_up_to_run_id": follow_up_to_run_id,
|
"follow_up_to_run_id": follow_up_to_run_id,
|
||||||
"owner_worker_id": owner_worker_id,
|
"owner_worker_id": owner_worker_id,
|
||||||
"lease_expires_at": lease_dt,
|
"lease_expires_at": lease_dt,
|
||||||
@ -203,10 +205,12 @@ class RunRepository(RunStore):
|
|||||||
result = await session.execute(stmt)
|
result = await session.execute(stmt)
|
||||||
return {row.run_id: self._row_to_dict(row) for row in result.scalars()}
|
return {row.run_id: self._row_to_dict(row) for row in result.scalars()}
|
||||||
|
|
||||||
async def update_status(self, run_id, status, *, error=None) -> bool:
|
async def update_status(self, run_id, status, *, error=None, stop_reason=None) -> bool:
|
||||||
values: dict[str, Any] = {"status": status, "updated_at": datetime.now(UTC)}
|
values: dict[str, Any] = {"status": status, "updated_at": datetime.now(UTC)}
|
||||||
if error is not None:
|
if error is not None:
|
||||||
values["error"] = error
|
values["error"] = error
|
||||||
|
if stop_reason is not None:
|
||||||
|
values["stop_reason"] = stop_reason
|
||||||
# Guard: only transition rows that are still active. ``interrupted`` is
|
# Guard: only transition rows that are still active. ``interrupted`` is
|
||||||
# included because the rollback path goes ``running → interrupted``
|
# included because the rollback path goes ``running → interrupted``
|
||||||
# (cancel acknowledged) then ``interrupted → error`` (task finalize).
|
# (cancel acknowledged) then ``interrupted → error`` (task finalize).
|
||||||
|
|||||||
@ -181,6 +181,7 @@ class RunRecord:
|
|||||||
finalizing: bool = False
|
finalizing: bool = False
|
||||||
owner_worker_id: str | None = None
|
owner_worker_id: str | None = None
|
||||||
lease_expires_at: str | None = None
|
lease_expires_at: str | None = None
|
||||||
|
stop_reason: str | None = None
|
||||||
|
|
||||||
|
|
||||||
class RunManager:
|
class RunManager:
|
||||||
@ -244,7 +245,7 @@ class RunManager:
|
|||||||
return [record for run_id in run_ids if (record := self._runs.get(run_id)) is not None]
|
return [record for run_id in run_ids if (record := self._runs.get(run_id)) is not None]
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _store_put_payload(record: RunRecord, *, error: str | None = None) -> dict[str, Any]:
|
def _store_put_payload(record: RunRecord, *, error: str | None = None, stop_reason: str | None = None) -> dict[str, Any]:
|
||||||
payload = {
|
payload = {
|
||||||
"thread_id": record.thread_id,
|
"thread_id": record.thread_id,
|
||||||
"assistant_id": record.assistant_id,
|
"assistant_id": record.assistant_id,
|
||||||
@ -260,6 +261,8 @@ class RunManager:
|
|||||||
}
|
}
|
||||||
if record.user_id is not None:
|
if record.user_id is not None:
|
||||||
payload["user_id"] = record.user_id
|
payload["user_id"] = record.user_id
|
||||||
|
if record.stop_reason is not None:
|
||||||
|
payload["stop_reason"] = record.stop_reason
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
async def _call_store_with_retry(
|
async def _call_store_with_retry(
|
||||||
@ -331,16 +334,16 @@ class RunManager:
|
|||||||
self._store_put_payload(record, error=error),
|
self._store_put_payload(record, error=error),
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _persist_status(self, record: RunRecord, status: RunStatus, *, error: str | None = None) -> bool:
|
async def _persist_status(self, record: RunRecord, status: RunStatus, *, error: str | None = None, stop_reason: str | None = None) -> bool:
|
||||||
"""Best-effort persist a status transition to the backing store."""
|
"""Best-effort persist a status transition to the backing store."""
|
||||||
if self._store is None:
|
if self._store is None:
|
||||||
return True
|
return True
|
||||||
row_recovery_payload = self._store_put_payload(record, error=error)
|
row_recovery_payload = self._store_put_payload(record, error=error, stop_reason=stop_reason)
|
||||||
try:
|
try:
|
||||||
updated = await self._call_store_with_retry(
|
updated = await self._call_store_with_retry(
|
||||||
"update_status",
|
"update_status",
|
||||||
record.run_id,
|
record.run_id,
|
||||||
lambda: self._store.update_status(record.run_id, status.value, error=error),
|
lambda: self._store.update_status(record.run_id, status.value, error=error, stop_reason=stop_reason),
|
||||||
)
|
)
|
||||||
if updated is False:
|
if updated is False:
|
||||||
# ``update_status`` is now guarded by ``status IN ('pending','running')``.
|
# ``update_status`` is now guarded by ``status IN ('pending','running')``.
|
||||||
@ -407,6 +410,7 @@ class RunManager:
|
|||||||
first_human_message=row.get("first_human_message"),
|
first_human_message=row.get("first_human_message"),
|
||||||
owner_worker_id=row.get("owner_worker_id"),
|
owner_worker_id=row.get("owner_worker_id"),
|
||||||
lease_expires_at=row.get("lease_expires_at"),
|
lease_expires_at=row.get("lease_expires_at"),
|
||||||
|
stop_reason=row.get("stop_reason"),
|
||||||
)
|
)
|
||||||
|
|
||||||
async def update_run_completion(self, run_id: str, **kwargs) -> None:
|
async def update_run_completion(self, run_id: str, **kwargs) -> None:
|
||||||
@ -656,7 +660,7 @@ class RunManager:
|
|||||||
logger.warning("Failed to map store row for run %s", run_id, exc_info=True)
|
logger.warning("Failed to map store row for run %s", run_id, exc_info=True)
|
||||||
return records_by_id
|
return records_by_id
|
||||||
|
|
||||||
async def set_status(self, run_id: str, status: RunStatus, *, error: str | None = None) -> None:
|
async def set_status(self, run_id: str, status: RunStatus, *, error: str | None = None, stop_reason: str | None = None) -> None:
|
||||||
"""Transition a run to a new status."""
|
"""Transition a run to a new status."""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
record = self._runs.get(run_id)
|
record = self._runs.get(run_id)
|
||||||
@ -667,7 +671,9 @@ class RunManager:
|
|||||||
record.updated_at = _now_iso()
|
record.updated_at = _now_iso()
|
||||||
if error is not None:
|
if error is not None:
|
||||||
record.error = error
|
record.error = error
|
||||||
await self._persist_status(record, status, error=error)
|
if stop_reason is not None:
|
||||||
|
record.stop_reason = stop_reason
|
||||||
|
await self._persist_status(record, status, error=error, stop_reason=stop_reason)
|
||||||
logger.info("Run %s -> %s", run_id, status.value)
|
logger.info("Run %s -> %s", run_id, status.value)
|
||||||
|
|
||||||
async def set_finalizing(self, run_id: str, finalizing: bool) -> None:
|
async def set_finalizing(self, run_id: str, finalizing: bool) -> None:
|
||||||
|
|||||||
@ -29,6 +29,7 @@ class RunStore(abc.ABC):
|
|||||||
metadata: dict[str, Any] | None = None,
|
metadata: dict[str, Any] | None = None,
|
||||||
kwargs: dict[str, Any] | None = None,
|
kwargs: dict[str, Any] | None = None,
|
||||||
error: str | None = None,
|
error: str | None = None,
|
||||||
|
stop_reason: str | None = None,
|
||||||
created_at: str | None = None,
|
created_at: str | None = None,
|
||||||
owner_worker_id: str | None = None,
|
owner_worker_id: str | None = None,
|
||||||
lease_expires_at: str | None = None,
|
lease_expires_at: str | None = None,
|
||||||
@ -84,6 +85,7 @@ class RunStore(abc.ABC):
|
|||||||
status: str,
|
status: str,
|
||||||
*,
|
*,
|
||||||
error: str | None = None,
|
error: str | None = None,
|
||||||
|
stop_reason: str | None = None,
|
||||||
) -> bool | None:
|
) -> bool | None:
|
||||||
"""Update a run status.
|
"""Update a run status.
|
||||||
|
|
||||||
|
|||||||
@ -45,6 +45,7 @@ class MemoryRunStore(RunStore):
|
|||||||
metadata=None,
|
metadata=None,
|
||||||
kwargs=None,
|
kwargs=None,
|
||||||
error=None,
|
error=None,
|
||||||
|
stop_reason=None,
|
||||||
created_at=None,
|
created_at=None,
|
||||||
owner_worker_id=None,
|
owner_worker_id=None,
|
||||||
lease_expires_at=None,
|
lease_expires_at=None,
|
||||||
@ -61,6 +62,7 @@ class MemoryRunStore(RunStore):
|
|||||||
"metadata": metadata or {},
|
"metadata": metadata or {},
|
||||||
"kwargs": kwargs or {},
|
"kwargs": kwargs or {},
|
||||||
"error": error,
|
"error": error,
|
||||||
|
"stop_reason": stop_reason,
|
||||||
"created_at": created_at or now,
|
"created_at": created_at or now,
|
||||||
"updated_at": now,
|
"updated_at": now,
|
||||||
"owner_worker_id": owner_worker_id,
|
"owner_worker_id": owner_worker_id,
|
||||||
@ -105,7 +107,7 @@ class MemoryRunStore(RunStore):
|
|||||||
thread_run_ids = self._runs_by_thread.get(thread_id) or ()
|
thread_run_ids = self._runs_by_thread.get(thread_id) or ()
|
||||||
return {run_id: run for run_id in thread_run_ids if run_id in run_ids and (run := self._runs.get(run_id)) is not None and (user_id is None or run.get("user_id") == user_id)}
|
return {run_id: run for run_id in thread_run_ids if run_id in run_ids and (run := self._runs.get(run_id)) is not None and (user_id is None or run.get("user_id") == user_id)}
|
||||||
|
|
||||||
async def update_status(self, run_id, status, *, error=None):
|
async def update_status(self, run_id, status, *, error=None, stop_reason=None):
|
||||||
run = self._runs.get(run_id)
|
run = self._runs.get(run_id)
|
||||||
if run is None:
|
if run is None:
|
||||||
return False
|
return False
|
||||||
@ -116,6 +118,8 @@ class MemoryRunStore(RunStore):
|
|||||||
run["status"] = status
|
run["status"] = status
|
||||||
if error is not None:
|
if error is not None:
|
||||||
run["error"] = error
|
run["error"] = error
|
||||||
|
if stop_reason is not None:
|
||||||
|
run["stop_reason"] = stop_reason
|
||||||
run["updated_at"] = datetime.now(UTC).isoformat()
|
run["updated_at"] = datetime.now(UTC).isoformat()
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|||||||
@ -517,6 +517,12 @@ async def run_agent(
|
|||||||
await subagent_events.add(chunk)
|
await subagent_events.add(chunk)
|
||||||
|
|
||||||
# 7. Stream the requested turn, then optionally continue hidden goal turns.
|
# 7. Stream the requested turn, then optionally continue hidden goal turns.
|
||||||
|
# Clear any stale stop_reason before the first (user-visible) turn only.
|
||||||
|
# Continuation turns preserve a cap reason from the user turn: a run that
|
||||||
|
# hits a cap during the user turn IS capped even if hidden goal-evaluator
|
||||||
|
# turns complete cleanly afterward (#4176 review).
|
||||||
|
if isinstance(runtime.context, dict):
|
||||||
|
runtime.context.pop("stop_reason", None)
|
||||||
await _stream_once(graph_input, initial_runnable_config)
|
await _stream_once(graph_input, initial_runnable_config)
|
||||||
while not record.abort_event.is_set() and not llm_error_fallback_message and (journal is None or not journal.had_llm_error_fallback):
|
while not record.abort_event.is_set() and not llm_error_fallback_message and (journal is None or not journal.had_llm_error_fallback):
|
||||||
continuation_input = await _prepare_goal_continuation_input(
|
continuation_input = await _prepare_goal_continuation_input(
|
||||||
@ -562,7 +568,22 @@ async def run_agent(
|
|||||||
error_msg = error_msg or "LLM provider failed after retries"
|
error_msg = error_msg or "LLM provider failed after retries"
|
||||||
await run_manager.set_status(run_id, RunStatus.error, error=error_msg)
|
await run_manager.set_status(run_id, RunStatus.error, error=error_msg)
|
||||||
else:
|
else:
|
||||||
await run_manager.set_status(run_id, RunStatus.success)
|
runtime_context = runtime.context if isinstance(runtime.context, dict) else None
|
||||||
|
# Guard middlewares that hard-stop a run by stripping tool_calls
|
||||||
|
# stamp stop_reason into runtime.context so the worker can surface
|
||||||
|
# it on the run record:
|
||||||
|
# loop_detection -> "loop_capped"
|
||||||
|
# token_budget -> "token_capped"
|
||||||
|
# safety_finish_reason -> "safety_capped"
|
||||||
|
# subagent_limit -> "subagent_limit_capped"
|
||||||
|
#
|
||||||
|
# If more guards grow stop_reason semantics, consider a publish/
|
||||||
|
# collect pattern (e.g. each guard middleware publishes its cap
|
||||||
|
# reason to a dedicated runtime.context channel, and the worker
|
||||||
|
# collects the most severe / first / all reasons) instead of each
|
||||||
|
# guard writing directly to the same key.
|
||||||
|
stop_reason = runtime_context.get("stop_reason") if runtime_context is not None else None
|
||||||
|
await run_manager.set_status(run_id, RunStatus.success, stop_reason=stop_reason)
|
||||||
|
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
await run_manager.set_finalizing(run_id, True)
|
await run_manager.set_finalizing(run_id, True)
|
||||||
|
|||||||
280
backend/tests/test_loop_detection_stop_reason.py
Normal file
280
backend/tests/test_loop_detection_stop_reason.py
Normal file
@ -0,0 +1,280 @@
|
|||||||
|
"""Integration tests: verify that guard middlewares write ``stop_reason`` to
|
||||||
|
``runtime.context`` and the worker surfaces it on the run record (#4176).
|
||||||
|
|
||||||
|
The lead worker calls ``agent.astream()``. During streaming, guard
|
||||||
|
middlewares (loop detection, token budget) may detect a cap and write
|
||||||
|
``stop_reason`` into ``runtime.context``. After streaming completes, the
|
||||||
|
worker reads ``runtime.context["stop_reason"]`` and persists it.
|
||||||
|
|
||||||
|
The key invariant: the middleware's ``runtime.context`` IS the worker's
|
||||||
|
``runtime.context`` — LangGraph surfaces the same dict — so the worker
|
||||||
|
sees whatever the middleware wrote.
|
||||||
|
|
||||||
|
These tests exercise that invariant end-to-end, using real middleware
|
||||||
|
instances (not hand-written simulations of the write) driven inside
|
||||||
|
``astream`` to prove the full middleware→context→worker pipeline.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langchain_core.messages import AIMessage
|
||||||
|
|
||||||
|
from deerflow.runtime import RunContext, RunManager, RunStatus
|
||||||
|
from deerflow.runtime.runs.worker import run_agent
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_worker_surfaces_stop_reason_from_loop_detection():
|
||||||
|
"""The worker persists ``stop_reason=loop_capped`` when the real
|
||||||
|
LoopDetectionMiddleware triggers a hard stop during streaming."""
|
||||||
|
from deerflow.agents.middlewares.loop_detection_middleware import LoopDetectionMiddleware
|
||||||
|
|
||||||
|
run_manager = RunManager()
|
||||||
|
record = await run_manager.create("thread-1")
|
||||||
|
|
||||||
|
mw = LoopDetectionMiddleware(warn_threshold=1, hard_limit=3, window_size=5)
|
||||||
|
captured_runtime: list[Any] = [None]
|
||||||
|
|
||||||
|
class DummyAgent:
|
||||||
|
metadata: dict[str, Any] = {"model_name": "test-model"}
|
||||||
|
|
||||||
|
async def astream(self, graph_input, config=None, stream_mode=None, subgraphs=False):
|
||||||
|
runtime = ((config or {}).get("configurable") or {}).get("__pregel_runtime")
|
||||||
|
assert runtime is not None, "LangGraph Runtime must be in configurable"
|
||||||
|
captured_runtime[0] = runtime
|
||||||
|
|
||||||
|
# Drive the real middleware to a hard stop with repeated identical
|
||||||
|
# tool calls. With hard_limit=3, the 3rd identical call fires the
|
||||||
|
# hard stop, triggering the runtime.context write.
|
||||||
|
tool_calls = [{"name": "bash", "args": {"command": "ls"}, "id": "c1", "type": "tool_call"}]
|
||||||
|
for _ in range(2):
|
||||||
|
mw._apply({"messages": [AIMessage(content="", tool_calls=tool_calls)]}, runtime)
|
||||||
|
# 3rd call — hard stop fires here.
|
||||||
|
mw._apply({"messages": [AIMessage(content="", tool_calls=tool_calls)]}, runtime)
|
||||||
|
|
||||||
|
yield {"messages": [AIMessage(content="Done.")]}
|
||||||
|
|
||||||
|
bridge = AsyncMock()
|
||||||
|
bridge.publish = AsyncMock()
|
||||||
|
bridge.publish_end = AsyncMock()
|
||||||
|
bridge.cleanup = AsyncMock()
|
||||||
|
|
||||||
|
def factory(*, config):
|
||||||
|
return DummyAgent()
|
||||||
|
|
||||||
|
await run_agent(
|
||||||
|
bridge,
|
||||||
|
run_manager,
|
||||||
|
record,
|
||||||
|
ctx=RunContext(checkpointer=None),
|
||||||
|
agent_factory=factory,
|
||||||
|
graph_input={"messages": []},
|
||||||
|
config={},
|
||||||
|
)
|
||||||
|
|
||||||
|
# Prove runtime object identity: the DummyAgent captured the same Runtime
|
||||||
|
# object the worker reads from; if LangGraph's merge created a copy, the
|
||||||
|
# worker would see a different context dict and stop_reason would be None.
|
||||||
|
assert captured_runtime[0] is not None, "DummyAgent never captured runtime"
|
||||||
|
runtime_ctx = captured_runtime[0].context
|
||||||
|
assert isinstance(runtime_ctx, dict)
|
||||||
|
assert runtime_ctx.get("stop_reason") == "loop_capped", "The runtime the DummyAgent wrote to is the same one the worker read from"
|
||||||
|
|
||||||
|
fetched = await run_manager.get(record.run_id)
|
||||||
|
assert fetched is not None
|
||||||
|
assert fetched.status == RunStatus.success
|
||||||
|
assert fetched.stop_reason == "loop_capped"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_worker_surfaces_stop_reason_from_token_budget():
|
||||||
|
"""The worker persists ``stop_reason=token_capped`` when the real
|
||||||
|
TokenBudgetMiddleware triggers a hard stop during streaming."""
|
||||||
|
from deerflow.agents.middlewares.token_budget_middleware import TokenBudgetMiddleware
|
||||||
|
from deerflow.config.token_budget_config import TokenBudgetConfig
|
||||||
|
|
||||||
|
run_manager = RunManager()
|
||||||
|
record = await run_manager.create("thread-1")
|
||||||
|
# Use a moderate budget with hard_stop_threshold=0.0 so even
|
||||||
|
# modest usage triggers the hard stop immediately.
|
||||||
|
config = TokenBudgetConfig(
|
||||||
|
enabled=True,
|
||||||
|
max_tokens=1000,
|
||||||
|
hard_stop_threshold=0.0,
|
||||||
|
warn_threshold=0.0,
|
||||||
|
)
|
||||||
|
mw = TokenBudgetMiddleware(config=config)
|
||||||
|
captured_runtime: list[Any] = [None]
|
||||||
|
|
||||||
|
class DummyAgent:
|
||||||
|
metadata: dict[str, Any] = {"model_name": "test-model"}
|
||||||
|
|
||||||
|
async def astream(self, graph_input, config=None, stream_mode=None, subgraphs=False):
|
||||||
|
runtime = ((config or {}).get("configurable") or {}).get("__pregel_runtime")
|
||||||
|
assert runtime is not None, "LangGraph Runtime must be in configurable"
|
||||||
|
captured_runtime[0] = runtime
|
||||||
|
|
||||||
|
# Feed a single AIMessage with token usage that exceeds the budget.
|
||||||
|
msg = AIMessage(
|
||||||
|
id="msg-budget",
|
||||||
|
content="hello",
|
||||||
|
usage_metadata={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150},
|
||||||
|
)
|
||||||
|
mw._apply({"messages": [msg]}, runtime)
|
||||||
|
|
||||||
|
yield {"messages": [AIMessage(content="Budget exceeded, wrapping up.")]}
|
||||||
|
|
||||||
|
bridge = AsyncMock()
|
||||||
|
bridge.publish = AsyncMock()
|
||||||
|
bridge.publish_end = AsyncMock()
|
||||||
|
bridge.cleanup = AsyncMock()
|
||||||
|
|
||||||
|
def factory(*, config):
|
||||||
|
return DummyAgent()
|
||||||
|
|
||||||
|
await run_agent(
|
||||||
|
bridge,
|
||||||
|
run_manager,
|
||||||
|
record,
|
||||||
|
ctx=RunContext(checkpointer=None),
|
||||||
|
agent_factory=factory,
|
||||||
|
graph_input={"messages": []},
|
||||||
|
config={},
|
||||||
|
)
|
||||||
|
# Prove runtime object identity (same rationale as the loop-detection test).
|
||||||
|
assert captured_runtime[0] is not None, "DummyAgent never captured runtime"
|
||||||
|
runtime_ctx = captured_runtime[0].context
|
||||||
|
assert isinstance(runtime_ctx, dict)
|
||||||
|
assert runtime_ctx.get("stop_reason") == "token_capped"
|
||||||
|
|
||||||
|
fetched = await run_manager.get(record.run_id)
|
||||||
|
assert fetched is not None
|
||||||
|
assert fetched.status == RunStatus.success
|
||||||
|
assert fetched.stop_reason == "token_capped"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_worker_surfaces_stop_reason_from_safety_finish_reason():
|
||||||
|
"""The worker persists ``stop_reason=safety_capped`` when the real
|
||||||
|
SafetyFinishReasonMiddleware strips tool_calls on a safety termination."""
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from deerflow.agents.middlewares.safety_finish_reason_middleware import SafetyFinishReasonMiddleware
|
||||||
|
from deerflow.agents.middlewares.safety_termination_detectors import SafetyTermination
|
||||||
|
|
||||||
|
# A detector that always fires, simulating any provider safety signal.
|
||||||
|
always_detector = MagicMock()
|
||||||
|
always_detector.name = "test-always-fire"
|
||||||
|
always_detector.detect.return_value = SafetyTermination(
|
||||||
|
detector="test-always-fire",
|
||||||
|
reason_field="finish_reason",
|
||||||
|
reason_value="content_filter",
|
||||||
|
)
|
||||||
|
mw = SafetyFinishReasonMiddleware(detectors=[always_detector])
|
||||||
|
captured_runtime: list[Any] = [None]
|
||||||
|
|
||||||
|
class DummyAgent:
|
||||||
|
metadata: dict[str, Any] = {"model_name": "test-model"}
|
||||||
|
|
||||||
|
async def astream(self, graph_input, config=None, stream_mode=None, subgraphs=False):
|
||||||
|
runtime = ((config or {}).get("configurable") or {}).get("__pregel_runtime")
|
||||||
|
assert runtime is not None
|
||||||
|
captured_runtime[0] = runtime
|
||||||
|
|
||||||
|
# Feed an AIMessage with tool_calls so the middleware triggers.
|
||||||
|
msg = AIMessage(
|
||||||
|
content="I can't do that.",
|
||||||
|
tool_calls=[{"name": "bash", "args": {}, "id": "c1", "type": "tool_call"}],
|
||||||
|
response_metadata={"finish_reason": "content_filter"},
|
||||||
|
)
|
||||||
|
mw._apply({"messages": [msg]}, runtime)
|
||||||
|
|
||||||
|
yield {"messages": [AIMessage(content="Safety filter triggered.")]}
|
||||||
|
|
||||||
|
run_manager = RunManager()
|
||||||
|
record = await run_manager.create("thread-1")
|
||||||
|
bridge = AsyncMock()
|
||||||
|
bridge.publish = AsyncMock()
|
||||||
|
bridge.publish_end = AsyncMock()
|
||||||
|
bridge.cleanup = AsyncMock()
|
||||||
|
|
||||||
|
await run_agent(
|
||||||
|
bridge,
|
||||||
|
run_manager,
|
||||||
|
record,
|
||||||
|
ctx=RunContext(checkpointer=None),
|
||||||
|
agent_factory=lambda *, config: DummyAgent(),
|
||||||
|
graph_input={"messages": []},
|
||||||
|
config={},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert captured_runtime[0] is not None
|
||||||
|
assert captured_runtime[0].context.get("stop_reason") == "safety_capped"
|
||||||
|
|
||||||
|
fetched = await run_manager.get(record.run_id)
|
||||||
|
assert fetched is not None
|
||||||
|
assert fetched.status == RunStatus.success
|
||||||
|
assert fetched.stop_reason == "safety_capped"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_worker_surfaces_stop_reason_from_subagent_limit():
|
||||||
|
"""The worker persists ``stop_reason=subagent_limit_capped`` when the
|
||||||
|
real SubagentLimitMiddleware hits the total per-run cap."""
|
||||||
|
from deerflow.agents.middlewares.subagent_limit_middleware import SubagentLimitMiddleware
|
||||||
|
|
||||||
|
# max_total=1: first delegation exhausts the cap.
|
||||||
|
mw = SubagentLimitMiddleware(max_concurrent=3, max_total=1)
|
||||||
|
captured_runtime: list[Any] = [None]
|
||||||
|
|
||||||
|
class DummyAgent:
|
||||||
|
metadata: dict[str, Any] = {"model_name": "test-model"}
|
||||||
|
|
||||||
|
async def astream(self, graph_input, config=None, stream_mode=None, subgraphs=False):
|
||||||
|
runtime = ((config or {}).get("configurable") or {}).get("__pregel_runtime")
|
||||||
|
assert runtime is not None
|
||||||
|
captured_runtime[0] = runtime
|
||||||
|
|
||||||
|
run_id = runtime.context.get("run_id")
|
||||||
|
# Simulate one prior delegation so remaining_total = 0.
|
||||||
|
state: dict[str, Any] = {
|
||||||
|
"messages": [
|
||||||
|
AIMessage(
|
||||||
|
content="Delegating...",
|
||||||
|
tool_calls=[{"name": "task", "args": {"subagent_type": "general-purpose"}, "id": "c1", "type": "tool_call"}],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
"delegations": [{"id": "prior-delegation", "run_id": run_id}],
|
||||||
|
}
|
||||||
|
mw._truncate_task_calls(state, runtime)
|
||||||
|
|
||||||
|
yield {"messages": [AIMessage(content="Subagent limit reached.")]}
|
||||||
|
|
||||||
|
run_manager = RunManager()
|
||||||
|
record = await run_manager.create("thread-1")
|
||||||
|
bridge = AsyncMock()
|
||||||
|
bridge.publish = AsyncMock()
|
||||||
|
bridge.publish_end = AsyncMock()
|
||||||
|
bridge.cleanup = AsyncMock()
|
||||||
|
|
||||||
|
await run_agent(
|
||||||
|
bridge,
|
||||||
|
run_manager,
|
||||||
|
record,
|
||||||
|
ctx=RunContext(checkpointer=None),
|
||||||
|
agent_factory=lambda *, config: DummyAgent(),
|
||||||
|
graph_input={"messages": []},
|
||||||
|
config={},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert captured_runtime[0] is not None
|
||||||
|
assert captured_runtime[0].context.get("stop_reason") == "subagent_limit_capped"
|
||||||
|
|
||||||
|
fetched = await run_manager.get(record.run_id)
|
||||||
|
assert fetched is not None
|
||||||
|
assert fetched.status == RunStatus.success
|
||||||
|
assert fetched.stop_reason == "subagent_limit_capped"
|
||||||
@ -156,7 +156,7 @@ async def test_migration_dedupes_duplicate_active_rows_before_unique_index(tmp_p
|
|||||||
|
|
||||||
with sqlite3.connect(db_path) as raw:
|
with sqlite3.connect(db_path) as raw:
|
||||||
version_row = raw.execute("SELECT version_num FROM alembic_version").fetchone()
|
version_row = raw.execute("SELECT version_num FROM alembic_version").fetchone()
|
||||||
assert version_row[0] == "0004_run_ownership"
|
assert version_row[0] == "0005_run_stop_reason"
|
||||||
|
|
||||||
# Sanity: the invariant the index enforces is now true — at most one
|
# Sanity: the invariant the index enforces is now true — at most one
|
||||||
# active row per thread.
|
# active row per thread.
|
||||||
|
|||||||
@ -1349,7 +1349,7 @@ async def test_cancel_returns_taken_over_when_peer_claims_during_local_cancel():
|
|||||||
# the race: in-memory cancel succeeds, but store write is blocked.
|
# the race: in-memory cancel succeeds, but store write is blocked.
|
||||||
original = store.update_status
|
original = store.update_status
|
||||||
|
|
||||||
async def race_update(run_id, status, *, error=None):
|
async def race_update(run_id, status, *, error=None, stop_reason=None):
|
||||||
# Simulate peer takeover: flip to error before our write lands
|
# Simulate peer takeover: flip to error before our write lands
|
||||||
run = store._runs.get(run_id)
|
run = store._runs.get(run_id)
|
||||||
if run and run["status"] == "running" and status == "interrupted":
|
if run and run["status"] == "running" and status == "interrupted":
|
||||||
@ -1357,7 +1357,7 @@ async def test_cancel_returns_taken_over_when_peer_claims_during_local_cancel():
|
|||||||
run["error"] = "peer takeover"
|
run["error"] = "peer takeover"
|
||||||
run["updated_at"] = datetime.now(UTC).isoformat()
|
run["updated_at"] = datetime.now(UTC).isoformat()
|
||||||
return False # our write was blocked
|
return False # our write was blocked
|
||||||
return await original(run_id, status, error=error)
|
return await original(run_id, status, error=error, stop_reason=stop_reason)
|
||||||
|
|
||||||
store.update_status = race_update
|
store.update_status = race_update
|
||||||
|
|
||||||
|
|||||||
@ -48,7 +48,7 @@ from deerflow.persistence.migrations._helpers import _normalize_default
|
|||||||
asyncio_test = pytest.mark.asyncio
|
asyncio_test = pytest.mark.asyncio
|
||||||
|
|
||||||
|
|
||||||
HEAD = "0004_run_ownership"
|
HEAD = "0005_run_stop_reason"
|
||||||
BASELINE = "0001_baseline"
|
BASELINE = "0001_baseline"
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -28,7 +28,7 @@ from deerflow.persistence.bootstrap import bootstrap_schema
|
|||||||
pytestmark = pytest.mark.asyncio
|
pytestmark = pytest.mark.asyncio
|
||||||
|
|
||||||
|
|
||||||
HEAD = "0004_run_ownership"
|
HEAD = "0005_run_stop_reason"
|
||||||
|
|
||||||
|
|
||||||
def _url(tmp_path: Path) -> str:
|
def _url(tmp_path: Path) -> str:
|
||||||
|
|||||||
@ -76,7 +76,7 @@ async def test_legacy_database_recovers_token_usage_column(tmp_path: Path) -> No
|
|||||||
cols = {row[1] for row in raw.execute("PRAGMA table_info(runs)").fetchall()}
|
cols = {row[1] for row in raw.execute("PRAGMA table_info(runs)").fetchall()}
|
||||||
assert "token_usage_by_model" in cols
|
assert "token_usage_by_model" in cols
|
||||||
version_row = raw.execute("SELECT version_num FROM alembic_version").fetchone()
|
version_row = raw.execute("SELECT version_num FROM alembic_version").fetchone()
|
||||||
assert version_row[0] == "0004_run_ownership"
|
assert version_row[0] == "0005_run_stop_reason"
|
||||||
|
|
||||||
# And the read path that originally 500'd must now succeed.
|
# And the read path that originally 500'd must now succeed.
|
||||||
sf = get_session_factory()
|
sf = get_session_factory()
|
||||||
@ -116,6 +116,6 @@ async def test_legacy_database_with_manual_alter_still_bootstraps(tmp_path: Path
|
|||||||
# No duplicate column -- list, not set, to catch dupes.
|
# No duplicate column -- list, not set, to catch dupes.
|
||||||
assert cols.count("token_usage_by_model") == 1
|
assert cols.count("token_usage_by_model") == 1
|
||||||
version_row = raw.execute("SELECT version_num FROM alembic_version").fetchone()
|
version_row = raw.execute("SELECT version_num FROM alembic_version").fetchone()
|
||||||
assert version_row[0] == "0004_run_ownership"
|
assert version_row[0] == "0005_run_stop_reason"
|
||||||
finally:
|
finally:
|
||||||
await close_engine()
|
await close_engine()
|
||||||
|
|||||||
@ -29,19 +29,19 @@ class FlakyStatusRunStore(MemoryRunStore):
|
|||||||
self.status_failures = status_failures
|
self.status_failures = status_failures
|
||||||
self.status_update_attempts = 0
|
self.status_update_attempts = 0
|
||||||
|
|
||||||
async def update_status(self, run_id, status, *, error=None):
|
async def update_status(self, run_id, status, *, error=None, stop_reason=None):
|
||||||
self.status_update_attempts += 1
|
self.status_update_attempts += 1
|
||||||
if self.status_failures > 0:
|
if self.status_failures > 0:
|
||||||
self.status_failures -= 1
|
self.status_failures -= 1
|
||||||
raise sqlite3.OperationalError("database is locked")
|
raise sqlite3.OperationalError("database is locked")
|
||||||
return await super().update_status(run_id, status, error=error)
|
return await super().update_status(run_id, status, error=error, stop_reason=stop_reason)
|
||||||
|
|
||||||
|
|
||||||
class MissingRowStatusRunStore(MemoryRunStore):
|
class MissingRowStatusRunStore(MemoryRunStore):
|
||||||
"""Memory run store that reports a missing row for status updates."""
|
"""Memory run store that reports a missing row for status updates."""
|
||||||
|
|
||||||
async def update_status(self, run_id, status, *, error=None):
|
async def update_status(self, run_id, status, *, error=None, stop_reason=None):
|
||||||
await super().update_status(run_id, status, error=error)
|
await super().update_status(run_id, status, error=error, stop_reason=stop_reason)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
@ -52,7 +52,7 @@ class PermanentStatusRunStore(MemoryRunStore):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.status_update_attempts = 0
|
self.status_update_attempts = 0
|
||||||
|
|
||||||
async def update_status(self, run_id, status, *, error=None):
|
async def update_status(self, run_id, status, *, error=None, stop_reason=None):
|
||||||
self.status_update_attempts += 1
|
self.status_update_attempts += 1
|
||||||
raise SQLAlchemyDatabaseError(
|
raise SQLAlchemyDatabaseError(
|
||||||
"UPDATE runs SET status = :status WHERE run_id = :run_id",
|
"UPDATE runs SET status = :status WHERE run_id = :run_id",
|
||||||
@ -68,7 +68,7 @@ class FailingStatusRunStore(MemoryRunStore):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.status_update_attempts = 0
|
self.status_update_attempts = 0
|
||||||
|
|
||||||
async def update_status(self, run_id, status, *, error=None):
|
async def update_status(self, run_id, status, *, error=None, stop_reason=None):
|
||||||
self.status_update_attempts += 1
|
self.status_update_attempts += 1
|
||||||
raise sqlite3.OperationalError("database is locked")
|
raise sqlite3.OperationalError("database is locked")
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user